[Quantization] Support config.json quantization_config format, fix exclude_modules matching, and fix KV cache scale loading for Nemotron (#18546)
Signed-off-by: root <dafrimi@nvidia.com>
This commit is contained in:
@@ -5,6 +5,7 @@ import logging
|
|||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import regex as re
|
||||||
import torch
|
import torch
|
||||||
from torch.nn.parameter import Parameter
|
from torch.nn.parameter import Parameter
|
||||||
|
|
||||||
@@ -295,6 +296,11 @@ class ModelOptQuantConfig(QuantizationConfig):
|
|||||||
elif self.kv_cache_quant_algo and isinstance(layer, RadixAttention):
|
elif self.kv_cache_quant_algo and isinstance(layer, RadixAttention):
|
||||||
return ModelOptFp8KVCacheMethod(self)
|
return ModelOptFp8KVCacheMethod(self)
|
||||||
elif isinstance(layer, FusedMoE):
|
elif isinstance(layer, FusedMoE):
|
||||||
|
# Check if MoE layer should be excluded from quantization
|
||||||
|
# (e.g., MTP layers that have no quantization scales in checkpoint)
|
||||||
|
if self.is_layer_excluded(prefix):
|
||||||
|
# Falls back to default unquantized MoE
|
||||||
|
return None
|
||||||
return Moe(self)
|
return Moe(self)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -321,6 +327,54 @@ class ModelOptQuantConfig(QuantizationConfig):
|
|||||||
# Preserve order, drop duplicates.
|
# Preserve order, drop duplicates.
|
||||||
self.exclude_modules = list(dict.fromkeys(expanded))
|
self.exclude_modules = list(dict.fromkeys(expanded))
|
||||||
|
|
||||||
|
def is_layer_excluded(self, prefix: str) -> bool:
|
||||||
|
"""Check if a layer should be excluded from quantization.
|
||||||
|
|
||||||
|
Handles:
|
||||||
|
- Exact matches (e.g., "lm_head" matching prefix "lm_head")
|
||||||
|
- Glob-style wildcards (e.g., "mtp*" matching "mtp_layers")
|
||||||
|
- Part-by-part matching (split prefix on "." and check each part)
|
||||||
|
- language_model. prefix stripping for vision-language models
|
||||||
|
- Fused module patterns (e.g., "q_a_proj" in "fused_qkv_a_proj_with_mqa")
|
||||||
|
"""
|
||||||
|
if not self.exclude_modules:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Build prefix variants: some models wrap layers under "language_model."
|
||||||
|
prefixes_to_check = [prefix]
|
||||||
|
if prefix.startswith("language_model."):
|
||||||
|
prefixes_to_check.append(prefix.removeprefix("language_model."))
|
||||||
|
|
||||||
|
# Fused module patterns: the exclude list may reference a sub-component
|
||||||
|
# (e.g., "q_a_proj") that is fused into a combined parameter name
|
||||||
|
# (e.g., "fused_qkv_a_proj_with_mqa"). We check if the last segment of
|
||||||
|
# the exclude pattern is a substring of the last segment of the prefix.
|
||||||
|
fused_patterns = {"q_a_proj", "q_b_proj", "kv_a_proj_with_mqa", "kv_b_proj"}
|
||||||
|
|
||||||
|
for pattern in self.exclude_modules:
|
||||||
|
# Convert glob-style wildcard to regex (e.g., "mtp*" -> "mtp.*")
|
||||||
|
regex_str = pattern.replace(".", r"\.").replace("*", r".*")
|
||||||
|
|
||||||
|
for pfx in prefixes_to_check:
|
||||||
|
if re.fullmatch(regex_str, pfx):
|
||||||
|
return True
|
||||||
|
# Part-by-part check: handles wildcards like "mtp*" matching
|
||||||
|
pfx_parts = pfx.split(".")
|
||||||
|
for part in pfx_parts:
|
||||||
|
if re.fullmatch(regex_str, part):
|
||||||
|
return True
|
||||||
|
|
||||||
|
# Check fused patterns: if the last segment of the exclude pattern
|
||||||
|
# is a known fused component, check if it appears in the prefix's
|
||||||
|
# last segment (handles fused_qkv_a_proj_with_mqa containing q_a_proj)
|
||||||
|
pattern_tail = pattern.rsplit(".", maxsplit=1)[-1]
|
||||||
|
if pattern_tail in fused_patterns:
|
||||||
|
for pfx in prefixes_to_check:
|
||||||
|
if pattern_tail in pfx.rsplit(".", maxsplit=1)[-1]:
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
class ModelOptFp8Config(ModelOptQuantConfig):
|
class ModelOptFp8Config(ModelOptQuantConfig):
|
||||||
"""Configuration for ModelOpt FP8 quantization, including serialization and compatibility checks."""
|
"""Configuration for ModelOpt FP8 quantization, including serialization and compatibility checks."""
|
||||||
@@ -376,11 +430,11 @@ class ModelOptFp8Config(ModelOptQuantConfig):
|
|||||||
quant_method = config.get("quant_algo")
|
quant_method = config.get("quant_algo")
|
||||||
if quant_method is not None:
|
if quant_method is not None:
|
||||||
# Flat format (config.json quantization_config)
|
# Flat format (config.json quantization_config)
|
||||||
# For kv_cache, check if kv_cache_scheme exists and extract algo
|
# Derive kv_cache quant from kv_cache_scheme dict
|
||||||
kv_cache_scheme = config.get("kv_cache_scheme")
|
kv_cache_scheme = config.get("kv_cache_scheme")
|
||||||
|
if isinstance(kv_cache_scheme, dict):
|
||||||
if (
|
if (
|
||||||
kv_cache_scheme
|
kv_cache_scheme.get("type") == "float"
|
||||||
and kv_cache_scheme.get("type") == "float"
|
|
||||||
and kv_cache_scheme.get("num_bits") == 8
|
and kv_cache_scheme.get("num_bits") == 8
|
||||||
):
|
):
|
||||||
kv_cache_quant_method = "FP8"
|
kv_cache_quant_method = "FP8"
|
||||||
@@ -388,7 +442,7 @@ class ModelOptFp8Config(ModelOptQuantConfig):
|
|||||||
# Map 'ignore' field to 'exclude_modules'
|
# Map 'ignore' field to 'exclude_modules'
|
||||||
exclude_modules = config.get("ignore")
|
exclude_modules = config.get("ignore")
|
||||||
else:
|
else:
|
||||||
# Fall back to nested format (hf_quant_config.json - legacy format)
|
# Fall back to nested format (hf_quant_config.json - will be deprecated)
|
||||||
try:
|
try:
|
||||||
quantization_section = cls.get_from_keys(config, ["quantization"])
|
quantization_section = cls.get_from_keys(config, ["quantization"])
|
||||||
quant_method = quantization_section.get("quant_algo")
|
quant_method = quantization_section.get("quant_algo")
|
||||||
@@ -417,18 +471,6 @@ class ModelOptFp8Config(ModelOptQuantConfig):
|
|||||||
packed_modules_mapping=config.get("packed_modules_mapping"),
|
packed_modules_mapping=config.get("packed_modules_mapping"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_layer_excluded(self, prefix: str) -> bool:
|
|
||||||
if len(self.exclude_modules) == 0:
|
|
||||||
return False
|
|
||||||
return any(
|
|
||||||
module in prefix
|
|
||||||
or (
|
|
||||||
prefix.startswith("language_model.")
|
|
||||||
and module in prefix.removeprefix("language_model.")
|
|
||||||
)
|
|
||||||
for module in self.exclude_modules
|
|
||||||
)
|
|
||||||
|
|
||||||
def get_quant_method(
|
def get_quant_method(
|
||||||
self, layer: torch.nn.Module, prefix: str
|
self, layer: torch.nn.Module, prefix: str
|
||||||
) -> Optional[QuantizeMethodBase]:
|
) -> Optional[QuantizeMethodBase]:
|
||||||
@@ -960,11 +1002,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
|||||||
quant_method = config.get("quant_algo")
|
quant_method = config.get("quant_algo")
|
||||||
if quant_method is not None:
|
if quant_method is not None:
|
||||||
# Flat format (config.json quantization_config)
|
# Flat format (config.json quantization_config)
|
||||||
# Note: FP4 models in config.json format may not have all the detailed fields
|
# Derive kv_cache_quant_algo from kv_cache_scheme dict
|
||||||
# that are present in hf_quant_config.json, so we need to handle defaults
|
|
||||||
kv_cache_quant_algo = config.get("kv_cache_quant_algo")
|
|
||||||
if not kv_cache_quant_algo:
|
|
||||||
# For config.json format, derive from kv_cache_scheme if available
|
|
||||||
kv_cache_scheme = config.get("kv_cache_scheme")
|
kv_cache_scheme = config.get("kv_cache_scheme")
|
||||||
if isinstance(kv_cache_scheme, dict):
|
if isinstance(kv_cache_scheme, dict):
|
||||||
if (
|
if (
|
||||||
@@ -1038,27 +1076,6 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
|||||||
config.get("packed_modules_mapping"),
|
config.get("packed_modules_mapping"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_layer_excluded(self, prefix: str):
|
|
||||||
import regex as re
|
|
||||||
|
|
||||||
fused_patterns = ["q_a_proj", "q_b_proj", "kv_a_proj_with_mqa", "kv_b_proj"]
|
|
||||||
prefix_split = prefix.split(".")
|
|
||||||
for pattern in self.exclude_modules:
|
|
||||||
regex_str = pattern.replace(".", r"\.").replace("*", r".*")
|
|
||||||
pattern_split = pattern.split(".")
|
|
||||||
if re.fullmatch(regex_str, prefix):
|
|
||||||
return True
|
|
||||||
elif (
|
|
||||||
pattern_split[-1] in fused_patterns
|
|
||||||
and pattern_split[-1] in prefix_split[-1]
|
|
||||||
):
|
|
||||||
# Check if the last part of the excluded pattern is contained in the last part of the prefix
|
|
||||||
# This handles fused modules like fused_qkv_a_proj_with_mqa that contain q_a_proj and kv_a_proj_with_mqa
|
|
||||||
# e.g., model.layers.{i}.self_attn.{fused_weight_name}
|
|
||||||
assert len(prefix_split) == 5 and len(pattern_split) == 5
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||||
return self._get_quant_method(
|
return self._get_quant_method(
|
||||||
layer,
|
layer,
|
||||||
|
|||||||
@@ -1214,17 +1214,22 @@ def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
|
|||||||
return remapped_name
|
return remapped_name
|
||||||
|
|
||||||
possible_scale_names = [".k_scale", ".v_scale"]
|
possible_scale_names = [".k_scale", ".v_scale"]
|
||||||
modelopt_scale_names = [".self_attn.k_proj.k_scale", ".self_attn.v_proj.v_scale"]
|
# Patterns where modelopt stores scales under k_proj/v_proj
|
||||||
|
# but the model expects them under attn (RadixAttention)
|
||||||
|
modelopt_attn_prefixes = [".self_attn.", ".mixer."]
|
||||||
for scale_name in possible_scale_names:
|
for scale_name in possible_scale_names:
|
||||||
if name.endswith(scale_name):
|
if name.endswith(scale_name):
|
||||||
# Check and remap the name based on modelopt scale names
|
# Check if this is a modelopt-style scale under k_proj/v_proj
|
||||||
if any(
|
matched_prefix = None
|
||||||
modelopt_scale_name in name
|
for attn_prefix in modelopt_attn_prefixes:
|
||||||
for modelopt_scale_name in modelopt_scale_names
|
if f"{attn_prefix}{scale_name[1]}_proj{scale_name}" in name:
|
||||||
):
|
matched_prefix = attn_prefix
|
||||||
|
break
|
||||||
|
|
||||||
|
if matched_prefix is not None:
|
||||||
remapped_name = name.replace(
|
remapped_name = name.replace(
|
||||||
f".self_attn.{scale_name[1]}_proj{scale_name}",
|
f"{matched_prefix}{scale_name[1]}_proj{scale_name}",
|
||||||
f".self_attn.attn{scale_name}",
|
f"{matched_prefix}attn{scale_name}",
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
remapped_name = name.replace(scale_name, f".attn{scale_name}")
|
remapped_name = name.replace(scale_name, f".attn{scale_name}")
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
replace_prefix,
|
replace_prefix,
|
||||||
replace_substrings,
|
replace_substrings,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
@@ -640,6 +641,12 @@ class NemotronHForCausalLM(nn.Module):
|
|||||||
"v_proj.v_scale": "attn.v_scale",
|
"v_proj.v_scale": "attn.v_scale",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
hf_to_sglang_mapper = WeightsMapper(
|
||||||
|
orig_to_new_prefix={
|
||||||
|
"backbone.": "model.",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
|
|||||||
Reference in New Issue
Block a user