mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do? Type of change: new feature. Follow-up to merged #2272. Adds composition of existing GEMM quantization with KV-cache AutoQuantize: - fixed FP8 GEMM PTQ followed by mixed-KV AutoQuantize; - gradient-based NVFP4/FP8 GEMM AutoQuantize followed by independent mixed-KV AutoQuantize; - an optional `kv_auto_quantize` recipe stage with independent method, constraints, candidates, and checkpoint path; - ordered `hf_ptq.py` orchestration that keeps selected weight/activation QDQ active while its calibration state remains frozen during KV candidate calibration; - fail-closed validation when a preceding stage leaves actual K/V quantizers enabled; and - unified export of a uniform-weight or mixed-weight checkpoint together with the selected per-layer KV map. The KV search still uses the public `mtq.auto_quantize(..., constraints={"cost_model": "kv_cache", ...})` API from #2272. On a converted model, the API preserves existing non-KV quantizers and requires K/V to be disabled before search. Fresh-model behavior is unchanged and starts from a deny-all quantizer baseline. #### Why a follow-up field instead of a generic stage list? This PR deliberately supports the two composition forms required by `hf_ptq.py` without replacing the stable recipe schema. Existing recipes already express a fixed `quantize` baseline plus one primary `auto_quantize` search. A generic ordered `stages` list would require a broader recipe/API migration, indexed checkpoint semantics, and compatibility rules for arbitrary stage sequences. There is not yet a demonstrated third search stage that justifies that surface-area change. The two searches are not combined inside `mtq.auto_quantize`: each invocation owns one search domain, constraint model, scoring method, and resumable checkpoint. Their ordering and independent checkpoint paths are orchestration concerns, while candidate calibration, scoring, selection, and state application remain in the shared public API. A general stage pipeline can be considered separately if more than this one optional KV follow-up is needed. Both solvers and scoring protocols are unchanged. The KV checkpoint compatibility signature additionally fingerprints the preceding quantizer configuration and calibrated state. Unsupported uniform-weight plus mixed-KV exports record `kv_cache_deployment_supported: false` in both ModelOpt and converted HF metadata. ### Usage Fixed FP8 GEMM PTQ followed by KV AutoQuantize: ```bash python examples/hf_ptq/hf_ptq.py \ --pyt_ckpt_path Qwen/Qwen3-8B \ --recipe general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \ --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \ --export_path /path/to/qwen3-8b-fp8-and-mixed-kv ``` Weight AutoQuantize followed by KV AutoQuantize: ```bash python examples/hf_ptq/hf_ptq.py \ --pyt_ckpt_path Qwen/Qwen3-8B \ --recipe general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \ --auto_quantize_checkpoint /path/to/weight_autoquant.pth \ --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \ --export_path /path/to/qwen3-8b-autoquant-and-mixed-kv ``` KV checkpoint resume requires identical preceding non-K/V quantizer configuration and calibrated state. If rerunning the preceding stage changes that state, use a new KV checkpoint path to recompute sensitivities; configuration identity alone is insufficient to reuse the scores safely. ### Testing - Latest changed-area validation: 126 tests passed across `hf_ptq.py` orchestration, KV checkpoint compatibility, export metadata, and HF configuration conversion. - A broader local run had 604 passes, one skip, and six failures: two socket-binding failures under the sandbox and four local Transformers API incompatibilities. This is not a full-suite pass. - The fixed-PTQ→KV recipe executes end to end on a tiny offline Qwen fixture. - Public API coverage verifies that composed KV search preserves preceding weight quantization and rejects enabled K/V state. - Changed-file pre-commit hooks passed; the isolated recipe validator also passed after dependency bootstrap. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A - Did you write any new necessary tests?: ✅ - Did you update Changelog?: ✅ (0.48.0 composition feature and KV checkpoint flag deprecation) - Did you get Claude approval on this PR?: ❌ ### Additional Information - This follow-up targets `main`, which contains merged #2272. - `--auto_quantize_checkpoint` and `--kv_auto_quantize_checkpoint` are intentionally separate because KV sensitivities depend on the preceding GEMM state. - Uniform-weight plus mixed-KV exports are for artifact inspection until the runtime's uniform-weight ModelOpt configuration consumes `kv_cache_quantized_layers`. Export emits an actionable warning and records `kv_cache_deployment_supported: false`; this marker does not itself add runtime support. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **New Features** - Added staged post-training quantization workflows for weights and KV caches, including dedicated KV-cache checkpoints. - Added FP8/NVFP4 recipes with configurable bit constraints and scoring. - KV-cache quantization now supports pre-quantized models. - **Bug Fixes** - Mixed weight and KV-cache quantization now exports with a warning instead of failing. - Improved validation and checkpoint compatibility for staged configurations. - Added safeguards for configurations without enabled weight quantizers. - **Documentation** - Clarified staged KV-cache workflows, checkpoint options, configuration behavior, and unsupported deployment combinations. - Documented deprecated legacy quantization options and their replacement behavior. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: weimingc <17592131+meenchen@users.noreply.github.com>
2033 lines
89 KiB
Python
Executable File
2033 lines
89 KiB
Python
Executable File
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
"""Utils for quantization including scaling factors adjustments."""
|
|
|
|
import fnmatch
|
|
import logging
|
|
from collections import defaultdict
|
|
from collections.abc import Generator
|
|
from typing import Any
|
|
from warnings import warn
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
|
|
from modelopt import __version__
|
|
from modelopt.torch.models import get_spec, list_all_possible
|
|
from modelopt.torch.quantization.ggml import IQ_FORMAT_REGISTRY
|
|
from modelopt.torch.quantization.model_calib import (
|
|
enable_stats_collection,
|
|
finish_stats_collection,
|
|
svd,
|
|
)
|
|
from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear
|
|
from modelopt.torch.quantization.qtensor import (
|
|
FP8QTensor,
|
|
MXFP4QTensor,
|
|
MXFP8QTensor,
|
|
NVFP4QTensor,
|
|
QTensorWrapper,
|
|
)
|
|
from modelopt.torch.quantization.utils import (
|
|
QuantizerAttrNames,
|
|
quantizer_attr_names,
|
|
representative_weight_quantizer,
|
|
weight_attr_names,
|
|
)
|
|
from modelopt.torch.utils import clear_cuda_cache
|
|
|
|
from ..quantization.nn import NVFP4StaticQuantizer, SequentialQuantizer, TensorQuantizer
|
|
from .model_utils import TiedWeightMap, get_language_model_from_vl
|
|
from .quant_format import (
|
|
IQ_FORMATS,
|
|
KV_CACHE_FP8,
|
|
KV_CACHE_FP8_K_NVFP4_V,
|
|
KV_CACHE_INT8,
|
|
KV_CACHE_NVFP4,
|
|
KV_CACHE_NVFP4_AFFINE,
|
|
QUANTIZATION_FP8,
|
|
QUANTIZATION_FP8_PB_REAL,
|
|
QUANTIZATION_FP8_PB_WO,
|
|
QUANTIZATION_FP8_PC_PT,
|
|
QUANTIZATION_INT4_AWQ,
|
|
QUANTIZATION_INT8_SQ,
|
|
QUANTIZATION_INT8_WO,
|
|
QUANTIZATION_MXFP4,
|
|
QUANTIZATION_MXFP8,
|
|
QUANTIZATION_NONE,
|
|
QUANTIZATION_NVFP4,
|
|
QUANTIZATION_NVFP4_AWQ,
|
|
QUANTIZATION_NVFP4_SVDQUANT,
|
|
QUANTIZATION_W4A8_AWQ,
|
|
QUANTIZATION_W4A8_MXFP4_FP8,
|
|
QUANTIZATION_W4A8_NVFP4_FP8,
|
|
QUANTIZATION_W4A16_NVFP4,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _has_large_fp8_scale(value: torch.Tensor) -> bool:
|
|
"""Return whether an FP8 scale exceeds the recommended threshold."""
|
|
return bool(torch.any(value > 0.5))
|
|
|
|
|
|
def maybe_transpose_expert_weight_dimensions(
|
|
weight: torch.Tensor,
|
|
weight_scale: torch.Tensor | None = None,
|
|
is_bmm_expert_weight: bool = True,
|
|
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
|
"""Transpose the last two dimensions of expert weights.
|
|
|
|
This function transposes expert weights between the two layouts:
|
|
- (num_experts, input_dim, output_dim) ↔ (num_experts, output_dim, input_dim)
|
|
|
|
Since transpose(-2, -1) is self-inverse, this function can be used for both
|
|
forward and backward transformations. This is needed for quantization functions
|
|
that expect the last dimension to be the input dimension for block quantization.
|
|
Specifically used for bmm-style expert weights in models like llama4 and gpt-oss.
|
|
|
|
Args:
|
|
weight: The weight tensor to transpose. Expected shape for experts: (num_experts, dim1, dim2)
|
|
weight_scale: Optional weight scaling factor tensor to transpose alongside weight
|
|
is_bmm_expert_weight: Whether this is an expert weight (3D tensor) that needs transposition
|
|
|
|
Returns:
|
|
Tuple of (transposed_weight, transposed_weight_scale)
|
|
"""
|
|
if not is_bmm_expert_weight or weight.dim() != 3:
|
|
return weight, weight_scale
|
|
|
|
transposed_weight = weight.transpose(-2, -1)
|
|
transposed_weight_scale = weight_scale.transpose(-2, -1) if weight_scale is not None else None
|
|
|
|
return transposed_weight, transposed_weight_scale
|
|
|
|
|
|
def adjust_attn_amax_values(module):
|
|
"""Adjusts the amax values for the attention layers."""
|
|
projection_prefixes = ["q", "k", "v"]
|
|
max_amax = float("-inf")
|
|
proj_layers = []
|
|
|
|
# Find all projection layers whose names contain 'q', 'k', or 'v'
|
|
for name, sub_module in module.named_children():
|
|
for prefix in projection_prefixes:
|
|
if (
|
|
prefix in name
|
|
and hasattr(sub_module, "weight_quantizer")
|
|
and hasattr(sub_module.weight_quantizer, "amax")
|
|
):
|
|
proj_layers.append(sub_module)
|
|
max_amax = max(max_amax, sub_module.weight_quantizer.amax.item())
|
|
|
|
if not proj_layers:
|
|
raise ValueError(
|
|
"No projection layers with the specified prefixes ('q', 'k', 'v') have amax attributes"
|
|
)
|
|
|
|
assert max_amax > 0, "max_amax must be positive."
|
|
|
|
# Set all amax values to the maximum found
|
|
for proj_layer in proj_layers:
|
|
proj_layer.weight_quantizer.amax.fill_(max_amax)
|
|
|
|
|
|
def get_scaling_factor(quantizer: TensorQuantizer) -> torch.Tensor:
|
|
"""Returns scaling factor from the quantizer as torch.Tensor."""
|
|
if not quantizer.is_enabled:
|
|
return None
|
|
|
|
amax = quantizer.export_amax()
|
|
if amax is None:
|
|
return None
|
|
|
|
# tensorrt_llm uses float as the scaling_factors.
|
|
if quantizer.num_bits == (2, 1):
|
|
scaling_factor = NVFP4QTensor.get_weights_scaling_factor_2_from_quantizer(quantizer)
|
|
else:
|
|
scaling_factor = amax.float() / quantizer.maxbound
|
|
|
|
assert torch.all(scaling_factor > 0), f"scaling factor {scaling_factor} not positive."
|
|
|
|
return scaling_factor
|
|
|
|
|
|
def get_activation_scaling_factor(
|
|
module: nn.Module, input_quantizer_name: str = "input_quantizer"
|
|
) -> torch.Tensor:
|
|
"""Returns the activation scaling factor."""
|
|
# If NVFP4, return activation scaling factor from NVFP4QTensor
|
|
input_quantizer = getattr(module, input_quantizer_name, None)
|
|
if input_quantizer is None:
|
|
return None
|
|
|
|
if get_quantization_format(module) in [
|
|
QUANTIZATION_NVFP4,
|
|
QUANTIZATION_NVFP4_AWQ,
|
|
QUANTIZATION_NVFP4_SVDQUANT,
|
|
]:
|
|
return NVFP4QTensor.get_activation_scaling_factor(input_quantizer)
|
|
return get_scaling_factor(input_quantizer)
|
|
|
|
|
|
def get_weight_scaling_factor(module: nn.Module, weight_name: str = "weight") -> torch.Tensor:
|
|
"""Returns the weight scaling factor."""
|
|
# module.weight_quantizer could be a TensorQuantizer (for algorithms except W4A8) or
|
|
# a SequentialQuantizer (for W4A8). In the latter case, we need to get the scaling factor from the
|
|
# first quantizer of the SequentialQuantizer instance.
|
|
|
|
weight: nn.Parameter = getattr(module, weight_name)
|
|
weight_quantizer: TensorQuantizer | SequentialQuantizer | None = getattr(
|
|
module, quantizer_attr_names(weight_name).weight_quantizer, None
|
|
)
|
|
|
|
if weight_quantizer is None:
|
|
return None
|
|
|
|
if isinstance(weight_quantizer, SequentialQuantizer):
|
|
return get_scaling_factor(weight_quantizer[0])
|
|
|
|
quantization_format = get_quantization_format(module)
|
|
|
|
if quantization_format in [
|
|
QUANTIZATION_NVFP4,
|
|
QUANTIZATION_NVFP4_AWQ,
|
|
QUANTIZATION_NVFP4_SVDQUANT,
|
|
QUANTIZATION_W4A16_NVFP4,
|
|
QUANTIZATION_W4A8_NVFP4_FP8,
|
|
]:
|
|
if quantization_format == QUANTIZATION_W4A8_NVFP4_FP8:
|
|
# weight_scaling_factor_2 for w4a8 needs to be amax/448, so that the wsf is in range 448/6.
|
|
# This is because the kernel dequantizes weight to fp8, which is in range 448.
|
|
weight_scaling_factor_2 = weight_quantizer._amax.float() / 448.0
|
|
else:
|
|
weight_scaling_factor_2 = NVFP4QTensor.get_weights_scaling_factor_2_from_quantizer(
|
|
weight_quantizer
|
|
)
|
|
# Unified method handles both static and dynamic quantizers
|
|
return NVFP4QTensor.get_weights_scaling_factor_from_quantizer(
|
|
weight_quantizer,
|
|
weight,
|
|
weight_scaling_factor_2.to(weight.device),
|
|
)[0]
|
|
|
|
if quantization_format in [QUANTIZATION_W4A8_MXFP4_FP8, QUANTIZATION_MXFP4]:
|
|
return MXFP4QTensor.quantize(weight, block_size=weight_quantizer.block_sizes[-1])[
|
|
1
|
|
].reshape(*weight.shape[:-1], -1)
|
|
|
|
if quantization_format == QUANTIZATION_MXFP8:
|
|
return MXFP8QTensor.get_weights_scaling_factor_from_quantizer(weight, weight_quantizer)
|
|
return get_scaling_factor(weight_quantizer)
|
|
|
|
|
|
def get_weight_scaling_factor_2(module: nn.Module, weight_name: str = "weight") -> torch.Tensor:
|
|
"""Returns the secondary weight scaling factor."""
|
|
weight_quantizer = getattr(module, quantizer_attr_names(weight_name).weight_quantizer, None)
|
|
|
|
if weight_quantizer is None:
|
|
return None
|
|
|
|
quantization_format = get_quantization_format(module)
|
|
|
|
if quantization_format in [
|
|
QUANTIZATION_NVFP4,
|
|
QUANTIZATION_NVFP4_AWQ,
|
|
QUANTIZATION_NVFP4_SVDQUANT,
|
|
QUANTIZATION_W4A16_NVFP4,
|
|
QUANTIZATION_W4A8_NVFP4_FP8,
|
|
]:
|
|
if quantization_format == QUANTIZATION_W4A8_NVFP4_FP8:
|
|
# weight_scaling_factor_2 for w4a8 needs to be amax/448, so that the wsf is in range 448/6.
|
|
# This is because the kernel dequantizes weight to fp8, which is in range 448.
|
|
return weight_quantizer._amax.float() / 448.0
|
|
else:
|
|
# Unified method handles both static and dynamic quantizers
|
|
return NVFP4QTensor.get_weights_scaling_factor_2_from_quantizer(weight_quantizer)
|
|
|
|
# SequentialQuantizer is required
|
|
if not isinstance(weight_quantizer, SequentialQuantizer) or not weight_quantizer[-1].is_enabled:
|
|
return None
|
|
|
|
assert len(weight_quantizer) == 2, (
|
|
"modelopt only supports 2 sequential quantization layers for now"
|
|
)
|
|
return get_scaling_factor(weight_quantizer[-1])
|
|
|
|
|
|
def get_prequant_scaling_factor(module: nn.Module) -> torch.Tensor:
|
|
"""Returns the prequant scaling factor."""
|
|
prequant_scaling_factor = (
|
|
module.input_quantizer._pre_quant_scale.squeeze()
|
|
if hasattr(module, "input_quantizer")
|
|
and hasattr(module.input_quantizer, "_pre_quant_scale")
|
|
else None
|
|
)
|
|
|
|
if prequant_scaling_factor is not None:
|
|
assert torch.all(prequant_scaling_factor > 0), (
|
|
f"prequant scaling factor {prequant_scaling_factor} not positive."
|
|
)
|
|
return prequant_scaling_factor
|
|
|
|
|
|
def get_kv_cache_bias(kv_module: nn.Module) -> list[torch.Tensor]:
|
|
"""Returns the kv_cache bias if _bias_value is set. Else returns None."""
|
|
kv_bias = []
|
|
for quantizer in ["k_bmm_quantizer", "v_bmm_quantizer"]:
|
|
quantizer_module = getattr(kv_module, quantizer, None)
|
|
kv_bias.append(getattr(quantizer_module, "_bias_value", None))
|
|
return kv_bias
|
|
|
|
|
|
def get_kv_cache_scaling_factor(
|
|
self_attention_module: nn.Module, clamp_fp8_scales: bool = True
|
|
) -> list[torch.Tensor | None]:
|
|
"""Get the K and V BMM scaling factors for the self attention module.
|
|
|
|
Args:
|
|
self_attention_module: The self attention module to get the K and V BMM scaling factors from.
|
|
clamp_fp8_scales: Whether to clamp FP8 KV cache scaling factors to at least 1.0.
|
|
|
|
Returns:
|
|
The K and V BMM scaling factors.
|
|
"""
|
|
if not hasattr(self_attention_module, "k_bmm_quantizer") or not hasattr(
|
|
self_attention_module, "v_bmm_quantizer"
|
|
):
|
|
return [None, None]
|
|
|
|
scaling_factors = [
|
|
get_scaling_factor(getattr(self_attention_module, quantizer))
|
|
for quantizer in ("k_bmm_quantizer", "v_bmm_quantizer")
|
|
]
|
|
|
|
# For FP8, we recommend default KV-cache scaling factor to be 1. The
|
|
# asymmetric format applies this only to K; V remains NVFP4.
|
|
kv_cache_dtype = get_kv_cache_dtype(self_attention_module)
|
|
if not clamp_fp8_scales:
|
|
fp8_indices = ()
|
|
elif kv_cache_dtype == KV_CACHE_FP8:
|
|
fp8_indices = range(len(scaling_factors))
|
|
elif kv_cache_dtype == KV_CACHE_FP8_K_NVFP4_V:
|
|
fp8_indices = (0,)
|
|
else:
|
|
fp8_indices = ()
|
|
for i in fp8_indices:
|
|
factor = scaling_factors[i]
|
|
if factor is None:
|
|
continue
|
|
if _has_large_fp8_scale(factor):
|
|
warn(
|
|
"Warning: Large KV activation detected. "
|
|
"Quantized KV cache may lead to higher accuracy drop."
|
|
)
|
|
scaling_factors[i] = torch.max(
|
|
factor, torch.tensor([1.0], dtype=torch.float, device=factor.device)
|
|
)
|
|
return scaling_factors
|
|
|
|
|
|
def get_kv_cache_dtype(modules: list[nn.Module] | nn.Module) -> str | None:
|
|
"""Returns the kv_cache dtype.
|
|
|
|
K/V quantizers are inspected as a pair so FP8 K with NVFP4 V remains
|
|
distinguishable from uniform FP8 or NVFP4. The output quantizer is retained
|
|
as a fallback for the unified Megatron export path.
|
|
|
|
Args:
|
|
modules: The module or list of modules to inspect.
|
|
|
|
Returns:
|
|
The kv_cache dtype.
|
|
"""
|
|
num_bits_list = []
|
|
is_affine = True
|
|
|
|
if isinstance(modules, nn.Module):
|
|
modules = [modules]
|
|
|
|
for module in modules:
|
|
k_quantizer = getattr(module, "k_bmm_quantizer", None)
|
|
v_quantizer = getattr(module, "v_bmm_quantizer", None)
|
|
if (
|
|
k_quantizer is not None
|
|
and v_quantizer is not None
|
|
and k_quantizer.is_enabled
|
|
and v_quantizer.is_enabled
|
|
):
|
|
k_dtype = _compute_kv_cache_dtype(
|
|
[k_quantizer.num_bits], hasattr(k_quantizer, "_bias_value")
|
|
)
|
|
v_dtype = _compute_kv_cache_dtype(
|
|
[v_quantizer.num_bits], hasattr(v_quantizer, "_bias_value")
|
|
)
|
|
if k_dtype == KV_CACHE_FP8 and v_dtype == KV_CACHE_NVFP4:
|
|
return KV_CACHE_FP8_K_NVFP4_V
|
|
if k_dtype == v_dtype:
|
|
return k_dtype
|
|
raise NotImplementedError(
|
|
"Unsupported mixed K/V cache quantization pair: "
|
|
f"K uses {k_dtype}, while V uses {v_dtype}."
|
|
)
|
|
|
|
# Case where the module has both k_bmm_quantizer and v_bmm_quantizer
|
|
# Still check for output quantizer for the unified_megatron_export path
|
|
for quantizer in ("k_bmm_quantizer", "v_bmm_quantizer", "output_quantizer"):
|
|
quantizer_attr = getattr(module, quantizer, None)
|
|
if quantizer_attr and quantizer_attr.is_enabled:
|
|
num_bits_list.append(quantizer_attr.num_bits)
|
|
is_affine &= hasattr(quantizer_attr, "_bias_value")
|
|
|
|
return _compute_kv_cache_dtype(num_bits_list, is_affine)
|
|
|
|
|
|
def _compute_kv_cache_dtype(
|
|
num_bits_list: list[int | tuple[int, int]], is_affine: bool = False
|
|
) -> str | None:
|
|
"""Returns the kv_cache dtype.
|
|
|
|
If num_bits of output_quantizer is (4, 3) then returns FP8; if it is 8, returns int8,
|
|
otherwise returns None.
|
|
|
|
Args:
|
|
num_bits_list: The list of num_bits from quantizers.
|
|
is_affine: Whether the quantizers have bias (affine mode).
|
|
|
|
Returns:
|
|
The kv_cache dtype.
|
|
"""
|
|
if (4, 3) in num_bits_list:
|
|
return KV_CACHE_FP8
|
|
elif 8 in num_bits_list:
|
|
return KV_CACHE_INT8
|
|
elif (2, 1) in num_bits_list and is_affine:
|
|
return KV_CACHE_NVFP4_AFFINE
|
|
elif (2, 1) in num_bits_list:
|
|
return KV_CACHE_NVFP4
|
|
else:
|
|
return QUANTIZATION_NONE
|
|
|
|
|
|
def get_weight_block_size(module: nn.Module, weight_name: str = "weight") -> int:
|
|
"""Returns the weight block size."""
|
|
weight_quantizer = representative_weight_quantizer(module, weight_name)
|
|
|
|
if weight_quantizer is None:
|
|
return 0
|
|
|
|
if isinstance(weight_quantizer, SequentialQuantizer):
|
|
weight_quantizer = weight_quantizer[0]
|
|
|
|
if not weight_quantizer.is_enabled:
|
|
return 0
|
|
|
|
block_sizes = weight_quantizer.block_sizes
|
|
|
|
if block_sizes:
|
|
return block_sizes[-1]
|
|
return 0
|
|
|
|
|
|
def uses_iq_quantization(module) -> bool:
|
|
"""Whether any weight quantizer in ``module`` or its children targets an IQ format.
|
|
|
|
``get_quantization_format`` returns the *first* non-``NONE`` format it finds, so in a
|
|
mixed-format model IQ layers sitting behind, say, an FP8 layer are invisible to it. Callers
|
|
that must reject IQ specifically need to see every layer.
|
|
|
|
This reads ``num_bits`` directly rather than resolving each layer's full format, so an
|
|
unrelated unsupported quantizer elsewhere in the model cannot turn the check into an error.
|
|
|
|
Known gap, shared with ``get_quantization_format``: ``weight_attr_names`` yields nothing for
|
|
a TEGroupedLinear, whose parameters are ``weight0..N`` while its quantizer is a single
|
|
``GroupedQuantizer`` under ``weight_quantizer``. Neither function sees such a module, so an
|
|
experts-only IQ model reports no format at all -- not just here. Closing it belongs in
|
|
``weight_attr_names``, where it affects every format, rather than in this helper.
|
|
"""
|
|
for weight_name in weight_attr_names(module):
|
|
weight_quantizer = representative_weight_quantizer(module, weight_name)
|
|
# getattr: a SequentialQuantizer has is_enabled but no num_bits, and is never IQ --
|
|
# IQ is a single quantizer with backend="ggml".
|
|
if (
|
|
weight_quantizer is not None
|
|
and weight_quantizer.is_enabled
|
|
and getattr(weight_quantizer, "num_bits", None) in IQ_FORMATS
|
|
):
|
|
return True
|
|
return any(uses_iq_quantization(child) for _, child in module.named_children())
|
|
|
|
|
|
def get_quantization_format(module) -> str | None:
|
|
"""Gets the quantization string.
|
|
|
|
Gets the quantization string by iterating through the module and its children.
|
|
The first non-None quantization string is returned.
|
|
"""
|
|
|
|
def _get_quantization_from_layer(layer, quantizer_attr_names: QuantizerAttrNames):
|
|
# Singular form first, plural ModuleList fallback (fused-experts).
|
|
# Strip the "_weight_quantizer" suffix to recover the weight attr name.
|
|
weight_attr = quantizer_attr_names.weight_quantizer
|
|
weight_name = weight_attr[: -len("_weight_quantizer")].rstrip("_") or "weight"
|
|
weight_quantizer = representative_weight_quantizer(layer, weight_name)
|
|
input_quantizer = getattr(layer, quantizer_attr_names.input_quantizer, None)
|
|
|
|
if weight_quantizer is None or not weight_quantizer.is_enabled:
|
|
return QUANTIZATION_NONE
|
|
|
|
# Handle SequentialQuantizer
|
|
if isinstance(weight_quantizer, SequentialQuantizer):
|
|
assert (
|
|
len(weight_quantizer) == 2
|
|
and weight_quantizer[0].num_bits == 4
|
|
and weight_quantizer[1].num_bits == (4, 3)
|
|
), "Unsupported SequentialQuantizer configuration"
|
|
assert (
|
|
weight_quantizer[0].block_sizes
|
|
and len(weight_quantizer[0].block_sizes) > 0
|
|
and weight_quantizer[0].block_sizes[-1] > 0
|
|
), "Invalid block_sizes for SequentialQuantizer"
|
|
|
|
return QUANTIZATION_W4A8_AWQ
|
|
|
|
# Handle individual num_bits cases
|
|
if weight_quantizer.num_bits in IQ_FORMATS:
|
|
if weight_quantizer.backend != "ggml":
|
|
raise ValueError("IQ formats require the built-in 'ggml' quantization backend")
|
|
# Both exporters return before collecting input_scale and before the pre_quant_scale
|
|
# handling below, so an enabled activation quantizer would be dropped without a trace
|
|
# and the checkpoint would load as weight-only. Refuse instead.
|
|
if input_quantizer is not None and input_quantizer.is_enabled:
|
|
raise NotImplementedError(
|
|
"IQ1_S/IQ2_XS export is weight-only, but this layer has an enabled input "
|
|
"quantizer. The GGML block payload carries no activation scale, so the "
|
|
"activation quantization would be silently lost."
|
|
)
|
|
if input_quantizer is not None and hasattr(input_quantizer, "_pre_quant_scale"):
|
|
raise NotImplementedError(
|
|
"IQ1_S/IQ2_XS export does not support an AWQ-style pre_quant_scale."
|
|
)
|
|
return weight_quantizer.num_bits
|
|
|
|
if weight_quantizer.num_bits == 4:
|
|
assert len(weight_quantizer.block_sizes) > 0 and weight_quantizer.block_sizes[-1] > 0, (
|
|
"Invalid block_sizes for INT4 quantizer"
|
|
)
|
|
return QUANTIZATION_INT4_AWQ
|
|
|
|
if weight_quantizer.num_bits == 8:
|
|
if input_quantizer is not None and input_quantizer.is_enabled:
|
|
return QUANTIZATION_INT8_SQ
|
|
else:
|
|
return QUANTIZATION_INT8_WO
|
|
|
|
if weight_quantizer.num_bits == (4, 3):
|
|
if weight_quantizer.block_sizes:
|
|
assert weight_quantizer.block_sizes[-1] > 0, "Invalid block_sizes for FP8 quantizer"
|
|
# Check if this is MXFP8 (dynamic block quantization with scale_bits (8, 0))
|
|
block_sizes = getattr(weight_quantizer, "block_sizes")
|
|
if (
|
|
isinstance(block_sizes, dict)
|
|
and block_sizes.get("type", "static") == "dynamic"
|
|
and block_sizes.get("scale_bits") == (8, 0)
|
|
):
|
|
return QUANTIZATION_MXFP8
|
|
if weight_quantizer.fake_quant:
|
|
return QUANTIZATION_FP8_PB_WO
|
|
else:
|
|
return QUANTIZATION_FP8_PB_REAL
|
|
if weight_quantizer.axis == 0:
|
|
return QUANTIZATION_FP8_PC_PT
|
|
return QUANTIZATION_FP8
|
|
|
|
if weight_quantizer.num_bits == (2, 1):
|
|
# FP4 formats are all block quantization
|
|
block_sizes = getattr(weight_quantizer, "block_sizes")
|
|
scale_bits = block_sizes.get("scale_bits")
|
|
|
|
if input_quantizer is not None and hasattr(weight_quantizer, "svdquant_lora_a"):
|
|
return QUANTIZATION_NVFP4_SVDQUANT
|
|
if input_quantizer is not None and hasattr(input_quantizer, "_pre_quant_scale"):
|
|
return QUANTIZATION_NVFP4_AWQ
|
|
if getattr(layer, "fused_with_prequant", False):
|
|
return QUANTIZATION_NVFP4_AWQ
|
|
if input_quantizer is None or not input_quantizer.is_enabled:
|
|
if scale_bits == (4, 3):
|
|
return QUANTIZATION_W4A16_NVFP4
|
|
assert input_quantizer is not None, (
|
|
f"input_quantizer is None for {quantizer_attr_names}"
|
|
)
|
|
if (
|
|
block_sizes.get("type", "static") == "dynamic"
|
|
and scale_bits == (8, 0)
|
|
and input_quantizer.is_enabled
|
|
and input_quantizer.num_bits == (4, 3)
|
|
and input_quantizer.block_sizes is None
|
|
):
|
|
return QUANTIZATION_W4A8_MXFP4_FP8
|
|
if (
|
|
block_sizes.get("type", "static") == "dynamic"
|
|
and scale_bits == (4, 3)
|
|
and input_quantizer.is_enabled
|
|
and input_quantizer.num_bits == (4, 3)
|
|
and input_quantizer.block_sizes is None
|
|
):
|
|
return QUANTIZATION_W4A8_NVFP4_FP8
|
|
if scale_bits == (4, 3):
|
|
return QUANTIZATION_NVFP4
|
|
elif scale_bits == (8, 0):
|
|
return QUANTIZATION_MXFP4
|
|
|
|
# Raise error for unsupported num_bits
|
|
raise NotImplementedError(
|
|
f"Unsupported quantizer with num_bits: {weight_quantizer.num_bits}"
|
|
)
|
|
|
|
for weight_name in weight_attr_names(module):
|
|
quantization = _get_quantization_from_layer(module, quantizer_attr_names(weight_name))
|
|
if quantization != QUANTIZATION_NONE:
|
|
return quantization
|
|
|
|
for _, layer in module.named_children():
|
|
format = get_quantization_format(layer)
|
|
if format != QUANTIZATION_NONE:
|
|
return format
|
|
|
|
return QUANTIZATION_NONE
|
|
|
|
|
|
def _prefix_wildcard_summarize_exclude_modules(unquantized_layers, quantized_layers):
|
|
"""Generate a summarization of the quantization layer configs using prefix wildcards.
|
|
|
|
Prefix wildcards means we only consider wildcards that is a prefix with a star in the end.
|
|
We do not consider other wildcards such as: a*b.
|
|
"""
|
|
|
|
def all_matching_prefix_wildcards(name):
|
|
# include all possible prefix wildcards, and the exact name itself
|
|
wildcards = {name}
|
|
for i in range(len(name) + 1):
|
|
wildcards.add(name[:i] + "*")
|
|
return wildcards
|
|
|
|
def next_formatted_matching_prefix_wildcards(name: str) -> Generator[list[str], None, None]:
|
|
"""Enumerate formatted prefix wildcards. A result may be a combination of prefix wildcards.
|
|
|
|
Formatted here means we only consider wildcards at dot split. We need two patterns.
|
|
|
|
1. a single wildcard: module_name*
|
|
2. a set of 2 wildcards: {module_name, module_name.*}. We need this pattern set because
|
|
module_name* may match other modules with module_name as a prefix.
|
|
"""
|
|
for i in range(len(name)):
|
|
if name[i] == ".":
|
|
yield [name[:i] + "*"]
|
|
yield [name[:i], name[:i] + ".*"]
|
|
# in the end, itself only is a wildcard
|
|
yield [name]
|
|
|
|
# any of the wildcard in this set cannot be present in the result
|
|
negative_wild_candidates = set()
|
|
for layer in quantized_layers:
|
|
negative = all_matching_prefix_wildcards(layer)
|
|
negative_wild_candidates.update(negative)
|
|
logger.debug(
|
|
f"Quantized layer {layer}, prefix wildcards {negative} identified as negative wildcards"
|
|
)
|
|
|
|
res_summary = set()
|
|
for layer in unquantized_layers:
|
|
candidate_wildcards = []
|
|
for wildcards in next_formatted_matching_prefix_wildcards(layer):
|
|
if any(wildcard in negative_wild_candidates for wildcard in wildcards):
|
|
# need a more specific wildcard
|
|
logger.debug(
|
|
f"Unquantized layer {layer}, prefix wildcards {wildcards} invalidated by negative wildcards"
|
|
)
|
|
continue
|
|
if all(wildcard in res_summary for wildcard in wildcards):
|
|
# we get covered already, do not need to move forward, and clear candidate
|
|
logger.debug(
|
|
f"Unquantized layer {layer}, prefix wildcards {wildcards} already covered"
|
|
)
|
|
candidate_wildcards = []
|
|
break
|
|
# find one, now terminate the search
|
|
candidate_wildcards = wildcards
|
|
logger.debug(
|
|
f"Unquantized layer {layer}, prefix wildcards {wildcards} identified as a new match"
|
|
)
|
|
break
|
|
# When candidate is the pair [prefix, prefix+".*"], emit only prefix+".*" for deployment.
|
|
if len(candidate_wildcards) == 2:
|
|
a, b = sorted(candidate_wildcards, key=len)
|
|
if b == a + ".*":
|
|
res_summary.add(b)
|
|
else:
|
|
res_summary.update(candidate_wildcards)
|
|
else:
|
|
res_summary.update(candidate_wildcards)
|
|
return res_summary
|
|
|
|
|
|
def process_layer_quant_config(layer_config_dict):
|
|
"""Processes per layer quantization information for TRTLLM export to quant_cfg.json."""
|
|
per_layer_config: dict[str, Any] = {
|
|
"quant_algo": None,
|
|
"kv_cache_quant_algo": None,
|
|
"quantized_layers": {},
|
|
}
|
|
layer_config: dict[str, Any] = {}
|
|
# Set of quantization formats used.
|
|
quantization_formats = set()
|
|
quantization_config = None
|
|
exclude_modules = []
|
|
|
|
for k, v in layer_config_dict.items():
|
|
if "awq_block_size" in k:
|
|
continue
|
|
|
|
# Get layer name for constructing quantized_layers dictionary under per_layer_config
|
|
prefix = ".".join(k.rsplit(".", 1)[:-1])
|
|
awq_key = prefix + ".awq_block_size"
|
|
|
|
# Get the corresponding AWQ block size
|
|
block_size_value = layer_config_dict.get(awq_key, 0)
|
|
|
|
if v == "fp8":
|
|
layer_config = {"quant_algo": "FP8"}
|
|
elif v == "fp8_pc_pt":
|
|
layer_config = {"quant_algo": "FP8_PER_CHANNEL_PER_TOKEN"}
|
|
elif v == "int4_awq":
|
|
layer_config = {
|
|
"quant_algo": "W4A16_AWQ",
|
|
"group_size": block_size_value,
|
|
"has_zero_point": False,
|
|
"pre_quant_scale": True,
|
|
}
|
|
elif v == "w4a8_awq":
|
|
layer_config = {
|
|
"quant_algo": "W4A8_AWQ",
|
|
"group_size": block_size_value,
|
|
"has_zero_point": False,
|
|
"pre_quant_scale": True,
|
|
}
|
|
elif v == "int8_wo":
|
|
layer_config = {"quant_algo": "W8A16"}
|
|
elif v == "int8_sq":
|
|
layer_config = {"quant_algo": "W8A8_SQ_PER_CHANNEL"}
|
|
elif v in ["nvfp4", "nvfp4_static"]:
|
|
layer_config = {
|
|
"quant_algo": "NVFP4",
|
|
"group_size": block_size_value,
|
|
}
|
|
elif v == "w4a16_nvfp4":
|
|
layer_config = {
|
|
"quant_algo": "W4A16_NVFP4",
|
|
"group_size": block_size_value,
|
|
}
|
|
elif v == "nvfp4_awq":
|
|
layer_config = {
|
|
"quant_algo": "NVFP4_AWQ",
|
|
"group_size": block_size_value,
|
|
"has_zero_point": False,
|
|
"pre_quant_scale": True,
|
|
}
|
|
elif v == "w4a8_nvfp4_fp8":
|
|
layer_config = {
|
|
"quant_algo": "W4A8_NVFP4_FP8",
|
|
"group_size": block_size_value,
|
|
}
|
|
elif v == "w4a8_mxfp4_fp8":
|
|
layer_config = {
|
|
"quant_algo": "W4A8_MXFP4_FP8",
|
|
"group_size": block_size_value,
|
|
}
|
|
elif v == "nvfp4_svdquant":
|
|
# SVDQuant builds on the AWQ-style pre_quant_scale smoothing, so its
|
|
# config mirrors nvfp4_awq (group_size + pre_quant_scale flag).
|
|
layer_config = {
|
|
"quant_algo": "NVFP4_SVD",
|
|
"group_size": block_size_value,
|
|
"has_zero_point": False,
|
|
"pre_quant_scale": True,
|
|
}
|
|
elif v == "mxfp8":
|
|
layer_config = {
|
|
"quant_algo": "MXFP8",
|
|
"group_size": block_size_value,
|
|
}
|
|
elif v in IQ_FORMATS:
|
|
iq_format = IQ_FORMAT_REGISTRY[v]
|
|
block_size, payload_bytes = iq_format.block_size, iq_format.block_bytes
|
|
effective_bits = iq_format.effective_bits
|
|
if block_size_value != block_size:
|
|
raise ValueError(
|
|
f"{v.upper()} requires block size {block_size}, got {block_size_value}"
|
|
)
|
|
layer_config = {
|
|
"quant_algo": v.upper(),
|
|
"group_size": block_size,
|
|
"effective_bits": effective_bits,
|
|
"block_payload_bytes": payload_bytes,
|
|
"packing": "ggml",
|
|
}
|
|
else:
|
|
layer_config = {"quant_algo": v}
|
|
|
|
if layer_config["quant_algo"] != QUANTIZATION_NONE:
|
|
quantization_formats.add(str(layer_config))
|
|
quantization_config = layer_config
|
|
per_layer_config["quantized_layers"].update({prefix: layer_config})
|
|
else:
|
|
exclude_modules.append(prefix)
|
|
|
|
# If we have more than one quantization format, infer MIXED_PRECISION
|
|
if len(quantization_formats) > 1:
|
|
per_layer_config["quant_algo"] = "MIXED_PRECISION"
|
|
elif len(quantization_formats) == 1 and quantization_config is not None:
|
|
per_layer_config.update(quantization_config)
|
|
per_layer_config["exclude_modules"] = sorted(
|
|
_prefix_wildcard_summarize_exclude_modules(
|
|
exclude_modules, per_layer_config["quantized_layers"].keys()
|
|
)
|
|
)
|
|
per_layer_config.pop("quantized_layers")
|
|
|
|
return per_layer_config
|
|
|
|
|
|
def _validate_int4_block_size(in_dim, block_size):
|
|
if not isinstance(block_size, int) or block_size <= 0:
|
|
raise ValueError(f"Block size must be a positive integer, got {block_size}.")
|
|
if in_dim % block_size != 0:
|
|
raise NotImplementedError(
|
|
f"Cannot pack weight with input dimension {in_dim} and block size {block_size}: "
|
|
"partial blocks are not supported."
|
|
)
|
|
|
|
|
|
def pack_int4_in_uint8(weight, weights_scaling_factor, block_size):
|
|
"""Packs the INT4 weights into uint8 tensor."""
|
|
out_dim = weight.shape[-2]
|
|
assert out_dim % 2 == 0, f"Cannot pack weight. Out dimension {out_dim} is not an even number."
|
|
in_dim = weight.shape[-1]
|
|
_validate_int4_block_size(in_dim, block_size)
|
|
|
|
# Scale, round, and clamp to the signed 4-bit range [-8..7].
|
|
int8_tensor = (
|
|
(weight / weights_scaling_factor[..., :, torch.arange(in_dim) // block_size])
|
|
.round()
|
|
.clamp(-8, 7)
|
|
.to(torch.int8)
|
|
)
|
|
|
|
# -- Handle the MoE (3D) case vs. the 2D case --
|
|
if int8_tensor.dim() == 3:
|
|
# Dimensions might be (experts, out_dim, in_dim)
|
|
transpose = int8_tensor.permute(0, 2, 1) # -> (experts, in_dim, out_dim)
|
|
# Reshape to group two output channels (out_dim // 2) and keep an extra dimension of size 2
|
|
transpose = transpose.reshape(-1, in_dim, out_dim // 2, 2) # (E, in_dim, out_dim//2, 2)
|
|
|
|
# Pack two 4-bit values (val0,val1) into a single byte:
|
|
val0 = transpose[..., 0] & 0x0F
|
|
val1 = transpose[..., 1] & 0x0F
|
|
packed_byte = val0 | (val1 << 4)
|
|
|
|
# Transpose back to the shape (experts, out_dim // 2, in_dim)
|
|
return packed_byte.permute(0, 2, 1).contiguous().view(torch.uint8)
|
|
|
|
else:
|
|
# 2D weights: shape typically (out_dim, in_dim)
|
|
# Transpose to (in_dim, out_dim)
|
|
reshaped = int8_tensor.T.reshape(in_dim, out_dim // 2, 2)
|
|
|
|
# Pack two 4-bit values into one byte
|
|
val0 = reshaped[..., 0] & 0x0F
|
|
val1 = reshaped[..., 1] & 0x0F
|
|
packed_byte = val0 | (val1 << 4)
|
|
|
|
# Return shape (out_dim // 2, in_dim)
|
|
return packed_byte.T.contiguous().view(torch.uint8)
|
|
|
|
|
|
def to_quantized_weight(
|
|
weight: torch.Tensor,
|
|
weights_scaling_factor: torch.Tensor,
|
|
quantization: str,
|
|
weights_scaling_factor2: torch.Tensor | None = None,
|
|
block_size: int | None = None,
|
|
):
|
|
"""Converts the weight to the quantized (packed) format."""
|
|
if weights_scaling_factor is not None:
|
|
weights_scaling_factor = weights_scaling_factor.to(weight.device)
|
|
|
|
if weights_scaling_factor2 is not None:
|
|
weights_scaling_factor2 = weights_scaling_factor2.to(weight.device)
|
|
|
|
# For compressed weights, we directly return the data from wrapper
|
|
if isinstance(weight, QTensorWrapper):
|
|
if quantization in [QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ]:
|
|
_validate_int4_block_size(weight.metadata["shape"][-1], block_size)
|
|
return weight.data
|
|
|
|
if quantization == QUANTIZATION_FP8:
|
|
# Fix RuntimeError: Promotion for Float8 Types is not supported, attempted to promote Float8_e4m3fn and Float
|
|
# in speculative decoding fp8 model export
|
|
if weight.dtype == torch.float8_e4m3fn:
|
|
warn("Skipping quantization: weight already in fp8 format")
|
|
return weight
|
|
|
|
if weight.dim() == 3:
|
|
# for MOE stacked weights
|
|
# Clear GPU cache to avoid pontential GPU OOM issues for large models.
|
|
clear_cuda_cache()
|
|
return (weight / weights_scaling_factor.unsqueeze(-1)).to(torch.float8_e4m3fn)
|
|
return (weight / weights_scaling_factor).to(torch.float8_e4m3fn)
|
|
|
|
if quantization in [QUANTIZATION_INT8_SQ, QUANTIZATION_INT8_WO]:
|
|
return (weight / weights_scaling_factor[:, None]).round().clamp(-128, 127).to(torch.int8)
|
|
|
|
if quantization == QUANTIZATION_MXFP8:
|
|
return MXFP8QTensor.quantize_with_scale(weight, weights_scaling_factor)
|
|
|
|
if quantization == QUANTIZATION_FP8_PB_WO:
|
|
return FP8QTensor.quantize(
|
|
weight, weights_scaling_factor.squeeze(), block_sizes={-1: block_size, -2: block_size}
|
|
)[0]._quantized_data
|
|
|
|
if quantization == QUANTIZATION_FP8_PC_PT:
|
|
if weight.dim() == 3:
|
|
# Handle different scale tensor shapes
|
|
if weights_scaling_factor.dim() == 1:
|
|
# Per-expert scaling only: (num_experts,) -> (num_experts, 1, 1)
|
|
return (weight / weights_scaling_factor[:, None, None]).to(torch.float8_e4m3fn)
|
|
elif weights_scaling_factor.dim() == 2:
|
|
# Per-channel scaling: check which dimension matches
|
|
if weights_scaling_factor.shape[0] != weight.shape[0]:
|
|
raise ValueError(
|
|
f"First dimension (num_experts) mismatch for FP8_PC_PT quantization. "
|
|
f"weight shape: {weight.shape}, scale shape: {weights_scaling_factor.shape}"
|
|
)
|
|
if weight.shape[-1] == weight.shape[-2]:
|
|
raise ValueError(
|
|
f"Ambiguous scaling dimension for FP8_PC_PT quantization with square weight matrix. "
|
|
f"weight shape: {weight.shape}, scale shape: {weights_scaling_factor.shape}. "
|
|
f"Cannot determine if scaling should be applied to input_dim or output_dim."
|
|
)
|
|
if weights_scaling_factor.shape[-1] == weight.shape[-1]:
|
|
# (num_experts, input_dim) -> (num_experts, 1, input_dim), BMM-style
|
|
return (weight / weights_scaling_factor.unsqueeze(-2)).to(torch.float8_e4m3fn)
|
|
elif weights_scaling_factor.shape[-1] == weight.shape[-2]:
|
|
# (num_experts, output_dim) -> (num_experts, output_dim, 1), Standard MoE case
|
|
return (weight / weights_scaling_factor.unsqueeze(-1)).to(torch.float8_e4m3fn)
|
|
else:
|
|
raise ValueError(
|
|
f"Cannot determine correct unsqueeze dimension for FP8_PC_PT quantization. "
|
|
f"weight shape: {weight.shape}, scale shape: {weights_scaling_factor.shape}"
|
|
)
|
|
return (weight / weights_scaling_factor[:, None]).to(torch.float8_e4m3fn)
|
|
|
|
if quantization in [QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ]:
|
|
return pack_int4_in_uint8(weight, weights_scaling_factor, block_size)
|
|
|
|
if quantization in [
|
|
QUANTIZATION_NVFP4,
|
|
QUANTIZATION_NVFP4_AWQ,
|
|
QUANTIZATION_W4A16_NVFP4,
|
|
QUANTIZATION_W4A8_NVFP4_FP8,
|
|
QUANTIZATION_NVFP4_SVDQUANT,
|
|
]:
|
|
assert block_size is not None, "Block size not passed. Unable to quantize to NVFP4 format."
|
|
assert weights_scaling_factor2 is not None, (
|
|
"Weights scaling factor 2 not passed. Unable to quantize to NVFP4 format"
|
|
)
|
|
# If MoE reshape weights_scaling_factor2 to enable quantize operations
|
|
return NVFP4QTensor.quantize(
|
|
weight,
|
|
block_size,
|
|
weights_scaling_factor,
|
|
weights_scaling_factor2.view(-1, 1, 1)
|
|
if weights_scaling_factor2.dim() != 0
|
|
else weights_scaling_factor2,
|
|
)[0]._quantized_data
|
|
|
|
if quantization in [QUANTIZATION_W4A8_MXFP4_FP8, QUANTIZATION_MXFP4]:
|
|
return MXFP4QTensor.quantize(weight, block_size=block_size)[0]._quantized_data
|
|
|
|
raise NotImplementedError(f"quantization format {quantization} not supported")
|
|
|
|
|
|
def from_quantized_weight(
|
|
weight: torch.Tensor,
|
|
weights_scaling_factor: torch.Tensor,
|
|
quantization: str,
|
|
torch_dtype,
|
|
):
|
|
"""Converts the quantized weight to the target torch_dtype format."""
|
|
if weight.element_size() >= 2 or weights_scaling_factor is None or not quantization:
|
|
# No need to unquantize the weight.
|
|
return weight.to(torch_dtype)
|
|
|
|
if quantization == QUANTIZATION_FP8:
|
|
# safe tensors does not support fp8 yet. So we pack the tensors as int8
|
|
return weight.view(torch.float8_e4m3fn).to(torch_dtype) * weights_scaling_factor.to(
|
|
torch_dtype
|
|
)
|
|
|
|
if quantization in [QUANTIZATION_INT8_SQ, QUANTIZATION_INT8_WO]:
|
|
return weight.to(torch_dtype) * weights_scaling_factor[:, None].to(torch_dtype)
|
|
|
|
raise NotImplementedError(f"quantization format {quantization} not supported")
|
|
|
|
|
|
_KV_CACHE_REPLACEMENTS: dict[str, str] = {
|
|
"k_bmm_quantizer._amax": "k_proj.k_scale",
|
|
"v_bmm_quantizer._amax": "v_proj.v_scale",
|
|
"k_bmm_quantizer._bias_value": "k_proj.k_bias",
|
|
"v_bmm_quantizer._bias_value": "v_proj.v_bias",
|
|
"input_quantizer._pre_quant_scale": "pre_quant_scale",
|
|
}
|
|
_BASE_SKIP_KEYS: tuple[str, ...] = (
|
|
"output_quantizer",
|
|
"_amax",
|
|
"_bias_value",
|
|
"input_quantizer._pre_quant_scale",
|
|
"weight_shape",
|
|
)
|
|
|
|
|
|
def _strip_base_layer(key: str, is_modelopt_qlora: bool) -> str:
|
|
"""Drop the `base_layer` component PEFT inserts, which deployment does not expect.
|
|
|
|
Stripping generically means new key types (bias, scales) need no enumeration here.
|
|
"""
|
|
return key.replace(".base_layer.", ".") if is_modelopt_qlora else key
|
|
|
|
|
|
def _maybe_squeeze_scale(key: str, value: Any) -> Any:
|
|
"""Squeeze a leading dim=1 from 3-D scale tensors of shape (1, n, m)."""
|
|
if (
|
|
"scale" in key
|
|
and isinstance(value, torch.Tensor)
|
|
and value.dim() == 3
|
|
and value.shape[0] == 1
|
|
):
|
|
return value.squeeze(0)
|
|
return value
|
|
|
|
|
|
def _postprocess_single_tensor(
|
|
key: str,
|
|
value: torch.Tensor,
|
|
kv_cache_max_bound: float,
|
|
kv_cache_format: str | dict[str, dict[str, str]] | None,
|
|
is_modelopt_qlora: bool = False,
|
|
) -> tuple[str | None, torch.Tensor | None]:
|
|
"""Per-tensor subset of :func:`postprocess_state_dict`, for streaming export.
|
|
|
|
Returns ``(new_key, new_value)`` to emit, or ``(None, None)`` to skip.
|
|
Tied-weight dedup is NOT performed here; callers should pre-compute alias
|
|
keys from ``model._tied_weights_keys`` and filter them at the call site.
|
|
"""
|
|
replacements = _KV_CACHE_REPLACEMENTS
|
|
skip_keys = _BASE_SKIP_KEYS
|
|
|
|
# Skip problematic VL model parameters
|
|
if key == "vision_model.radio_model.summary_idxs":
|
|
return None, None
|
|
|
|
# Skip real quant parameters
|
|
if any(key.endswith("weight_quantizer." + q) for q in RealQuantLinear.list_of_scale_tensors):
|
|
return None, None
|
|
|
|
# Skip LoRA adapters for QLoRA models
|
|
if is_modelopt_qlora and "lora" in key:
|
|
return None, None
|
|
|
|
# Keys not related to quantizers: keep as-is
|
|
if all(sk not in key for sk in skip_keys):
|
|
new_key = _strip_base_layer(key, is_modelopt_qlora)
|
|
return new_key, _maybe_squeeze_scale(new_key, value)
|
|
|
|
# Apply replacements if the key matches any suffix in the replacements dict
|
|
for old_suffix, new_suffix in replacements.items():
|
|
if key.endswith(old_suffix):
|
|
prefix = key[: -len(old_suffix)]
|
|
if "_amax" in key:
|
|
layer_quantization = _resolve_kv_cache_format_for_key(key, kv_cache_format)
|
|
assert layer_quantization in [
|
|
KV_CACHE_FP8,
|
|
KV_CACHE_NVFP4,
|
|
KV_CACHE_NVFP4_AFFINE,
|
|
], "Invalid KV cache quantization format."
|
|
assert kv_cache_max_bound > 0, "Maxbound must be greater than zero."
|
|
value = value.float() / kv_cache_max_bound
|
|
if layer_quantization == KV_CACHE_FP8 and _has_large_fp8_scale(value):
|
|
logger.warning(
|
|
"Large KV activations detected. Quantized KV cache may lead to higher accuracy drop."
|
|
)
|
|
new_key = _strip_base_layer(prefix + new_suffix, is_modelopt_qlora)
|
|
return new_key, _maybe_squeeze_scale(new_key, value)
|
|
|
|
# Key has a skip_key but no replacement matched — drop it
|
|
return None, None
|
|
|
|
|
|
def _resolve_kv_cache_format_for_key(
|
|
key: str, quantization: str | dict[str, dict[str, str]] | None
|
|
) -> str | None:
|
|
"""Resolve uniform or per-layer metadata to the format for this K or V tensor."""
|
|
if isinstance(quantization, dict):
|
|
matches = [
|
|
(layer_name, layer_config.get("quant_algo"))
|
|
for layer_name, layer_config in quantization.items()
|
|
if key == layer_name or key.startswith(layer_name + ".")
|
|
]
|
|
quantization = max(matches, key=lambda item: len(item[0]))[1] if matches else None
|
|
if quantization == KV_CACHE_FP8_K_NVFP4_V:
|
|
if key.endswith("k_bmm_quantizer._amax"):
|
|
return KV_CACHE_FP8
|
|
if key.endswith("v_bmm_quantizer._amax"):
|
|
return KV_CACHE_NVFP4
|
|
return None
|
|
return quantization
|
|
|
|
|
|
def _get_kv_cache_postprocess_config(
|
|
quantization_details: dict[str, Any],
|
|
) -> str | dict[str, dict[str, str]] | None:
|
|
"""Return the uniform format or layer map consumed by both HF exporters."""
|
|
kv_cache_format = quantization_details.get("kv_cache_quant_algo")
|
|
if kv_cache_format == "MIXED_PRECISION":
|
|
return quantization_details.get("kv_cache_quantized_layers", {})
|
|
return kv_cache_format
|
|
|
|
|
|
def postprocess_state_dict(
|
|
state_dict: dict,
|
|
maxbound: float,
|
|
quantization: str | dict[str, dict[str, str]] | None,
|
|
is_modelopt_qlora: bool = False,
|
|
tied_map: "TiedWeightMap | None" = None,
|
|
) -> dict:
|
|
"""Filters out keys related to weight quantizers and updates KV cache related keys.
|
|
|
|
Args:
|
|
state_dict: The full model state_dict.
|
|
maxbound: The maximum bound value for the output quantizer.
|
|
quantization: The uniform KV cache quantization format, or a per-attention-layer
|
|
``{layer_name: {"quant_algo": ...}}`` mapping for mixed precision.
|
|
is_modelopt_qlora: Whether the model is a modelopt-trained QLoRA model.
|
|
tied_map: Optional :class:`TiedWeightMap`. When provided, tied-weight
|
|
dedup is authoritative and name-based: a declared alias key whose canonical
|
|
counterpart is present is dropped, independent of tensor address. This is
|
|
what makes dedup correct under the FSDP full-state-dict gather (and offload),
|
|
where tied tensors are materialized at distinct addresses and the address
|
|
pass below cannot see the tie. The address pass is retained as a backstop for
|
|
undeclared genuine shares and coincidental collisions.
|
|
|
|
Returns:
|
|
The filtered state_dict without unnecessary keys like '_amax' and non KV cache output quantizers.
|
|
"""
|
|
replacements = _KV_CACHE_REPLACEMENTS
|
|
skip_keys = _BASE_SKIP_KEYS
|
|
|
|
def _export_key(key: str) -> str:
|
|
return _strip_base_layer(key, is_modelopt_qlora)
|
|
|
|
post_state_dict = {}
|
|
|
|
for key, value in state_dict.items():
|
|
# Skip problematic parameters for specific model architectures, e.g., Nemotron Nano VL models
|
|
if key == "vision_model.radio_model.summary_idxs":
|
|
logger.info(f"Removing problematic parameter: {key}")
|
|
continue
|
|
|
|
# Skip keys not related to quantizers
|
|
if all(skip_key not in key for skip_key in skip_keys):
|
|
post_state_dict[_export_key(key)] = value
|
|
continue
|
|
|
|
# Apply replacements if the key matches any suffix in the replacements dict
|
|
for old_suffix, new_suffix in replacements.items():
|
|
if key.endswith(old_suffix):
|
|
prefix = key[: -len(old_suffix)]
|
|
|
|
if "_amax" in key:
|
|
layer_quantization = _resolve_kv_cache_format_for_key(key, quantization)
|
|
assert layer_quantization in [
|
|
KV_CACHE_FP8,
|
|
KV_CACHE_NVFP4,
|
|
KV_CACHE_NVFP4_AFFINE,
|
|
], "Invalid KV cache quantization format."
|
|
assert maxbound > 0, "Maxbound must be greater than zero."
|
|
|
|
value = value.float() / maxbound
|
|
|
|
# Warn if scale exceeds threshold
|
|
if layer_quantization == KV_CACHE_FP8 and _has_large_fp8_scale(value):
|
|
logger.warning(
|
|
"Large KV activations detected. Quantized KV cache may lead to higher accuracy drop."
|
|
)
|
|
post_state_dict[_export_key(prefix + new_suffix)] = value
|
|
break
|
|
|
|
post_state_dict = {k: _maybe_squeeze_scale(k, v) for k, v in post_state_dict.items()}
|
|
|
|
# remove real quant parameters from the state dict
|
|
keys_to_delete = []
|
|
for key, value in post_state_dict.items():
|
|
if any(
|
|
key.endswith("weight_quantizer." + q_key)
|
|
for q_key in RealQuantLinear.list_of_scale_tensors
|
|
):
|
|
keys_to_delete.append(key)
|
|
|
|
# remove LoRA adapters from state dict
|
|
if is_modelopt_qlora:
|
|
for key in post_state_dict:
|
|
if "lora" in key and key not in keys_to_delete:
|
|
keys_to_delete.append(key)
|
|
# Name-based tied-weight dedup (authoritative, address-independent): for each declared
|
|
# {alias: canonical} tie, drop the alias's own exported keys when the canonical is present.
|
|
# Per-parameter (not per-prefix), so an untied sibling (a bias, an untied projection) is kept.
|
|
# This is what makes it correct under the FSDP gather / offload, where tied tensors land at
|
|
# distinct addresses.
|
|
dropped_dense_prefixes: set[str] = set()
|
|
if tied_map is not None and tied_map.alias_to_canonical:
|
|
# Expand each tied *pre-pack* parameter name into the concrete keys export produced. If
|
|
# export starts emitting a new scale companion or fused-projection name, extend
|
|
# `weight_suffixes` / `proj_splits` here -- nothing else needs to change.
|
|
# (pre_quant_scale is the AWQ / NVFP4_AWQ / SVDQuant companion, renamed in the KV-cache pass.)
|
|
weight_suffixes = (
|
|
"weight",
|
|
"weight_scale",
|
|
"weight_scale_2",
|
|
"input_scale",
|
|
"pre_quant_scale",
|
|
)
|
|
# A tied 3-D fused projection splits into these per-expert 2-D projection names.
|
|
proj_splits = {
|
|
"gate_up_proj": ("gate_proj", "up_proj"),
|
|
"up_proj": ("up_proj",),
|
|
"down_proj": ("down_proj",),
|
|
}
|
|
|
|
# group name -> [(alias_key, canonical_key), ...], dropped all-or-none.
|
|
alias_groups: dict[str, list[tuple[str, str]]] = {}
|
|
# dense group name -> module prefix, for the leftover-companion check below.
|
|
dense_prefixes: dict[str, str] = {}
|
|
# fused-MoE container (alias_prefix, canonical_prefix) -> its tied per-expert projection names.
|
|
moe_containers: dict[tuple[str, str], set[str]] = {}
|
|
|
|
for alias, canonical in tied_map.alias_to_canonical.items():
|
|
a_pre, _, a_name = alias.rpartition(".")
|
|
c_pre, _, c_name = canonical.rpartition(".")
|
|
if a_name == c_name == "weight":
|
|
# Dense tie: the weight and its own scale companions only (never a sibling bias).
|
|
members = [
|
|
(f"{a_pre}.{s}" if a_pre else s, f"{c_pre}.{s}" if c_pre else s)
|
|
for s in weight_suffixes
|
|
]
|
|
members = [(ak, ck) for ak, ck in members if ak in post_state_dict]
|
|
if members:
|
|
alias_groups[alias] = members
|
|
dense_prefixes[alias] = a_pre
|
|
elif a_name in proj_splits and a_name == c_name:
|
|
# Fused-MoE: export splits the 3-D container into per-expert keys; record those names.
|
|
moe_containers.setdefault((a_pre, c_pre), set()).update(proj_splits[a_name])
|
|
elif alias in post_state_dict:
|
|
alias_groups[alias] = [(alias, canonical)]
|
|
|
|
for (a_pre, c_pre), tied_proj_names in moe_containers.items():
|
|
# Rewrite only the tied projections' per-expert keys, so an untied projection (e.g.
|
|
# down_proj when only gate_up_proj is tied) or a router/bias child is left alone.
|
|
prefix = f"{a_pre}." if a_pre else ""
|
|
members = [
|
|
(key, c_pre + key[len(a_pre) :])
|
|
for key in post_state_dict
|
|
if key.startswith(prefix)
|
|
and any(part in tied_proj_names for part in key[len(prefix) :].split("."))
|
|
]
|
|
if members:
|
|
# Key by the full (alias, canonical) pair: one alias container's projections may
|
|
# tie to different canonical containers, and keying by a_pre alone would overwrite.
|
|
alias_groups[f"{a_pre} -> {c_pre}"] = members
|
|
|
|
# Atomic drop: remove a group only when every alias key has its canonical twin present,
|
|
# so mixed quant state (one side quantized, the other not) never orphans a scale.
|
|
for gk, members in alias_groups.items():
|
|
missing = [ak for ak, ck in members if ck not in post_state_dict]
|
|
if missing:
|
|
logger.warning(
|
|
f"Skipping name-based dedup of '{gk}': tied sides have mismatched keys "
|
|
f"(e.g. '{missing[0]}' has no canonical counterpart, likely differing "
|
|
f"quantization state); keeping both sides to avoid orphaned tensors."
|
|
)
|
|
continue
|
|
# Safety: a declared tie must export identical bytes on both sides. If quantization
|
|
# diverged them, dropping the alias would silently corrupt it (HF re-ties canonical over
|
|
# it on load) -- so raise. Mirrors HF's own torch.equal decline-to-tie.
|
|
for ak, ck in members:
|
|
av, cv = post_state_dict[ak], post_state_dict[ck]
|
|
if (
|
|
isinstance(av, torch.Tensor)
|
|
and isinstance(cv, torch.Tensor)
|
|
and not av.is_meta
|
|
and not cv.is_meta
|
|
and not torch.equal(av, cv)
|
|
):
|
|
raise RuntimeError(
|
|
f"Tied-weight export mismatch: '{ak}' differs from its canonical '{ck}'. "
|
|
f"The tie is declared but the two sides quantized to different values, so "
|
|
f"deduplicating would corrupt '{ak}'. Ensure both tied modules use the same "
|
|
f"quantization format/config."
|
|
)
|
|
for ak, _ in members:
|
|
keys_to_delete.append(ak)
|
|
if gk in dense_prefixes:
|
|
dropped_dense_prefixes.add(dense_prefixes[gk])
|
|
logger.warning(
|
|
f"Tied weight (declared): dropping {len(members)} alias key(s) for '{gk}'; "
|
|
f"canonical kept."
|
|
)
|
|
|
|
# Apply the name drops first so the address pass below runs on the reduced dict and never
|
|
# re-processes a declared alias. dict.fromkeys dedups the delete list while preserving order.
|
|
for key in dict.fromkeys(keys_to_delete):
|
|
post_state_dict.pop(key, None)
|
|
|
|
# Address backstop (pre-existing, unchanged): drop any remaining keys that still share a
|
|
# data_ptr. Declared ties are already gone above, so this only catches an undeclared
|
|
# same-address share (an unquantized/unpacked tie); a quantized tie packs to distinct storage
|
|
# and never reaches here. Zero-pointer (meta) tensors are left to serialization.
|
|
seen_tensors: dict = {}
|
|
backstop_delete = []
|
|
for key, value in post_state_dict.items():
|
|
if isinstance(value, torch.Tensor) and value.data_ptr() != 0:
|
|
tensor_id = value.data_ptr()
|
|
if tensor_id in seen_tensors:
|
|
backstop_delete.append(key)
|
|
logger.warning(
|
|
f"Found tied weight: '{key}' is tied to '{seen_tensors[tensor_id]}'. "
|
|
f"Removing duplicate '{key}' from the exported state dict."
|
|
)
|
|
else:
|
|
seen_tensors[tensor_id] = key
|
|
for key in backstop_delete:
|
|
del post_state_dict[key]
|
|
|
|
# If a dense tie was dropped but a non-`bias` key still remains under its prefix, it is likely
|
|
# a new quantizer companion missing from `weight_suffixes` -- warn so it gets added.
|
|
for pre in dropped_dense_prefixes:
|
|
leftovers = [
|
|
k for k in post_state_dict if k.startswith(f"{pre}.") and k.rsplit(".", 1)[-1] != "bias"
|
|
]
|
|
if leftovers:
|
|
logger.warning(
|
|
f"Tied weight '{pre}.weight' was deduped but companion key(s) {leftovers} remain "
|
|
f"under '{pre}.'; likely an un-enumerated quantizer companion -- add its suffix to "
|
|
f"weight_suffixes in postprocess_state_dict so it is dropped with the tie."
|
|
)
|
|
|
|
return post_state_dict
|
|
|
|
|
|
def all_items_same(item_list):
|
|
"""Checks if all elements in the provided list are the same."""
|
|
return all(x == item_list[0] for x in item_list)
|
|
|
|
|
|
def _update_pre_quant_scale(module, new_pre_quant_scale):
|
|
old_pre_quant_scale = module.input_quantizer._pre_quant_scale
|
|
# do the processing in fp32 for numerical stability
|
|
dtype = module.weight.dtype
|
|
module.weight = nn.Parameter(
|
|
(
|
|
module.weight.to(torch.float32)
|
|
* old_pre_quant_scale.to(dtype=torch.float32, device=module.weight.device)
|
|
/ new_pre_quant_scale.to(dtype=torch.float32, device=module.weight.device)
|
|
).to(dtype)
|
|
)
|
|
module.input_quantizer.pre_quant_scale = new_pre_quant_scale
|
|
|
|
# Redo weights collection
|
|
module.weight_quantizer.reset_amax()
|
|
enable_stats_collection(module.weight_quantizer)
|
|
module.weight_quantizer(module.weight)
|
|
finish_stats_collection(module.weight_quantizer)
|
|
|
|
|
|
def _update_svdquant(modules, new_pre_quant_scale):
|
|
"""Updates the pre_quant_scale, svdquant_lora_a and svdquant_lora_b matrices when pre_quant_scale is changed."""
|
|
new_pre_quant_scale = new_pre_quant_scale.to(torch.float32)
|
|
lora_a = [m.weight_quantizer.svdquant_lora_a.to(torch.float32) for m in modules]
|
|
lora_b = [m.weight_quantizer.svdquant_lora_b.to(torch.float32) for m in modules]
|
|
weight = [m.weight.to(torch.float32) for m in modules]
|
|
old_pre_quant_scale = [m.input_quantizer._pre_quant_scale.to(torch.float32) for m in modules]
|
|
weight = [
|
|
(w + (lb @ la)) * (s / new_pre_quant_scale)
|
|
for w, la, lb, s in zip(weight, lora_a, lora_b, old_pre_quant_scale)
|
|
]
|
|
weight_concatenated = torch.cat(weight, dim=0)
|
|
lb, la = svd(weight_concatenated, rank=lora_a[0].shape[0])
|
|
weight_concatenated -= lb @ la
|
|
weight_concatenated = weight_concatenated.to(modules[0].weight.dtype)
|
|
la = la.to(modules[0].weight_quantizer.svdquant_lora_a.dtype)
|
|
lb = lb.to(modules[0].weight_quantizer.svdquant_lora_b.dtype)
|
|
new_pre_quant_scale = new_pre_quant_scale.to(modules[0].input_quantizer.pre_quant_scale.dtype)
|
|
|
|
index = 0
|
|
for i, module in enumerate(modules):
|
|
module.input_quantizer.pre_quant_scale = new_pre_quant_scale
|
|
module.weight_quantizer.svdquant_lora_a = la
|
|
assert lora_b[i].shape[0] == module.weight.shape[0]
|
|
module.weight_quantizer.svdquant_lora_b = lb[index : index + lora_b[i].shape[0], :]
|
|
module.weight = nn.Parameter(weight_concatenated[index : index + lora_b[i].shape[0], :])
|
|
index += lora_b[i].shape[0]
|
|
# Redo weights collection
|
|
module.weight_quantizer.reset_amax()
|
|
enable_stats_collection(module.weight_quantizer)
|
|
module.weight_quantizer(module.weight)
|
|
finish_stats_collection(module.weight_quantizer)
|
|
|
|
|
|
# AWQ pre_quant_scale fusion rules are per-model data and live in modelopt/torch/models/*:
|
|
# - Attention: fold o_proj's pre_quant_scale into v_proj's output dimension.
|
|
# Before: o_proj_out = [attn @ (v_proj_in @ v_proj.W^T)^T * scale] @ o_proj.W^T
|
|
# After: o_proj_out = [attn @ (v_proj_in @ (v_proj.W * scale)^T)^T] @ o_proj.W^T
|
|
# - MLP: fold down_proj's pre_quant_scale into up_proj's output dimension.
|
|
# Before: down_proj_out = {[act_fn(gate_proj(x)) * up_proj(x)] * scale} @ down_proj.W^T
|
|
# After: down_proj_out = {[act_fn(gate_proj(x)) * (up_proj(x) * scale)]} @ down_proj.W^T
|
|
# Each rule is a (module_class_substrings, fuse_into, fuse_from) triple.
|
|
|
|
|
|
def _pqs_fuse_rules(model_type: str | None):
|
|
"""The AWQ pre_quant_scale fusion rules to try, preferring the model's own.
|
|
|
|
A rule asserts a mathematical equivalence for one model's modules, so a registered
|
|
model uses only what its own spec declares. The aggregate across every spec is the
|
|
fallback for a model with no spec, or one whose spec declares no rules -- which is
|
|
what this did for every model before. That fallback is safe rather than merely
|
|
tolerated, because each rule is keyed on class-name substrings (``LlamaAttention``,
|
|
``Qwen3MoeMLP``) that cannot match another family's modules.
|
|
"""
|
|
spec = get_spec(model_type) if model_type else None
|
|
export_spec = spec.export_spec if spec is not None else None
|
|
if export_spec is not None and export_spec.pqs_fuse_rules:
|
|
return export_spec.pqs_fuse_rules
|
|
return list_all_possible("pqs_fuse_rules")
|
|
|
|
|
|
def fuse_prequant_to_linear(
|
|
model: torch.nn.Module, fuse_grouped_heads=False, model_type: str | None = None
|
|
):
|
|
"""Fuse pre_quant_scale to the linear weights if possible.
|
|
|
|
Args:
|
|
model: The model to fuse pre_quant_scale to.
|
|
fuse_grouped_heads: If True, fuse the pre_quant_scale even if dimension between pre_quant_scale
|
|
and linear weights is not the same.
|
|
model_type: The root model's HF model type, used to prefer its own fusion rules.
|
|
|
|
Returns:
|
|
fused_modules: A list of modules of which pre_quant_scale is fused to the previous linear layer.
|
|
"""
|
|
# Resolved once: this is a fixed vocabulary for the whole walk, and recomputing it per
|
|
# module rescans every registered spec.
|
|
fuse_rules = _pqs_fuse_rules(model_type)
|
|
|
|
# Fuse pre_quant_scale to the linear weights
|
|
for _, module in model.named_modules():
|
|
for target_module_list, fuse_into, fuse_from in fuse_rules:
|
|
if any(module_name in type(module).__name__ for module_name in target_module_list):
|
|
linear_fuse_into = module.get_submodule(fuse_into)
|
|
linear_pqs_from = module.get_submodule(fuse_from)
|
|
if hasattr(linear_pqs_from, "input_quantizer") and hasattr(
|
|
linear_pqs_from.input_quantizer, "_pre_quant_scale"
|
|
):
|
|
pre_quant_scale = linear_pqs_from.input_quantizer._pre_quant_scale
|
|
|
|
# for GQA/MQA models, we can apply averaging to the pre_quant_scale for shared head groups
|
|
if pre_quant_scale.numel() != linear_fuse_into.weight.shape[-2]:
|
|
if (
|
|
not fuse_grouped_heads
|
|
or "attention" not in type(module).__name__.lower()
|
|
):
|
|
warn(
|
|
f"Skipping pattern fuse prequant for {type(module).__name__}"
|
|
f"pre_quant_scale dim {pre_quant_scale.numel()} != "
|
|
f"out_channel dim {linear_fuse_into.weight.shape[-2]}"
|
|
)
|
|
continue
|
|
config = module.config
|
|
num_kv_heads = config.num_key_value_heads
|
|
kv_head_dim = linear_fuse_into.weight.shape[0] // num_kv_heads
|
|
n_rep = pre_quant_scale.numel() // num_kv_heads // kv_head_dim
|
|
|
|
# Reshape:(num_kv_heads, n_rep, kv_head_dim)
|
|
# n_rep is the number of query group
|
|
averaged_scale = pre_quant_scale.view(
|
|
num_kv_heads, n_rep, kv_head_dim
|
|
).mean(dim=1)
|
|
|
|
# To update o_proj, we need to repeat back to original shape
|
|
repeated_scale = (
|
|
averaged_scale.unsqueeze(1)
|
|
.expand(num_kv_heads, n_rep, kv_head_dim)
|
|
.reshape(-1)
|
|
)
|
|
# Update o_proj's pre_quant_scale
|
|
_update_pre_quant_scale(linear_pqs_from, repeated_scale)
|
|
|
|
# Use averaged scale (flattened) for v_proj fusion
|
|
pre_quant_scale = averaged_scale.reshape(-1)
|
|
|
|
# Fuse the pre_quant_scale to weight
|
|
linear_fuse_into.weight = torch.nn.Parameter(
|
|
linear_fuse_into.weight * pre_quant_scale.view(-1, 1)
|
|
)
|
|
if hasattr(linear_fuse_into, "bias") and linear_fuse_into.bias is not None:
|
|
linear_fuse_into.bias = torch.nn.Parameter(
|
|
linear_fuse_into.bias * pre_quant_scale
|
|
)
|
|
|
|
# Recalibrate the weight quantizer for linear_fuse_into
|
|
linear_fuse_into.weight_quantizer.reset_amax()
|
|
enable_stats_collection(linear_fuse_into.weight_quantizer)
|
|
linear_fuse_into.weight_quantizer(linear_fuse_into.weight)
|
|
finish_stats_collection(linear_fuse_into.weight_quantizer)
|
|
|
|
delattr(linear_pqs_from.input_quantizer, "_pre_quant_scale")
|
|
setattr(linear_pqs_from, "fused_with_prequant", True)
|
|
|
|
|
|
def _layernorm_uses_weight_plus_one(module: torch.nn.Module) -> bool:
|
|
"""Whether this norm stores ``w - 1``, so export must fold scales into ``weight + 1``.
|
|
|
|
The names are per-model data (``ExportSpec.weight_plus_one_norm_names``) but the match
|
|
is by *substring*, not exact name, which is what the hardcoded list this replaced did.
|
|
That is load-bearing rather than sloppy: the convention travels by family, and
|
|
transformers derives several norms whose names embed a registered one --
|
|
``DiffusionGemmaRMSNorm``, ``RecurrentGemmaRMSNorm``, ``T5GemmaRMSNorm``,
|
|
``T5Gemma2RMSNorm``, ``VaultGemmaRMSNorm`` all carry the Gemma convention. Exact
|
|
matching would drop them silently, and the failure is wrong numerics in an exported
|
|
checkpoint rather than an error.
|
|
|
|
Checked against every class in the MRO so quantized subclasses still match.
|
|
``zero_centered_gamma`` is the structural fallback for norms that announce it.
|
|
"""
|
|
registered = [n.lower() for n in list_all_possible("weight_plus_one_norm_names")]
|
|
mro_names = [cls.__name__.lower() for cls in type(module).__mro__]
|
|
if any(name in cls_name for cls_name in mro_names for name in registered):
|
|
return True
|
|
|
|
return bool(hasattr(module, "zero_centered_gamma") and module.zero_centered_gamma)
|
|
|
|
|
|
def fuse_prequant_layernorm(
|
|
layernorm_module: torch.nn.Module,
|
|
modules: list[torch.Tensor],
|
|
):
|
|
"""Scales layernorm weights with avg_pre_quant_scale of the modules list and sets pre_quant_scales to be deleted.
|
|
|
|
original:
|
|
layernorm_output = (normalization(input) * weight) + bias
|
|
layernorm_output_scaled = layernorm_output * pre_quant_scale
|
|
|
|
fused:
|
|
fused_weight = weight * avg_pre_quant_scale
|
|
fused_bias = bias * avg_pre_quant_scale
|
|
layernorm_output_scaled = (normalization(input) * fused_weight) + fused_bias
|
|
"""
|
|
if not hasattr(modules[0].input_quantizer, "_pre_quant_scale"):
|
|
return
|
|
|
|
pre_quant_scale = modules[0].input_quantizer._pre_quant_scale.to(layernorm_module.weight.device)
|
|
if _layernorm_uses_weight_plus_one(layernorm_module):
|
|
# For norms that use (1 + weight) in forward, fold pre_quant_scale into the effective weight.
|
|
fused_weight = (layernorm_module.weight + 1.0) * pre_quant_scale - 1.0
|
|
else:
|
|
fused_weight = layernorm_module.weight * pre_quant_scale
|
|
layernorm_module.weight = torch.nn.Parameter(fused_weight.to(layernorm_module.weight.dtype))
|
|
if hasattr(layernorm_module, "bias") and layernorm_module.bias is not None:
|
|
layernorm_module.bias = torch.nn.Parameter(layernorm_module.bias * pre_quant_scale)
|
|
# Pre_quant_scales of modules must not be exported, since they have been fused with layernorm
|
|
for module in modules:
|
|
delattr(module.input_quantizer, "_pre_quant_scale")
|
|
setattr(module, "fused_with_prequant", True)
|
|
|
|
|
|
def preprocess_linear_fusion(modules: list[torch.nn.Module], resmooth_only=False):
|
|
"""Preprocess the quantized linears that we plan to fuse.
|
|
|
|
Use resmooth_only for MOE experts as each individual expert is not fused.
|
|
"""
|
|
quantization_format_list = [get_quantization_format(module) for module in modules]
|
|
assert all_items_same(quantization_format_list), "Modules have different quantization formats"
|
|
|
|
# Activation
|
|
if hasattr(modules[0], "input_quantizer"):
|
|
# Resmooth
|
|
if modules[0].input_quantizer.pre_quant_scale is not None:
|
|
avg_prequant_scale = torch.mean(
|
|
torch.stack([module.input_quantizer.pre_quant_scale for module in modules]),
|
|
dim=0,
|
|
)
|
|
|
|
if all(
|
|
getattr(m.weight_quantizer, "svdquant_lora_a", None) is not None for m in modules
|
|
):
|
|
_update_svdquant(modules, avg_prequant_scale)
|
|
else:
|
|
for module in modules:
|
|
if not torch.equal(module.input_quantizer.pre_quant_scale, avg_prequant_scale):
|
|
_update_pre_quant_scale(module, avg_prequant_scale)
|
|
|
|
if resmooth_only:
|
|
return
|
|
|
|
if modules[0].input_quantizer.is_enabled and modules[0].input_quantizer.amax is not None:
|
|
assert modules[0].input_quantizer.amax.numel() == 1, (
|
|
"Only support scalar input quant amax"
|
|
)
|
|
|
|
input_amax = torch.max(torch.stack([module.input_quantizer.amax for module in modules]))
|
|
for module in modules:
|
|
module.input_quantizer.amax = input_amax
|
|
|
|
# Weight
|
|
if hasattr(modules[0], "weight_quantizer"):
|
|
is_seq_quant = isinstance(modules[0].weight_quantizer, SequentialQuantizer)
|
|
|
|
if is_seq_quant:
|
|
if modules[0].weight_quantizer[-1].is_enabled:
|
|
assert len(modules[0].weight_quantizer) == 2
|
|
weight_amax = torch.max(
|
|
torch.stack([module.weight_quantizer[-1].amax for module in modules])
|
|
)
|
|
for module in modules:
|
|
module.weight_quantizer[-1].amax = weight_amax
|
|
|
|
# Handle NVFP4StaticQuantizer: unify global_amax for fused layers
|
|
elif isinstance(modules[0].weight_quantizer, NVFP4StaticQuantizer):
|
|
global_amax_list = [
|
|
m.weight_quantizer.global_amax
|
|
for m in modules
|
|
if m.weight_quantizer.global_amax is not None
|
|
]
|
|
if global_amax_list:
|
|
unified_global_amax = torch.max(torch.stack(global_amax_list))
|
|
for module in modules:
|
|
module.weight_quantizer.global_amax = unified_global_amax
|
|
|
|
elif (
|
|
modules[0].weight_quantizer.is_enabled
|
|
and modules[0].weight_quantizer.amax is not None
|
|
and modules[0].weight_quantizer.amax.numel() == 1
|
|
):
|
|
weight_amax = torch.max(
|
|
torch.stack([module.weight_quantizer.amax for module in modules])
|
|
)
|
|
for module in modules:
|
|
module.weight_quantizer.amax = weight_amax
|
|
|
|
|
|
def seed_carried_over_exclusions(model: nn.Module, quant_config: dict) -> list[str]:
|
|
"""Add carried-weight module names to an already-built ``quant_config``'s exclusions.
|
|
|
|
The single place carried weights reach ``exclude_modules``, for both exporters.
|
|
:func:`get_quant_config` calls it once the per-layer pass is done, and the layerwise exporter
|
|
calls it again from ``finalize()`` -- it snapshots its config during ``bind()``, while
|
|
calibration is still running and long before ``export_hf_checkpoint`` records what it carried,
|
|
so it has no chance to see them any earlier. Without it a layerwise export copies GLM-4.7's
|
|
``mtp.safetensors`` into the checkpoint with nothing in ``exclude_modules``: the same
|
|
NVBug 5718750 failure this pass exists to prevent, reached through the other exporter.
|
|
|
|
Exclusions are exact module names rather than prefix wildcards. That keeps both exporters
|
|
emitting the same thing, and a literal can never over-match a module the export did in fact
|
|
quantize -- the risk :func:`_prefix_wildcard_summarize_exclude_modules` has to guard against by
|
|
consulting ``quantized_layers``, which is unavailable by the time the layerwise path runs.
|
|
|
|
Returns the names it added. No-op when the export is not uniformly quantized -- there is no
|
|
single ``quant_algo`` for a deployment framework to misapply, so there is nothing to exclude
|
|
a weight from.
|
|
"""
|
|
names = _get_carried_over_module_names(model)
|
|
if not names:
|
|
return []
|
|
quantization = quant_config.get("quantization")
|
|
if not isinstance(quantization, dict):
|
|
return []
|
|
if quantization.get("quant_algo") in (None, QUANTIZATION_NONE, "MIXED_PRECISION"):
|
|
return []
|
|
exclude_modules = quantization.setdefault("exclude_modules", [])
|
|
added = [n for n in names if not any(fnmatch.fnmatch(n, p) for p in exclude_modules)]
|
|
if added:
|
|
exclude_modules.extend(added)
|
|
exclude_modules.sort()
|
|
return added
|
|
|
|
|
|
def _get_carried_over_module_names(model: nn.Module) -> list[str]:
|
|
"""Return module names for checkpoint weights carried over without a module.
|
|
|
|
Weights the loader could not place -- an MTP head, an auxiliary tower -- are copied into
|
|
the export verbatim from the source checkpoint (see
|
|
:func:`modelopt.torch.export.unified_export_hf.read_unplaced_weights`). They
|
|
have no module in the live model, so the quantizer walk in :func:`get_quant_config` cannot
|
|
see them and would leave them out of ``exclude_modules`` even though their original-precision
|
|
weight is written to the checkpoint. A deployment framework then reads the top-level
|
|
``quant_algo`` and tries to load e.g. an MTP ``eh_proj`` as an FP8 weight.
|
|
|
|
This is the same failure the MoE-router pass above exists to prevent -- only the reason the
|
|
module is invisible differs (no quantizer there, no module at all here).
|
|
|
|
Prefers ``_modelopt_carried_over_names``, which the export records once it knows what it
|
|
actually wrote -- carried tensors plus the off-index sidecars copied verbatim. Those sidecars
|
|
are never ``unexpected_keys``, so the unplaced list alone would miss GLM-4.7's
|
|
``mtp.safetensors`` and leave its tensors in the export with nothing in ``exclude_modules``.
|
|
|
|
A state-dict key is ``<module path>.<parameter name>``, so the owning module is the key with
|
|
its last component removed. Keys without a dot are top-level tensors with no module and are
|
|
skipped.
|
|
"""
|
|
keys = getattr(model, "_modelopt_carried_over_names", None)
|
|
if keys is None:
|
|
# Export has not recorded yet (or this model never went through it). The recorded unplaced
|
|
# list is the best available answer; it is wider than what gets written, so it can name a
|
|
# module the export did not emit. That way round is harmless -- a deployment framework
|
|
# ignores an exclusion it finds no weight for, but fails loading one it was never told
|
|
# about.
|
|
keys = getattr(model, "_modelopt_unplaced_source_keys", None) or []
|
|
return sorted({key.rsplit(".", 1)[0] for key in keys if "." in key})
|
|
|
|
|
|
def _get_unquantized_moe_router_names(model: nn.Module) -> list[str]:
|
|
"""Return the names of MoE router/gate submodules left in original precision.
|
|
|
|
A module is added to ``exclude_modules`` during unified HF export only if it carries a
|
|
quantizer (even a disabled one) -- see :func:`get_quant_config`. MoE routers are kept
|
|
unquantized on purpose, but on ``transformers>=5.0`` they are no longer ``nn.Linear``
|
|
modules (e.g. ``TopKRouter``), so ``mtq.quantize`` never attaches a quantizer to them.
|
|
Without this, the BF16 router weight is written to the checkpoint but omitted from
|
|
``exclude_modules``, and deployment frameworks (vLLM / SGLang) then try to load it as a
|
|
quantized weight -- e.g. ``AssertionError: Tried to load weights of size [E, H] to a
|
|
parameter of size [E, H/2]`` for Qwen3-MoE.
|
|
|
|
Routers are detected structurally: an MoE block exposes an ``experts`` container plus a
|
|
``gate`` / ``router`` (or ``shared_expert_gate``) submodule that owns a weight tensor.
|
|
Routers that the user opted to quantize (a non-NONE format) are skipped.
|
|
"""
|
|
router_attrs = ("gate", "router", "shared_expert_gate")
|
|
router_names = []
|
|
for name, module in model.named_modules():
|
|
if not hasattr(module, "experts"):
|
|
continue
|
|
for attr in router_attrs:
|
|
router = getattr(module, attr, None)
|
|
if not isinstance(router, nn.Module):
|
|
continue
|
|
if not isinstance(getattr(router, "weight", None), torch.Tensor):
|
|
continue
|
|
if get_quantization_format(router) != QUANTIZATION_NONE:
|
|
continue
|
|
router_names.append(f"{name + '.' if name else ''}{attr}")
|
|
return router_names
|
|
|
|
|
|
def get_quant_config(
|
|
model: nn.Module,
|
|
is_modelopt_qlora: bool = False,
|
|
) -> dict[str, Any]:
|
|
"""Generate quantization config for a model.
|
|
|
|
The model should be the root model. It can be fully quantized, partially quantized or
|
|
mixed-precision quantized.
|
|
|
|
Args:
|
|
model: The PyTorch model to make config for.
|
|
is_modelopt_qlora: Whether the model is a modelopt-trained QLoRA model.
|
|
|
|
Returns:
|
|
Dictionary containing the quantization configuration
|
|
"""
|
|
# Find first quantized linear layer to determine quantization format
|
|
quantization_format = None
|
|
block_size = None
|
|
|
|
# Create base config
|
|
quant_config: dict[str, Any] = {
|
|
"producer": {
|
|
"name": "modelopt",
|
|
"version": __version__,
|
|
},
|
|
}
|
|
|
|
default_quantization = {
|
|
"quant_algo": None,
|
|
"kv_cache_quant_algo": None,
|
|
}
|
|
|
|
quant_config["quantization"] = default_quantization
|
|
|
|
# Layer config dict holds quantization format of each layer.
|
|
# It also holds awq_block_size information for applicable layers.
|
|
layer_config_dict = {}
|
|
|
|
kv_cache_formats: set[str] = set()
|
|
kv_cache_quantized_layers: dict[str, dict[str, str]] = {}
|
|
language_model_lineage = get_language_model_from_vl(model)
|
|
language_model_modules = (
|
|
None
|
|
if language_model_lineage is None
|
|
else {id(module) for module in language_model_lineage[-1].modules()}
|
|
)
|
|
for name, module in dict(model.named_modules()).items():
|
|
# Check for standard quantizers or any quantizers from weight attributes
|
|
weight_names = list(weight_attr_names(module))
|
|
has_quantizers = any(
|
|
hasattr(module, quantizer_attr_names(weight_name).weight_quantizer)
|
|
or hasattr(module, quantizer_attr_names(weight_name).input_quantizer)
|
|
for weight_name in weight_names
|
|
)
|
|
|
|
# Skip LORA module and adapters.
|
|
# ModelOpt does not currently quantize these layers in QLoRA path.
|
|
skip_layer = is_modelopt_qlora and (
|
|
hasattr(module, "base_layer") or "lora_A" in name or "lora_B" in name
|
|
)
|
|
|
|
if has_quantizers and not skip_layer:
|
|
quantization_format = get_quantization_format(module)
|
|
|
|
# For MoE expert modules, we need to extract block size from the correct weight quantizer
|
|
# Try to get block size from each weight attribute (e.g., gate_up_proj, down_proj)
|
|
block_size = 0
|
|
|
|
for weight_name in weight_names:
|
|
weight_block_size = get_weight_block_size(module, weight_name)
|
|
if weight_block_size > 0:
|
|
block_size = weight_block_size
|
|
break
|
|
|
|
# Fallback to default weight quantizer if no specific weight quantizer found
|
|
if block_size == 0:
|
|
block_size = get_weight_block_size(module)
|
|
|
|
# Static NVFP4 uses pre-computed per-block scales from MSE calibration
|
|
if quantization_format == QUANTIZATION_NVFP4:
|
|
weight_quantizer = getattr(module, "weight_quantizer", None)
|
|
if weight_quantizer is None:
|
|
# Try to get from first weight attribute
|
|
for wn in weight_names:
|
|
weight_quantizer = getattr(
|
|
module, quantizer_attr_names(wn).weight_quantizer, None
|
|
)
|
|
if weight_quantizer is not None:
|
|
break
|
|
if weight_quantizer is not None:
|
|
is_static = isinstance(weight_quantizer, NVFP4StaticQuantizer)
|
|
if is_static:
|
|
quantization_format = "nvfp4_static"
|
|
|
|
# Construct per layer config dictionary
|
|
layer_config_dict[name + ".quantization"] = quantization_format
|
|
layer_config_dict[name + ".awq_block_size"] = block_size
|
|
|
|
# Find kv cache quant format
|
|
is_language_model_module = (
|
|
language_model_modules is None or id(module) in language_model_modules
|
|
)
|
|
k_quantizer = getattr(module, "k_bmm_quantizer", None)
|
|
v_quantizer = getattr(module, "v_bmm_quantizer", None)
|
|
output_quantizer = getattr(module, "output_quantizer", None)
|
|
# Projection modules also expose output_quantizer; it is a KV fallback only on an
|
|
# attention boundary that has at least one K/V quantizer attribute.
|
|
is_kv_boundary = k_quantizer is not None or v_quantizer is not None
|
|
enabled_kv_quantizers = [
|
|
quantizer
|
|
for quantizer in (k_quantizer, v_quantizer, output_quantizer)
|
|
if is_kv_boundary and quantizer is not None and quantizer.is_enabled
|
|
]
|
|
if enabled_kv_quantizers and is_language_model_module:
|
|
module_kv_quant = get_kv_cache_dtype(module)
|
|
if module_kv_quant != QUANTIZATION_NONE:
|
|
kv_cache_formats.add(module_kv_quant)
|
|
# Per-layer mixed-KV metadata is defined only for an actual enabled K/V pair.
|
|
# Keep the output-quantizer/single-sided Megatron fallback top-level only.
|
|
if (
|
|
k_quantizer is not None
|
|
and v_quantizer is not None
|
|
and k_quantizer.is_enabled
|
|
and v_quantizer.is_enabled
|
|
):
|
|
kv_cache_quantized_layers[name] = {"quant_algo": module_kv_quant}
|
|
|
|
# MoE routers/gates are intentionally kept in original precision. On transformers>=5.0 they
|
|
# are not nn.Linear modules (e.g. TopKRouter), never receive a quantizer, and would otherwise
|
|
# be missing from exclude_modules even though their BF16 weight is exported -- causing
|
|
# deployment frameworks to load them as quantized weights. Record them explicitly as
|
|
# unquantized so they land in exclude_modules.
|
|
for router_name in _get_unquantized_moe_router_names(model):
|
|
layer_config_dict.setdefault(router_name + ".quantization", QUANTIZATION_NONE)
|
|
|
|
# Process per layer quantization config dict
|
|
quant_config["quantization"].update(process_layer_quant_config(layer_config_dict))
|
|
|
|
# Carried weights are seeded AFTER the per-layer pass, through the same helper the layerwise
|
|
# exporter calls. Seeding them into layer_config_dict instead would route them through
|
|
# _prefix_wildcard_summarize_exclude_modules and emit wildcards here, while the layerwise path
|
|
# -- which can only act once its config is already built -- emits literals: one model, two
|
|
# exporters, two different-looking quantization_config.ignore. The summarizer cannot serve both,
|
|
# because it needs `quantized_layers` to avoid a wildcard swallowing a quantized module and
|
|
# process_layer_quant_config pops that key before returning.
|
|
seed_carried_over_exclusions(model, quant_config)
|
|
|
|
weight_quant_algo = quant_config["quantization"].get("quant_algo")
|
|
needs_layerwise_kv_metadata = bool(kv_cache_quantized_layers) and (
|
|
weight_quant_algo is None or len(kv_cache_formats) > 1
|
|
)
|
|
if needs_layerwise_kv_metadata:
|
|
if weight_quant_algo not in (None, "MIXED_PRECISION"):
|
|
warn(
|
|
"The exported checkpoint combines uniform quantized weights with a mixed-precision "
|
|
"KV-cache layer map. Released runtimes do not yet consume "
|
|
"kv_cache_quantized_layers for uniform-weight ModelOpt checkpoints. Export succeeds "
|
|
"for artifact inspection only; do not deploy this checkpoint until the runtime "
|
|
"adds that metadata path. The exported metadata records "
|
|
"kv_cache_deployment_supported=false.",
|
|
stacklevel=2,
|
|
)
|
|
quant_config["quantization"]["kv_cache_deployment_supported"] = False
|
|
# KV metadata is orthogonal to weight metadata. In particular, a KV-only search
|
|
# must preserve BF16 weights instead of synthesizing a weight quantization algorithm.
|
|
quant_config["quantization"]["kv_cache_quant_algo"] = (
|
|
next(iter(kv_cache_formats)) if len(kv_cache_formats) == 1 else "MIXED_PRECISION"
|
|
)
|
|
quant_config["quantization"]["kv_cache_quantized_layers"] = kv_cache_quantized_layers
|
|
quant_config["quantization"]["kv_cache_schema_version"] = 1
|
|
elif len(kv_cache_formats) == 1:
|
|
# Preserve the pre-AutoQuantize uniform KV schema, including partial coverage.
|
|
quant_config["quantization"]["kv_cache_quant_algo"] = next(iter(kv_cache_formats))
|
|
|
|
return quant_config
|
|
|
|
|
|
def has_quantized_modules(model: nn.Module) -> bool:
|
|
"""Check if a model has any quantized modules.
|
|
|
|
Args:
|
|
model: The model to check.
|
|
|
|
Returns:
|
|
True if the model contains quantized modules, False otherwise.
|
|
"""
|
|
return any(
|
|
get_quantization_format(sub_module) != QUANTIZATION_NONE
|
|
for _, sub_module in model.named_modules()
|
|
)
|
|
|
|
|
|
def sync_tied_input_amax(model: nn.Module, tied_map: "TiedWeightMap | None" = None) -> int:
|
|
"""Max-merge ``input_quantizer`` amaxes across modules that share a weight, in place.
|
|
|
|
Tied modules whose forward paths see different activation ranges (encoder vs decoder in
|
|
YOCO-style models) must end up with one ``input_scale`` covering every side. Run BEFORE
|
|
per-module export so the merged amax flows into ``input_scale`` derivation; the model is
|
|
not expected to be reused afterward.
|
|
|
|
Declared ties are grouped by name via :class:`TiedWeightMap` (dense Linears and fused-MoE
|
|
containers). A physically shared but *undeclared* weight is grouped by ``id(weight)`` as a
|
|
fallback, so the side the address backstop later drops still had its amax merged in here.
|
|
Returns the number of groups merged; pass ``tied_map`` to reuse one, else it is built here.
|
|
"""
|
|
if tied_map is None:
|
|
tied_map = TiedWeightMap(model)
|
|
|
|
by_group: dict = defaultdict(list)
|
|
for name, m in model.named_modules():
|
|
# Fused MoE: 3-D source tensors with shared input quantizers
|
|
first_proj_attr = getattr(m, "_first_proj_attr", "gate_up_proj")
|
|
first_proj = getattr(m, first_proj_attr, None)
|
|
first_proj_input_quantizer_attr = f"{first_proj_attr}_input_quantizer"
|
|
if (
|
|
hasattr(m, first_proj_input_quantizer_attr)
|
|
and first_proj is not None
|
|
and hasattr(m, "down_proj")
|
|
and first_proj.dim() == 3
|
|
):
|
|
gk = tied_map.container_group_key(name, first_proj_attr)
|
|
if gk is not None:
|
|
by_group[("moe", gk)].append(m)
|
|
else:
|
|
# Undeclared share: group by projection identity so its amaxes still merge
|
|
# (symmetric with the dense fallback below).
|
|
by_group[("moe_shared", id(first_proj))].append(m)
|
|
# Dense quantized Linear with an input_quantizer
|
|
elif (
|
|
hasattr(m, "input_quantizer")
|
|
and hasattr(m, "weight")
|
|
and isinstance(m.weight, torch.nn.Parameter)
|
|
):
|
|
gk = tied_map.group_key(f"{name}.weight" if name else "weight")
|
|
if gk is not None:
|
|
by_group[("dense", gk)].append(m)
|
|
else:
|
|
# Undeclared share: group by object identity so its amaxes still merge.
|
|
by_group[("dense_shared", id(m.weight))].append(m)
|
|
|
|
def _merge(quantizers: list) -> bool:
|
|
"""Max-merge amaxes across the quantizer list. Returns True on merge."""
|
|
valid = [
|
|
q
|
|
for q in quantizers
|
|
if q is not None
|
|
and getattr(q, "is_enabled", False)
|
|
and getattr(q, "_amax", None) is not None
|
|
and not q._amax.is_meta
|
|
]
|
|
if len(valid) < 2:
|
|
return False
|
|
# Require scalar (per-tensor) amax — matches preprocess_linear_fusion.
|
|
if any(q._amax.numel() != 1 for q in valid):
|
|
warn(
|
|
"sync_tied_input_amax: non-scalar input_quantizer amax encountered "
|
|
"in a tied group; skipping. Only per-tensor input quantizers are "
|
|
"supported for tied-modules merging."
|
|
)
|
|
return False
|
|
merged = torch.max(torch.stack([q.amax for q in valid]))
|
|
for q in valid:
|
|
q.amax = merged.clone()
|
|
return True
|
|
|
|
synced = 0
|
|
for key, modules in by_group.items():
|
|
if len(modules) < 2:
|
|
continue
|
|
if key[0] == "moe":
|
|
first_proj_attr = getattr(modules[0], "_first_proj_attr", "gate_up_proj")
|
|
for q_name in (f"{first_proj_attr}_input_quantizer", "down_proj_input_quantizer"):
|
|
if _merge([getattr(m, q_name, None) for m in modules]):
|
|
synced += 1
|
|
elif _merge([m.input_quantizer for m in modules]):
|
|
synced += 1
|
|
return synced
|