[diffusion] chore: detect quantized transformer replacements (#36916)

This commit is contained in:
Mick
2026-08-31 13:55:50 +08:00
committed by GitHub
parent 3a6ed55999
commit 28690f5aa5
4 changed files with 292 additions and 20 deletions
@@ -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
@@ -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
@@ -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
@@ -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,