From 0dd66def7c50081034ea45dc75a9ea2f01608894 Mon Sep 17 00:00:00 2001 From: Mick Date: Thu, 3 Sep 2026 09:27:22 +0800 Subject: [PATCH] [chore] harden checkpoint quantization metadata parsing (#36922) --- .../test/unit/test_image_encoder_loader.py | 7 +- .../model_loader/checkpoint_quantization.py | 64 +++++++++---------- .../unit/model_loader/test_modelopt_loader.py | 7 +- .../unit/test_checkpoint_quantization.py | 35 ++-------- 4 files changed, 45 insertions(+), 68 deletions(-) diff --git a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py index ec629ee37..e0c4ef961 100644 --- a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py @@ -4,6 +4,7 @@ from unittest import mock import torch from torch import nn +from transformers import PretrainedConfig from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( @@ -135,7 +136,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): class TestImageEncoderNativeLoading(unittest.TestCase): def test_bnb4_uses_shared_transformers_path_and_image_precision(self): - component_config = SimpleNamespace( + component_config = PretrainedConfig( is_encoder_decoder=False, architectures=["CLIPVisionModelWithProjection"], quantization_config={ @@ -193,7 +194,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): ) def test_explicit_offload_is_rejected_before_transformers_load(self): - component_config = SimpleNamespace( + component_config = PretrainedConfig( is_encoder_decoder=False, architectures=["ThirdPartyVisionModel"], quantization_config={"quant_method": "fp8"}, @@ -236,7 +237,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): def to(self, *args, **kwargs): raise AssertionError("quantized component must not be moved again") - component_config = SimpleNamespace( + component_config = PretrainedConfig( is_encoder_decoder=False, architectures=["ThirdPartyVisionModel"], quantization_config={"quant_method": "fp8"}, diff --git a/python/sglang/srt/model_loader/checkpoint_quantization.py b/python/sglang/srt/model_loader/checkpoint_quantization.py index 3578e4f1b..21bddf73c 100644 --- a/python/sglang/srt/model_loader/checkpoint_quantization.py +++ b/python/sglang/srt/model_loader/checkpoint_quantization.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 -"""Pure-data helpers for quantization metadata in Hugging Face configs.""" +"""Read quantization metadata declared by Hugging Face configurations.""" from __future__ import annotations @@ -8,6 +8,8 @@ from copy import deepcopy from dataclasses import dataclass from typing import Any, Literal, Mapping, TypeAlias +from transformers import PretrainedConfig + __all__ = [ "CheckpointQuantSpec", "QuantMetadataSource", @@ -20,11 +22,12 @@ QuantMetadataSource: TypeAlias = Literal[ "text_config.quantization_config", "compression_config", ] +ConfigMapping: TypeAlias = Mapping[str, Any] | PretrainedConfig @dataclass(slots=True) class CheckpointQuantSpec: - """Quantization metadata declared by a checkpoint. + """Quantization metadata declared by a checkpoint configuration. ``declared_method`` preserves ``quant_method`` verbatim and is never inferred from backend-specific fields. This intentionally contains no runtime @@ -36,54 +39,48 @@ class CheckpointQuantSpec: source: QuantMetadataSource -def _get_field(config: object, name: str) -> Any: - if isinstance(config, Mapping): - return config.get(name) - return getattr(config, name, None) - - -def _to_metadata_dict(value: object, source: QuantMetadataSource) -> dict[str, Any]: +def _as_config_mapping(value: ConfigMapping, source: str) -> Mapping[str, Any]: if isinstance(value, Mapping): - return deepcopy(dict(value)) - - to_dict = getattr(value, "to_dict", None) - if callable(to_dict): - metadata = to_dict() - if isinstance(metadata, Mapping): - return deepcopy(dict(metadata)) - + return value + if isinstance(value, PretrainedConfig): + return value.to_dict() raise TypeError( - f"{source} must be a mapping or expose to_dict(), " + f"{source} must be a mapping or transformers.PretrainedConfig, " f"got {type(value).__name__}" ) def _select_hf_quant_metadata( - hf_config: object, + hf_config: ConfigMapping, ) -> tuple[QuantMetadataSource, object] | None: - value = _get_field(hf_config, "quantization_config") + config = _as_config_mapping(hf_config, "HF config") + value = config.get("quantization_config") if value is not None: return "quantization_config", value - text_config = _get_field(hf_config, "text_config") - value = _get_field(text_config, "quantization_config") - if value is not None: - return "text_config.quantization_config", value + text_config = config.get("text_config") + if text_config is not None: + text_config_mapping = _as_config_mapping(text_config, "text_config") + value = text_config_mapping.get("quantization_config") + if value is not None: + return "text_config.quantization_config", value - value = _get_field(hf_config, "compression_config") + value = config.get("compression_config") if value is not None: return "compression_config", value return None -def resolve_checkpoint_quant_spec(hf_config: object) -> CheckpointQuantSpec | None: - """Resolve checkpoint quantization metadata from an HF config. +def resolve_checkpoint_quant_spec( + hf_config: ConfigMapping, +) -> CheckpointQuantSpec | None: + """Resolve quantization metadata from an HF configuration. - The lookup order matches SRT's checkpoint loader: top-level + The lookup order matches both serving runtimes: top-level ``quantization_config``, the text sub-config used by some multimodal - checkpoints, then ``compression_config``. The returned metadata is deep-copied - so callers can attach runtime-only fields without mutating the HF config. + checkpoints, then ``compression_config``. Returned metadata is deep-copied + so callers can attach runtime-only fields without mutating the source config. """ selected = _select_hf_quant_metadata(hf_config) @@ -91,10 +88,11 @@ def resolve_checkpoint_quant_spec(hf_config: object) -> CheckpointQuantSpec | No return None source, value = selected - config = _to_metadata_dict(value, source) - declared_method = config.get("quant_method") + config = _as_config_mapping(value, source) + copied_config = deepcopy(dict(config)) + declared_method = copied_config.get("quant_method") return CheckpointQuantSpec( declared_method=(declared_method if isinstance(declared_method, str) else None), - config=config, + config=copied_config, source=source, ) diff --git a/test/registered/unit/model_loader/test_modelopt_loader.py b/test/registered/unit/model_loader/test_modelopt_loader.py index dc5eee372..613960f5d 100644 --- a/test/registered/unit/model_loader/test_modelopt_loader.py +++ b/test/registered/unit/model_loader/test_modelopt_loader.py @@ -14,6 +14,7 @@ from unittest.mock import MagicMock, patch import torch import torch.nn as nn +from transformers import PretrainedConfig from sglang.srt.configs.device_config import DeviceConfig from sglang.srt.configs.load_config import LoadConfig @@ -657,7 +658,7 @@ class TestModelOptFp4LoaderSelection(CustomTestCase): quantization="modelopt_fp4", is_draft_model=True, is_draft_quantization_explicit=is_explicit, - hf_config=SimpleNamespace( + hf_config=PretrainedConfig( quantization_config={ "quant_algo": "NVFP4", "group_size": 16, @@ -757,7 +758,7 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase): with self.subTest(inline_config=inline_config): model_config = SimpleNamespace( quantization="modelopt_mixed", - hf_config=SimpleNamespace( + hf_config=PretrainedConfig( quantization_config=inline_config, ), model_path=model_path, @@ -787,7 +788,7 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase): } model_config = SimpleNamespace( quantization="modelopt_mixed", - hf_config=SimpleNamespace( + hf_config=PretrainedConfig( quantization_config={ "quant_method": "modelopt_mixed", "quant_algo": "MIXED_PRECISION", diff --git a/test/registered/unit/test_checkpoint_quantization.py b/test/registered/unit/test_checkpoint_quantization.py index 8b6b11c84..5b26d3083 100644 --- a/test/registered/unit/test_checkpoint_quantization.py +++ b/test/registered/unit/test_checkpoint_quantization.py @@ -2,6 +2,8 @@ import unittest +from transformers import PretrainedConfig + from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.model_loader.checkpoint_quantization import ( @@ -14,19 +16,6 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=5, suite="base-a-test-cpu") -class _ConfigObject: - def __init__(self, **values): - self.__dict__.update(values) - - -class _QuantConfigObject: - def __init__(self, values): - self._values = values - - def to_dict(self): - return self._values - - class TestResolveCheckpointQuantSpec(CustomTestCase): def test_modelopt_quant_algo_canonicalization(self): cases = { @@ -75,9 +64,9 @@ class TestResolveCheckpointQuantSpec(CustomTestCase): ), ) - def test_text_config_fallback_supports_config_objects(self): - config = _ConfigObject( - text_config=_ConfigObject( + def test_text_config_fallback_supports_pretrained_configs(self): + config = PretrainedConfig( + text_config=PretrainedConfig( quantization_config={"quant_method": "gptq", "bits": 4} ), compression_config={"quant_method": "compressed-tensors"}, @@ -90,7 +79,7 @@ class TestResolveCheckpointQuantSpec(CustomTestCase): self.assertEqual(spec.source, "text_config.quantization_config") def test_compression_config_fallback(self): - config = _ConfigObject( + config = PretrainedConfig( compression_config={"quant_method": "compressed-tensors"} ) @@ -114,18 +103,6 @@ class TestResolveCheckpointQuantSpec(CustomTestCase): self.assertIsNone(spec.declared_method) self.assertEqual(spec.config["quant_algo"], "FP8") - def test_quant_config_object_is_converted(self): - config = _ConfigObject( - quantization_config=_QuantConfigObject( - {"quant_method": "bitsandbytes", "load_in_4bit": True} - ) - ) - - spec = resolve_checkpoint_quant_spec(config) - - self.assertIsNotNone(spec) - self.assertEqual(spec.config["load_in_4bit"], True) - def test_lookup_priority_matches_srt_loader(self): config = { "quantization_config": {},