[diffusion] chore: reuse shared checkpoint quant metadata resolver (#35174)

This commit is contained in:
Mick
2026-08-19 13:36:55 +08:00
committed by GitHub
parent 593b1a9b8a
commit a5c96362b6
3 changed files with 38 additions and 14 deletions
@@ -17,6 +17,7 @@ import torch
from diffusers.utils import SAFE_WEIGHTS_INDEX_NAME
from torch import nn
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig,
_patch_nunchaku_scales,
@@ -42,7 +43,6 @@ from sglang.multimodal_gen.runtime.utils.quantization_utils import (
get_quant_config,
get_quant_config_from_safetensors_metadata,
)
from sglang.srt.layers.quantization import QuantizationConfig
logger = init_logger(__name__)
@@ -13,6 +13,9 @@ from sglang.multimodal_gen.runtime.layers.quantization import (
get_quantization_config,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.model_loader.checkpoint_quantization import (
resolve_checkpoint_quant_spec,
)
logger = init_logger(__name__)
@@ -169,27 +172,17 @@ def get_quant_config(
quant_cls = _load_quant_cls(quant_cfg)
return quant_cls.from_config(quant_cfg, reverse_param_names_mapping)
if "quantization_config" not in model_config:
checkpoint_quant_spec = resolve_checkpoint_quant_spec(model_config)
if checkpoint_quant_spec is None:
return None
hf_quant_config = normalize_flat_modelopt_quant_config(
model_config["quantization_config"]
)
if hf_quant_config is not None and not isinstance(hf_quant_config, dict):
hf_quant_config = hf_quant_config.to_dict()
hf_quant_config = normalize_flat_modelopt_quant_config(checkpoint_quant_spec.config)
quant_cls = _load_quant_cls(hf_quant_config)
# GGUF doesn't have config file
if hf_quant_config["quant_method"] == "gguf":
return quant_cls.from_config({})
# some vision model may keep quantization_config in their text_config
hf_text_config = getattr(model_config, "text_config", None)
if hf_quant_config is None and hf_text_config is not None:
hf_quant_config = getattr(hf_text_config, "quantization_config", None)
if hf_quant_config is None:
# compressed-tensors uses a compressions_config
hf_quant_config = getattr(model_config, "compression_config", None)
if hf_quant_config is not None:
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
is_modelopt_fp8 = (
@@ -490,6 +490,37 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertTrue(config.load_in_4bit)
self.assertEqual(config.bnb_4bit_quant_type, "nf4")
def test_fp8_quant_config_resolves_from_text_config(self):
config = get_quant_config(
{
"text_config": {
"quantization_config": {
"quant_method": "fp8",
"activation_scheme": "dynamic",
}
}
},
"/unused/component/path",
)
self.assertIsInstance(config, Fp8Config)
self.assertTrue(config.is_checkpoint_fp8_serialized)
def test_bitsandbytes_quant_config_resolves_from_compression_config(self):
config = get_quant_config(
{
"compression_config": {
"quant_method": "bitsandbytes",
"load_in_4bit": True,
"bnb_4bit_quant_storage": "uint8",
}
},
"/unused/component/path",
)
self.assertEqual(config.get_name(), "bitsandbytes")
self.assertTrue(config.load_in_4bit)
def test_nvfp4_safetensors_inference_ignores_fp8_fallback_scales(self):
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
save_file(