mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Define kv cache scaling factor as amax / 448 (#790)
## What does this PR do? **Overview:** ? Unified the FP8 and NVFP4 kv cache scaling factor definition so the same checkpoint can be used for both FP8 and NVFP4 kv cache quantization deployment ## Testing Unit test ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes/No <!--- If No, explain why. --> - **Did you write any new necessary tests?**: Yes/No - **Did you add or update any necessary documentation?**: Yes/No - **Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**: Yes/No <!--- Only for new features, API changes, critical bug fixes or bw breaking changes. --> ## Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit ## Release Notes * **Refactor** * Fixed KV cache maximum bound to 448 for FP8 and NVFP4 quantization, simplifying configuration logic. * **Chores** * Removed internal constants from public exports. <sub>✏️ Tip: You can customize this high-level summary in your review settings.</sub> <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Chenjie Luo <108829653+cjluo-nv@users.noreply.github.com> Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
This commit is contained in:
@@ -56,9 +56,6 @@ from .layer_utils import (
|
||||
set_expert_quantizer_amax,
|
||||
)
|
||||
from .model_config import (
|
||||
KV_CACHE_FP8,
|
||||
KV_CACHE_NVFP4,
|
||||
KV_CACHE_NVFP4_AFFINE,
|
||||
QUANTIZATION_FP8,
|
||||
QUANTIZATION_FP8_PB_REAL,
|
||||
QUANTIZATION_FP8_PC_PT,
|
||||
@@ -647,19 +644,6 @@ def _export_transformers_checkpoint(
|
||||
|
||||
quant_config = get_quant_config(model, is_modelopt_qlora=is_modelopt_qlora)
|
||||
|
||||
kv_cache_max_bound = 0
|
||||
kv_cache_format = quant_config["quantization"]["kv_cache_quant_algo"]
|
||||
|
||||
cache_bound_mapping = {
|
||||
KV_CACHE_NVFP4: 6 * 448,
|
||||
KV_CACHE_NVFP4_AFFINE: 6 * 448,
|
||||
KV_CACHE_FP8: 448,
|
||||
}
|
||||
|
||||
# Only update kv_cache_max_bound if a quantization is applied.
|
||||
if kv_cache_format != QUANTIZATION_NONE:
|
||||
kv_cache_max_bound = cache_bound_mapping.get(kv_cache_format)
|
||||
|
||||
# Process all quantized modules and export weights
|
||||
_process_quantized_modules(model, dtype, is_modelopt_qlora)
|
||||
|
||||
@@ -669,6 +653,9 @@ def _export_transformers_checkpoint(
|
||||
else:
|
||||
quantized_state_dict = model.state_dict()
|
||||
|
||||
# We define kv cache scale as amax / 448 for both FP8 and NVFP4 KV cache quantization.
|
||||
kv_cache_max_bound = 448
|
||||
kv_cache_format = quant_config["quantization"]["kv_cache_quant_algo"]
|
||||
quantized_state_dict = postprocess_state_dict(
|
||||
quantized_state_dict, kv_cache_max_bound, kv_cache_format, is_modelopt_qlora
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user