From a5c96362b6626a996b5a0af3fe06c5ef1ae4779b Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 19 Aug 2026 13:36:55 +0800 Subject: [PATCH] [diffusion] chore: reuse shared checkpoint quant metadata resolver (#35174) --- .../runtime/loader/transformer_load_utils.py | 2 +- .../runtime/utils/quantization_utils.py | 19 ++++-------- .../test/unit/test_transformer_quant.py | 31 +++++++++++++++++++ 3 files changed, 38 insertions(+), 14 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 052545dd1..adebc8c2e 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index 10859b05c..e3a49f3f4 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -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 = ( diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index d8d1131af..cb85d61de 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -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(