From d1ce017665aa66990939a8c73a92d8d17be00a2b Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 29 Aug 2026 14:26:50 +0800 Subject: [PATCH] [diffusion] feat: delegate recognized quantized components to transformers (#36902) --- docs/docs/sglang-diffusion/quantization.mdx | 27 ++-- .../component_loaders/component_loader.py | 65 +++++--- .../component_loaders/text_encoder_loader.py | 33 ++++- .../test/unit/test_image_encoder_loader.py | 140 +++++++++++++++++- .../test/unit/test_text_encoder_loader.py | 40 ++++- 5 files changed, 260 insertions(+), 45 deletions(-) diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index b526a0059..8bad3ff07 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -72,7 +72,7 @@ paths: | Component path | Quantized checkpoint behavior | | --- | --- | | `transformer`, `transformer_2`, `unconditional_transformer`, `audio_dit`, `video_dit` | Uses the SGLang transformer quantization adapters documented below. | -| `text_encoder*`, `image_encoder*` | Passes detected metadata to the registered native encoder. Loading continues only when that implementation constructs compatible quantized layers; unknown or incompatible combinations fail closed. | +| `text_encoder*`, `image_encoder*` | Prefers a compatible registered native encoder. Otherwise, a standard top-level `quantization_config` recognized by the installed Transformers version is delegated to its `from_pretrained` path; unsupported formats and metadata locations fail closed. | | `vae`, `video_vae`, `audio_vae` | A standard top-level Diffusers `quantization_config` is delegated to `AutoModel.from_pretrained`. Native-only VAEs and nested/compression metadata fail closed. | | Library-managed Transformers or Diffusers components | Delegates to the upstream `from_pretrained` path and inherits its format support and validation behavior. The local PE model uses this path; compatible formats remain model-specific. | | Native auxiliary components that load raw state dicts | Quantized checkpoints are rejected before model construction until that component has a quantized materialization implementation. This includes connectors, duration heads, bridges, diffusion decoders, sound tokenizers, spatial upsamplers, and vocoders. | @@ -466,13 +466,15 @@ sglang generate \ Each pattern is matched against the full layer prefix (e.g. `layers.0.attention.to_q`). A layer is skipped and left unquantized if its prefix contains any of the given patterns. -## Transformers Component BnB4 +## Transformers-managed Quantized Components -Model components that already have a native Transformers loading path can load -serialized BitsAndBytes 4-bit checkpoints with a standard top-level -`quantization_config`. Plain checkpoints keep using an available native SGLang -implementation. For example, replace FLUX's T5 component with the official -Diffusers checkpoint: +Model components that already have a native Transformers loading path delegate +self-describing checkpoints whose standard top-level `quantization_config` is +recognized by the installed Transformers version. Successful loading still +depends on that backend's optional dependencies, platform, and model/checkpoint +compatibility. An available native SGLang implementation remains preferred when +it can restore the same format. For example, replace FLUX's T5 component with +the official BitsAndBytes checkpoint: ```bash Command sglang serve \ @@ -481,11 +483,12 @@ sglang serve \ diffusers/FLUX.1-dev-bnb-4bit/text_encoder_2 ``` -The quantized component must stay resident. SGLang rejects component or -layerwise offload, nonstandard metadata locations, and native-only component -fallbacks for this path instead of silently changing the checkpoint contract. -Diffusion DiT components declared under the Diffusers library use the separate -quantization backends documented above. +Transformers owns format validation and places the quantized component directly +on its resident device. SGLang rejects component/layerwise offload, FSDP, +nonstandard metadata locations, unsupported upstream formats, and native-only +fallbacks instead of silently changing the checkpoint contract. Diffusion DiT +components declared under the Diffusers library use the separate quantization +backends documented above. ## Validated ModelOpt Checkpoints diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index ac3360d81..9ee324480 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -19,6 +19,7 @@ from transformers import ( AutoTokenizer, PretrainedConfig, ) +from transformers.quantizers import AutoHfQuantizer from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.layers.attention.selector import ( @@ -32,6 +33,7 @@ from sglang.multimodal_gen.runtime.loader.utils import ( get_memory_usage_of_component, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( + RESIDENT, ComponentResidencyError, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import ( @@ -65,15 +67,15 @@ class NativeComponentLoaderRequired(RuntimeError): """The customized loader must defer to the native library loader.""" -def uses_native_transformers_bnb4(config: object, component_name: str) -> bool: - """Validate a serialized BnB4 checkpoint owned by Transformers.""" +def uses_native_transformers_quantization(config: object, component_name: str) -> bool: + """Validate quantization metadata that Transformers can restore itself.""" try: quant_spec = resolve_checkpoint_quant_spec(config) except (TypeError, ValueError) as error: raise ComponentCheckpointUnsupportedError( f"Cannot parse checkpoint quantization for {component_name!r}: {error}" ) from error - if quant_spec is None or quant_spec.declared_method != "bitsandbytes": + if quant_spec is None: return False if quant_spec.source != "quantization_config": raise ComponentCheckpointUnsupportedError( @@ -82,16 +84,18 @@ def uses_native_transformers_bnb4(config: object, component_name: str) -> bool: f"got metadata from {quant_spec.source!r}" ) - load_in_4bit = quant_spec.config.get( - "load_in_4bit", quant_spec.config.get("_load_in_4bit") - ) - load_in_8bit = quant_spec.config.get( - "load_in_8bit", quant_spec.config.get("_load_in_8bit", False) - ) - if load_in_4bit is not True or load_in_8bit is True: + try: + supported = AutoHfQuantizer.supports_quant_method(dict(quant_spec.config)) + except (TypeError, ValueError) as error: raise ComponentCheckpointUnsupportedError( - f"Transformers-managed {component_name!r} quantization supports only " - "serialized BitsAndBytes 4-bit checkpoints" + f"Cannot configure Transformers-managed quantization for " + f"{component_name!r}: {error}" + ) from error + if not supported: + method = quant_spec.declared_method or "unspecified" + raise ComponentCheckpointUnsupportedError( + f"Transformers does not support quant_method={method!r} declared by " + f"{component_name!r}" ) return True @@ -147,6 +151,7 @@ class ComponentLoader(ABC): def __init__(self, device=None) -> None: self.device = device self.component_architecture: str | None = None + self._native_load_manages_placement = False @staticmethod def target_device(component_starts_on_cpu: bool) -> torch.device: @@ -230,6 +235,7 @@ class ComponentLoader(ABC): If all of the above methods failed, an error will be thrown """ + self._native_load_manages_placement = False component_quantization = server_args.component_quantizations.get(component_name) if ( component_quantization is not None @@ -317,7 +323,10 @@ class ComponentLoader(ABC): else: if isinstance(component, nn.Module): component = component.eval() - if not is_fsdp_managed_module(component): + if ( + not is_fsdp_managed_module(component) + and not self._native_load_manages_placement + ): component = component.to( self.target_device( server_args.should_start_component_on_cpu(component_name) @@ -356,16 +365,38 @@ class ComponentLoader(ABC): load_kwargs["torch_dtype"] = precision if transformers_or_diffusers == "transformers": + self._native_load_manages_placement = False config = get_hf_config( component_model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, ) - if uses_native_transformers_bnb4(config, component_name or "component"): - server_args.require_component_resident( - component_name or "component", - feature_name="Transformers bitsandbytes component", + if uses_native_transformers_quantization( + config, component_name or "component" + ): + resolved_component_name = component_name or "component" + explicit_residency = server_args.explicit_residency_mode( + resolved_component_name ) + if explicit_residency is not None and explicit_residency != RESIDENT: + raise ComponentCheckpointUnsupportedError( + "Transformers-managed quantized component " + f"{resolved_component_name!r} requires resident placement; " + f"got explicit mode {explicit_residency!r}" + ) + server_args.require_component_resident( + resolved_component_name, + feature_name="Transformers quantized component", + ) + if server_args.should_use_fsdp_for_component(resolved_component_name): + raise ComponentCheckpointUnsupportedError( + "Transformers-managed quantized components do not support " + "SGLang FSDP loading" + ) + load_kwargs["device_map"] = { + "": self.target_device(component_starts_on_cpu=False) + } + self._native_load_manages_placement = True model_class = self.resolve_native_transformers_model_class(config) return model_class.from_pretrained( component_model_path, diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index d9a3d2021..c2f5eeb25 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -58,7 +58,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp ComponentCheckpointUnsupportedError, ComponentLoader, NativeComponentLoaderRequired, - uses_native_transformers_bnb4, + uses_native_transformers_quantization, ) from sglang.multimodal_gen.runtime.loader.gguf_weights import ( gguf_weights_iterator, @@ -135,15 +135,23 @@ _TRANSFORMERS_ENCODER_ONLY_CLASSES = { } -def _delegate_standard_bnb4_to_transformers( +def _delegate_quantized_checkpoint_to_transformers( component_config: dict, component_name: str, + *, + methods: frozenset[str] | None = None, ) -> None: - """Use Transformers when it owns a standard serialized BnB4 checkpoint.""" - if uses_native_transformers_bnb4(component_config, component_name): + """Use Transformers when it owns the checkpoint's serialized format.""" + quant_spec = resolve_checkpoint_quant_spec(component_config) + if quant_spec is None or ( + methods is not None and quant_spec.declared_method not in methods + ): + return + if uses_native_transformers_quantization(component_config, component_name): + method = quant_spec.declared_method or "unspecified" raise NativeComponentLoaderRequired( - f"{component_name!r} delegates serialized bitsandbytes checkpoint " - "loading to Transformers" + f"{component_name!r} delegates serialized quant_method={method!r} " + "checkpoint loading to Transformers" ) @@ -268,9 +276,10 @@ def _configure_encoder_quantization( # themselves; running the generic lifecycle as well would process twice. return - _delegate_standard_bnb4_to_transformers( + _delegate_quantized_checkpoint_to_transformers( component_config, component_name, + methods=frozenset({"bitsandbytes"}), ) try: quant_config = _get_encoder_quant_config( @@ -280,6 +289,10 @@ def _configure_encoder_quantization( model_cls, ) except (KeyError, NotImplementedError, TypeError, ValueError) as error: + _delegate_quantized_checkpoint_to_transformers( + component_config, + component_name, + ) raise ComponentCheckpointUnsupportedError( f"Cannot configure checkpoint quantization for {component_name!r}: {error}" ) from error @@ -305,6 +318,10 @@ def _configure_encoder_quantization( ) quant_config = model_config.quant_config if quant_config is None: + _delegate_quantized_checkpoint_to_transformers( + component_config, + component_name, + ) return if not issubclass(model_cls, EncoderTensorParallelMixin): raise ComponentCheckpointUnsupportedError( @@ -327,7 +344,7 @@ def _resolve_and_configure_encoder_quantization( try: model_cls, _ = ModelRegistry.resolve_model_cls(architectures) except Exception as resolution_error: - _delegate_standard_bnb4_to_transformers( + _delegate_quantized_checkpoint_to_transformers( component_config, component_name, ) 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 17c8b4952..a8eaeaf82 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 @@ -3,10 +3,12 @@ from types import SimpleNamespace from unittest import mock import torch +from torch import nn from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, + NativeComponentLoaderRequired, ) from sglang.multimodal_gen.runtime.loader.component_loaders.image_encoder_loader import ( ImageEncoderLoader, @@ -14,6 +16,9 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.image_encoder_loader from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( _configure_encoder_quantization, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( + COMPONENT_OFFLOAD, +) from sglang.multimodal_gen.runtime.models.encoders.clip import CLIPVisionModel from sglang.srt.layers.quantization.fp8 import Fp8Config as SRTFp8Config @@ -76,8 +81,26 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}, ) - def test_unknown_quantized_architecture_does_not_fall_back(self): + def test_unknown_transformers_quantized_architecture_falls_back(self): config = self._component_config("UnknownVisionModel", quantized=True) + with self._config_patch(config): + self._load() + self.load_native.assert_called_once() + + def test_native_only_quantized_architecture_does_not_fall_back(self): + self.server_args.pipeline_config.native_only_components = ("image_encoder",) + config = self._component_config("UnknownVisionModel", quantized=True) + with self._config_patch(config), self.assertRaises( + NativeComponentLoaderRequired + ): + self._load() + self.load_native.assert_not_called() + + def test_unknown_unsupported_quantized_architecture_does_not_fall_back(self): + config = { + "architectures": ["UnknownVisionModel"], + "quantization_config": {"quant_method": "not-a-format"}, + } with self._config_patch(config), self.assertRaises( ComponentCheckpointUnsupportedError ): @@ -107,7 +130,9 @@ class TestImageEncoderNativeLoading(unittest.TestCase): ) server_args = SimpleNamespace( pipeline_config=SimpleNamespace(image_encoder_precision="bf16"), + explicit_residency_mode=mock.Mock(return_value=None), require_component_resident=mock.Mock(), + should_use_fsdp_for_component=mock.Mock(return_value=False), revision=None, trust_remote_code=False, ) @@ -121,6 +146,10 @@ class TestImageEncoderNativeLoading(unittest.TestCase): loader, "resolve_native_transformers_model_class", return_value=model_class, + ), mock.patch.object( + loader, + "target_device", + return_value=torch.device("cuda:0"), ): component = loader.load_native( "/model/image_encoder", @@ -132,7 +161,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): self.assertIs(component, loaded_encoder) server_args.require_component_resident.assert_called_once_with( "image_encoder", - feature_name="Transformers bitsandbytes component", + feature_name="Transformers quantized component", ) model_class.from_pretrained.assert_called_once_with( "/model/image_encoder", @@ -140,4 +169,111 @@ class TestImageEncoderNativeLoading(unittest.TestCase): trust_remote_code=False, revision=None, torch_dtype=torch.bfloat16, + device_map={"": torch.device("cuda:0")}, ) + + def test_explicit_offload_is_rejected_before_transformers_load(self): + component_config = SimpleNamespace( + is_encoder_decoder=False, + architectures=["ThirdPartyVisionModel"], + quantization_config={"quant_method": "fp8"}, + ) + model_class = SimpleNamespace(from_pretrained=mock.Mock()) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace(image_encoder_precision="bf16"), + explicit_residency_mode=mock.Mock(return_value=COMPONENT_OFFLOAD), + require_component_resident=mock.Mock(), + should_use_fsdp_for_component=mock.Mock(return_value=False), + revision=None, + trust_remote_code=False, + ) + loader = ImageEncoderLoader() + + with mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.get_hf_config", + return_value=component_config, + ), mock.patch.object( + loader, + "resolve_native_transformers_model_class", + return_value=model_class, + ), self.assertRaisesRegex( + ComponentCheckpointUnsupportedError, "requires resident placement" + ): + loader.load_native( + "/model/image_encoder", + server_args, + "transformers", + "image_encoder", + ) + + model_class.from_pretrained.assert_not_called() + server_args.require_component_resident.assert_not_called() + + def test_full_loader_does_not_move_quantized_component_again(self): + class RejectMoveModule(nn.Module): + def to(self, *args, **kwargs): + raise AssertionError("quantized component must not be moved again") + + component_config = SimpleNamespace( + is_encoder_decoder=False, + architectures=["ThirdPartyVisionModel"], + quantization_config={"quant_method": "fp8"}, + ) + loaded_encoder = RejectMoveModule() + model_class = SimpleNamespace( + from_pretrained=mock.Mock(return_value=loaded_encoder) + ) + server_args = SimpleNamespace( + component_quantizations={}, + pipeline_config=SimpleNamespace( + image_encoder_precision="bf16", + native_only_components=(), + ), + resolve_component_attention_backend=lambda _name: (None, None), + explicit_residency_mode=lambda _name: None, + require_component_resident=mock.Mock(), + should_use_fsdp_for_component=lambda _name: False, + should_start_component_on_cpu=lambda _name: False, + revision=None, + trust_remote_code=False, + ) + loader = ImageEncoderLoader() + + with mock.patch.object( + loader, + "load_customized", + side_effect=NativeComponentLoaderRequired("use Transformers"), + ), mock.patch.object( + loader, + "resolve_native_transformers_model_class", + return_value=model_class, + ), mock.patch.object( + loader, + "target_device", + return_value=torch.device("cuda:0"), + ), mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.get_hf_config", + return_value=component_config, + ), mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.current_platform.get_available_gpu_memory", + return_value=10.0, + ), mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.get_memory_usage_of_component", + return_value=0.0, + ), mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.format_component_residency", + return_value="resident", + ): + component, _ = loader.load( + "/model/image_encoder", + server_args, + "image_encoder", + "transformers", + ) + + self.assertIs(component, loaded_encoder) diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py index c75bac370..371876947 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py @@ -151,7 +151,9 @@ class TestTextEncoderClassResolution(unittest.TestCase): ) server_args = SimpleNamespace( pipeline_config=SimpleNamespace(text_encoder_precisions=["bf16"]), + explicit_residency_mode=mock.Mock(return_value=None), require_component_resident=mock.Mock(), + should_use_fsdp_for_component=mock.Mock(return_value=False), revision=None, trust_remote_code=False, ) @@ -162,16 +164,21 @@ class TestTextEncoderClassResolution(unittest.TestCase): } } + loader = TextEncoderLoader() with mock.patch.object( TextEncoderLoader, "resolve_native_transformers_model_class", return_value=transformers_model_class, + ), mock.patch.object( + loader, + "target_device", + return_value=torch.device("cuda:0"), ), mock.patch( "sglang.multimodal_gen.runtime.loader.component_loaders." "component_loader.get_hf_config", return_value=component_config, ): - encoder = TextEncoderLoader().load_native( + encoder = loader.load_native( "/model/text_encoder", server_args, "transformers", @@ -181,7 +188,7 @@ class TestTextEncoderClassResolution(unittest.TestCase): self.assertIs(encoder, loaded_encoder) server_args.require_component_resident.assert_called_once_with( "text_encoder", - feature_name="Transformers bitsandbytes component", + feature_name="Transformers quantized component", ) transformers_model_class.from_pretrained.assert_called_once_with( "/model/text_encoder", @@ -189,6 +196,7 @@ class TestTextEncoderClassResolution(unittest.TestCase): trust_remote_code=False, revision=None, torch_dtype=torch.bfloat16, + device_map={"": torch.device("cuda:0")}, ) @@ -712,7 +720,7 @@ class TestTextEncoderQuantization(unittest.TestCase): ): with self.subTest(architecture=architecture), self.assertRaisesRegex( NativeComponentLoaderRequired, - "delegates serialized bitsandbytes checkpoint loading to Transformers", + "delegates serialized quant_method='bitsandbytes' checkpoint", ): _resolve_and_configure_encoder_quantization( SimpleNamespace(architectures=[architecture], quant_config=None), @@ -742,10 +750,10 @@ class TestTextEncoderQuantization(unittest.TestCase): "text_encoder", ) - def test_rejects_bitsandbytes_8bit(self): + def test_bitsandbytes_8bit_delegates_to_transformers(self): with self.assertRaisesRegex( - ComponentCheckpointUnsupportedError, - "supports only serialized BitsAndBytes 4-bit checkpoints", + NativeComponentLoaderRequired, + "delegates serialized quant_method='bitsandbytes' checkpoint", ): _resolve_and_configure_encoder_quantization( SimpleNamespace( @@ -763,6 +771,26 @@ class TestTextEncoderQuantization(unittest.TestCase): "text_encoder", ) + def test_unknown_fp8_architecture_delegates_to_transformers(self): + with self.assertRaisesRegex( + NativeComponentLoaderRequired, + "delegates serialized quant_method='fp8' checkpoint", + ): + _resolve_and_configure_encoder_quantization( + SimpleNamespace( + architectures=["ThirdPartyTextEncoder"], quant_config=None + ), + { + "quantization_config": { + "quant_method": "fp8", + "activation_scheme": "dynamic", + } + }, + "/model/text_encoder", + "/model/text_encoder", + "text_encoder", + ) + def test_model_managed_quantization_bypasses_generic_lifecycle(self): model_config = SimpleNamespace(quant_config=None) with mock.patch.object(