From 28690f5aa57e67023ccbdcf842984a27a8cd13d4 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 31 Aug 2026 13:55:50 +0800 Subject: [PATCH] [diffusion] chore: detect quantized transformer replacements (#36916) --- .../configs/models/dits/base.py | 1 + .../component_loaders/transformer_loader.py | 1 + .../runtime/loader/transformer_load_utils.py | 176 ++++++++++++++++-- .../test/unit/test_transformer_quant.py | 134 ++++++++++++- 4 files changed, 292 insertions(+), 20 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/models/dits/base.py b/python/sglang/multimodal_gen/configs/models/dits/base.py index e32cc44ee..3c2ce00ab 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/base.py +++ b/python/sglang/multimodal_gen/configs/models/dits/base.py @@ -26,6 +26,7 @@ class DiTArchConfig(ArchConfig): # Reverse mapping for saving checkpoints: custom -> hf reverse_param_names_mapping: dict = field(default_factory=dict) + quant_ignore_remap: dict = field(default_factory=dict) hidden_size: int = 0 num_attention_heads: int = 0 num_channels_latents: int = 0 diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 7e24bc197..6fcf1f482 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -327,6 +327,7 @@ class TransformerLoader(ComponentLoader): gguf_file=gguf_file, checkpoint_quant_config=checkpoint_quant_config, transformer_override_config_path=transformer_override_config_path, + arch_config=dit_config.arch_config, ) if quant_spec.gguf_file is not None and is_minimax_h3: assert quant_spec.quant_config is not None 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 8694ad8a4..e040f53e8 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -15,8 +15,10 @@ from typing import Callable, Optional import torch from diffusers.utils import SAFE_WEIGHTS_INDEX_NAME +from safetensors import safe_open from torch import nn +from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, @@ -58,6 +60,9 @@ from sglang.multimodal_gen.runtime.weights.source import ( materialize_weight_set_config, resolve_safetensors_weight_set, ) +from sglang.srt.model_loader.checkpoint_quantization import ( + resolve_checkpoint_quant_spec, +) from sglang.srt.utils.hf_transformers import ( check_gguf_file, resolve_hf_gguf_reference, @@ -745,6 +750,7 @@ def resolve_transformer_quant_load_spec( gguf_file: str | None = None, checkpoint_quant_config: QuantizationConfig | None = None, transformer_override_config_path: str | None = None, + arch_config: DiTArchConfig | None = None, ) -> TransformerQuantLoadSpec: if gguf_file is not None: if checkpoint_quant_config is not None: @@ -777,6 +783,7 @@ def resolve_transformer_quant_load_spec( safetensors_list=safetensors_list, component_model_path=component_model_path, transformer_override_config_path=transformer_override_config_path, + arch_config=arch_config, ) if quant_config is not None: @@ -788,6 +795,10 @@ def resolve_transformer_quant_load_spec( ) nunchaku_config = server_args.nunchaku_config + if quant_config is not None and nunchaku_config is not None: + raise ValueError( + "Replacement checkpoint quantization and Nunchaku are mutually exclusive" + ) # resolve target param dtype param_dtype = _resolve_target_param_dtype( @@ -918,6 +929,110 @@ def _build_transformer_quant_adapters( return adapters +def _merge_quant_declaration(base: dict, incoming: dict) -> dict: + """Merge compatible checkpoint declarations and reject conflicts.""" + merged = dict(base) + for key, value in incoming.items(): + previous = merged.get(key) + if isinstance(previous, dict) and isinstance(value, dict): + merged[key] = _merge_quant_declaration(previous, value) + elif key in merged and previous != value: + raise ValueError(f"Conflicting checkpoint quantization field {key!r}") + else: + merged[key] = value + return merged + + +def _resolve_weight_override_quantization( + safetensors_list: list[str], + reverse_param_names_mapping: dict, + quant_ignore_remap: dict, +) -> tuple[Optional[QuantizationConfig], bool]: + """Resolve declarations carried by the materialized replacement weight set.""" + component_model_path = os.path.dirname(safetensors_list[0]) + component_config = {} + component_config_path = os.path.join(component_model_path, "config.json") + if os.path.isfile(component_config_path): + with open(component_config_path, encoding="utf-8") as config_stream: + component_config = json.load(config_stream) + + config_spec = resolve_checkpoint_quant_spec(component_config) + declaration = config_spec.config if config_spec is not None else None + header_quant_config = None + detected_quantized_tensors = False + + for safetensors_file in safetensors_list: + metadata = get_metadata_from_safetensors_file(safetensors_file) or {} + file_quant_config = get_quant_config_from_safetensors_metadata(safetensors_file) + if file_quant_config is not None: + if header_quant_config is not None and _get_quant_config_name( + header_quant_config + ) != _get_quant_config_name(file_quant_config): + raise ValueError("Conflicting safetensors quantization declarations") + header_quant_config = file_quant_config + + for metadata_key in ("_quantization_metadata", "quantization_config"): + serialized = metadata.get(metadata_key) + if serialized is None: + continue + try: + metadata_config = json.loads(serialized) + except json.JSONDecodeError as error: + raise ValueError( + f"Invalid {metadata_key} in {safetensors_file}" + ) from error + if not isinstance(metadata_config, dict): + raise ValueError( + f"Invalid {metadata_key} in {safetensors_file}: expected an object" + ) + metadata_spec = resolve_checkpoint_quant_spec( + {"quantization_config": metadata_config} + ) + assert metadata_spec is not None + declaration = ( + metadata_spec.config + if declaration is None + else _merge_quant_declaration(declaration, metadata_spec.config) + ) + + with safe_open(safetensors_file, framework="pt", device="cpu") as checkpoint: + for key in checkpoint.keys(): + if key.endswith((".weight_scale", ".input_scale", ".comfy_quant")): + detected_quantized_tensors = True + break + if key.endswith(".weight") and checkpoint.get_slice( + key + ).get_dtype() in ("F8_E4M3", "I8", "U8"): + detected_quantized_tensors = True + break + + if declaration is not None: + if "quant_method" not in declaration: + return header_quant_config, True + return ( + get_quant_config( + {"quantization_config": declaration}, + component_model_path, + reverse_param_names_mapping=reverse_param_names_mapping, + quant_ignore_remap=quant_ignore_remap, + ), + True, + ) + if header_quant_config is not None: + return header_quant_config, True + + description_config = get_quant_config( + component_config, + component_model_path, + reverse_param_names_mapping=reverse_param_names_mapping, + quant_ignore_remap=quant_ignore_remap, + ) + return ( + description_config, + detected_quantized_tensors or description_config is not None, + ) + + def _resolve_quant_config_from_transformer_override( override_config_path: str, ) -> Optional[QuantizationConfig]: @@ -938,13 +1053,37 @@ def _resolve_quant_config( safetensors_list: list[str], component_model_path: str, transformer_override_config_path: str | None = None, + arch_config: DiTArchConfig | None = None, ) -> Optional[QuantizationConfig]: """ resolve quant config from checkpoints' metadata priority: explicit --quantization flag -> model config.json -> safetensors metadata -> format-specific fallback """ + if arch_config is None: + arch_config = server_args.pipeline_config.dit_config.arch_config + param_names_mapping_dict = arch_config.param_names_mapping + reverse_param_names_mapping_dict = arch_config.reverse_param_names_mapping + quant_ignore_remap_dict = arch_config.quant_ignore_remap + + override_quant_config = None + override_declares_quantization = False + if server_args.transformer_weights_path: + ( + override_quant_config, + override_declares_quantization, + ) = _resolve_weight_override_quantization( + safetensors_list, + reverse_param_names_mapping_dict, + quant_ignore_remap_dict, + ) + # priority: explicit --quantization flag (e.g. mxfp8, mxfp4_npu, modelslim) if server_args.quantization is not None: + if override_declares_quantization: + raise ValueError( + "The replacement checkpoint already contains or declares " + "quantization; do not also set an online --quantization override" + ) from sglang.multimodal_gen.runtime.layers.quantization import ( get_quantization_config, ) @@ -976,24 +1115,15 @@ def _resolve_quant_config( ) return quant_cls(**quant_kwargs) - quant_config = get_quant_config(hf_config, component_model_path) - if quant_config is None and server_args.transformer_weights_path: - for safetensors_file in safetensors_list: - quant_config = get_quant_config_from_safetensors_metadata(safetensors_file) - if quant_config is not None: - return quant_config - - arch_config = server_args.pipeline_config.dit_config.arch_config - param_names_mapping_dict = arch_config.param_names_mapping - reverse_param_names_mapping_dict = getattr( - arch_config, "reverse_param_names_mapping", None - ) - quant_ignore_remap_dict = getattr(arch_config, "quant_ignore_remap", None) - quant_config = get_quant_config( - hf_config, - component_model_path, - reverse_param_names_mapping=reverse_param_names_mapping_dict, - quant_ignore_remap=quant_ignore_remap_dict, + quant_config = ( + override_quant_config + if server_args.transformer_weights_path + else get_quant_config( + hf_config, + component_model_path, + reverse_param_names_mapping=reverse_param_names_mapping_dict, + quant_ignore_remap=quant_ignore_remap_dict, + ) ) quant_config_name = _get_quant_config_name(quant_config) inferred_nvfp4_config = None @@ -1007,7 +1137,15 @@ def _resolve_quant_config( reverse_param_names_mapping_dict, fallback_group_size, ) - quant_config = _merge_modelopt_fp4_configs(quant_config, inferred_nvfp4_config) + if override_declares_quantization and override_quant_config is None: + if inferred_nvfp4_config is None: + raise ValueError( + "Replacement checkpoint contains quantized tensors but no supported " + "native quantization declaration" + ) + quant_config = inferred_nvfp4_config + else: + quant_config = _merge_modelopt_fp4_configs(quant_config, inferred_nvfp4_config) if quant_config is not None or transformer_override_config_path is None: return quant_config 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 18c75ba27..23f557c0f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -99,6 +99,7 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( _Flux2Nvfp4FallbackAdapter, _needs_device_weight_postprocess, _resolve_quant_config, + _resolve_weight_override_quantization, resolve_transformer_checkpoint_files, resolve_transformer_quant_load_spec, ) @@ -148,6 +149,132 @@ def _make_quant_config(name: str, **attrs): class TestTransformerQuantHelpers(unittest.TestCase): + @patch( + "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", + return_value=None, + ) + def test_weight_override_uses_adjacent_quantization_config(self, _build_nvfp4): + with tempfile.TemporaryDirectory() as directory: + weights = f"{directory}/model.safetensors" + save_file({"block.weight": torch.ones((2, 2))}, weights) + with open(f"{directory}/config.json", "w", encoding="utf-8") as stream: + json.dump( + { + "quantization_config": { + "quant_method": "fp8", + "activation_scheme": "dynamic", + } + }, + stream, + ) + server_args = self._make_server_args(transformer_weights_path=weights) + + quant_config = _resolve_quant_config( + hf_config={ + "quantization_config": { + "quant_method": "fp8", + "activation_scheme": "static", + } + }, + server_args=server_args, + safetensors_list=[weights], + component_model_path="/base", + ) + + self.assertIsInstance(quant_config, Fp8Config) + self.assertEqual(quant_config.activation_scheme, "dynamic") + self.assertTrue(quant_config.is_checkpoint_fp8_serialized) + + @patch( + "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", + return_value=None, + ) + def test_unquantized_weight_override_does_not_inherit_base_config( + self, _build_nvfp4 + ): + with tempfile.TemporaryDirectory() as directory: + weights = f"{directory}/model.safetensors" + save_file({"block.weight": torch.ones((2, 2))}, weights) + server_args = self._make_server_args(transformer_weights_path=weights) + + quant_config = _resolve_quant_config( + hf_config={"quantization_config": {"quant_method": "fp8"}}, + server_args=server_args, + safetensors_list=[weights], + component_model_path="/base", + ) + + self.assertIsNone(quant_config) + + def test_weight_override_defers_header_without_quant_method_to_layout(self): + with tempfile.TemporaryDirectory() as directory: + weights = f"{directory}/model.safetensors" + save_file( + {"block.weight": torch.ones((2, 2))}, + weights, + metadata={"quantization_config": json.dumps({"quant_algo": "NVFP4"})}, + ) + + quant_config, declared = _resolve_weight_override_quantization( + [weights], {}, {} + ) + + self.assertIsNone(quant_config) + self.assertTrue(declared) + + @patch( + "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", + return_value=None, + ) + def test_declared_weight_override_rejects_online_quantization(self, _build_nvfp4): + with tempfile.TemporaryDirectory() as directory: + weights = f"{directory}/model.safetensors" + save_file( + {"block.weight": torch.ones((2, 2))}, + weights, + metadata={ + "quantization_config": json.dumps( + {"quant_method": "fp8", "activation_scheme": "dynamic"} + ) + }, + ) + server_args = self._make_server_args( + transformer_weights_path=weights, + quantization="fp8", + ) + + with self.assertRaisesRegex(ValueError, "online --quantization"): + _resolve_quant_config( + hf_config={}, + server_args=server_args, + safetensors_list=[weights], + component_model_path="/base", + ) + + @patch( + "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", + return_value=None, + ) + def test_undeclared_quantized_weight_override_fails_closed(self, _build_nvfp4): + with tempfile.TemporaryDirectory() as directory: + weights = f"{directory}/model.safetensors" + save_file( + { + "block.weight": torch.ones((2, 2)), + "block.weight_scale": torch.ones(2), + }, + weights, + ) + server_args = self._make_server_args(transformer_weights_path=weights) + + with self.assertRaisesRegex(ValueError, "no supported native"): + _resolve_quant_config( + hf_config={}, + server_args=server_args, + safetensors_list=[weights], + component_model_path="/base", + ) + def test_autoround_config_is_inferred_and_remapped_to_native_prefixes(self): layer_config = { "bits": 4, @@ -212,7 +339,11 @@ class TestTransformerQuantHelpers(unittest.TestCase): pipeline_config=SimpleNamespace( dit_precision="bf16", dit_config=SimpleNamespace( - arch_config=SimpleNamespace(param_names_mapping={}) + arch_config=SimpleNamespace( + param_names_mapping={}, + reverse_param_names_mapping={}, + quant_ignore_remap={}, + ) ), ), nunchaku_config=None, @@ -968,6 +1099,7 @@ class TestTransformerQuantHelpers(unittest.TestCase): "config": json.dumps({"_class_name": _FakeFluxTransformer.__name__}) } with tempfile.NamedTemporaryFile(suffix=".safetensors") as f: + save_file({"block.weight": torch.ones((2, 2))}, f.name) nunchaku_config = NunchakuConfig(transformer_weights_path=f.name) server_args = self._make_server_args( transformer_weights_path=nunchaku_config.transformer_weights_path,