quant: extract shared checkpoint quant metadata resolver (#35172)

This commit is contained in:
Mick
2026-08-19 08:26:41 +08:00
committed by GitHub
parent e73201e462
commit ef490853bb
4 changed files with 332 additions and 15 deletions
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
import unittest
from sglang.srt.model_loader.checkpoint_quantization import (
CheckpointQuantSpec,
resolve_checkpoint_quant_spec,
)
from sglang.test.ci.ci_register import register_cpu_ci
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_top_level_quantization_config(self):
config = {
"quantization_config": {
"quant_method": "fp8",
"activation_scheme": "dynamic",
}
}
spec = resolve_checkpoint_quant_spec(config)
self.assertEqual(
spec,
CheckpointQuantSpec(
declared_method="fp8",
config={"quant_method": "fp8", "activation_scheme": "dynamic"},
source="quantization_config",
),
)
def test_text_config_fallback_supports_config_objects(self):
config = _ConfigObject(
text_config=_ConfigObject(
quantization_config={"quant_method": "gptq", "bits": 4}
),
compression_config={"quant_method": "compressed-tensors"},
)
spec = resolve_checkpoint_quant_spec(config)
self.assertIsNotNone(spec)
self.assertEqual(spec.declared_method, "gptq")
self.assertEqual(spec.source, "text_config.quantization_config")
def test_compression_config_fallback(self):
config = _ConfigObject(
compression_config={"quant_method": "compressed-tensors"}
)
spec = resolve_checkpoint_quant_spec(config)
self.assertIsNotNone(spec)
self.assertEqual(spec.declared_method, "compressed-tensors")
self.assertEqual(spec.source, "compression_config")
def test_modelopt_quant_algo_does_not_infer_declared_method(self):
config = {
"quantization_config": {
"quant_algo": "FP8",
"exclude_modules": ["lm_head"],
}
}
spec = resolve_checkpoint_quant_spec(config)
self.assertIsNotNone(spec)
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": {},
"text_config": {"quantization_config": {"quant_method": "gptq"}},
"compression_config": {"quant_method": "compressed-tensors"},
}
spec = resolve_checkpoint_quant_spec(config)
self.assertIsNotNone(spec)
self.assertEqual(spec.config, {})
self.assertEqual(spec.source, "quantization_config")
def test_metadata_is_deep_copied(self):
metadata = {"quant_method": "fp8", "modules_to_not_convert": ["lm_head"]}
spec = resolve_checkpoint_quant_spec({"quantization_config": metadata})
self.assertIsNotNone(spec)
spec.config["modules_to_not_convert"].append("embed_tokens")
self.assertEqual(metadata["modules_to_not_convert"], ["lm_head"])
def test_missing_metadata_returns_none(self):
self.assertIsNone(resolve_checkpoint_quant_spec({"model_type": "qwen3_vl"}))
def test_invalid_metadata_type_has_clear_error(self):
with self.assertRaisesRegex(TypeError, "quantization_config must be a mapping"):
resolve_checkpoint_quant_spec({"quantization_config": "fp8"})
if __name__ == "__main__":
unittest.main()