diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 88bed2a8e..2a94b8288 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -423,6 +423,27 @@ the native encoder explicitly supports. Text-encoder FP8 is approximate, is not enabled by default, and is rejected by MiniMax-H3's strict `quality="high"` deployment contract. +## Transformers Component BnB4 + +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: + +```bash Command +sglang serve \ + --model-path black-forest-labs/FLUX.1-dev \ + --component-paths.text_encoder_2 \ + 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. + ## Validated ModelOpt Checkpoints This section is the canonical support matrix for the thirteen published 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 f1755aa0c..931ac03bc 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 @@ -10,9 +10,15 @@ from abc import ABC from typing import Any, Type import torch +import transformers from diffusers import AutoModel from torch import nn -from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer +from transformers import ( + AutoImageProcessor, + AutoProcessor, + AutoTokenizer, + PretrainedConfig, +) from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.layers.attention.selector import ( @@ -55,6 +61,37 @@ 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.""" + 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": + return False + if quant_spec.source != "quantization_config": + raise ComponentCheckpointUnsupportedError( + f"Transformers-managed {component_name!r} quantization requires " + "a top-level quantization_config; " + 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: + raise ComponentCheckpointUnsupportedError( + f"Transformers-managed {component_name!r} quantization supports only " + "serialized BitsAndBytes 4-bit checkpoints" + ) + return True + + def _load_auto_tokenizer_with_roberta_processing_compat(*args, **kwargs): from tokenizers import processors @@ -301,14 +338,18 @@ class ComponentLoader(ABC): load_kwargs["torch_dtype"] = precision if transformers_or_diffusers == "transformers": - from transformers import AutoModel - config = get_hf_config( component_model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, ) - return AutoModel.from_pretrained( + if uses_native_transformers_bnb4(config, component_name or "component"): + server_args.require_component_resident( + component_name or "component", + feature_name="Transformers bitsandbytes component", + ) + model_class = self.resolve_native_transformers_model_class(config) + return model_class.from_pretrained( component_model_path, config=config, trust_remote_code=server_args.trust_remote_code, @@ -330,6 +371,9 @@ class ComponentLoader(ABC): else: raise ValueError(f"Unsupported library: {transformers_or_diffusers}") + def resolve_native_transformers_model_class(self, config: PretrainedConfig) -> type: + return transformers.AutoModel + def load_customized( self, component_model_path: str, server_args: ServerArgs, component_name: str ): 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 dd870b338..3a3d07d6b 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 @@ -7,7 +7,9 @@ from itertools import chain from typing import cast import torch +import transformers from torch import nn +from transformers import PretrainedConfig from transformers.utils import SAFE_WEIGHTS_INDEX_NAME from sglang.multimodal_gen.configs.models import EncoderConfig @@ -28,6 +30,8 @@ from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, ComponentLoader, + NativeComponentLoaderRequired, + uses_native_transformers_bnb4, ) from sglang.multimodal_gen.runtime.loader.utils import ( set_default_torch_dtype, @@ -58,13 +62,36 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( load_dict, ) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.runtime.utils.precision import precision_to_dtype from sglang.multimodal_gen.runtime.utils.quantization_utils import get_quant_config from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.srt.environ import envs logger = init_logger(__name__) +_TRANSFORMERS_ENCODER_ONLY_CLASSES = { + "T5EncoderModel": transformers.T5EncoderModel, + "T5Model": transformers.T5EncoderModel, + "T5ForConditionalGeneration": transformers.T5EncoderModel, + "UMT5EncoderModel": transformers.UMT5EncoderModel, + "UMT5Model": transformers.UMT5EncoderModel, + "UMT5ForConditionalGeneration": transformers.UMT5EncoderModel, + "MT5EncoderModel": transformers.MT5EncoderModel, + "MT5Model": transformers.MT5EncoderModel, + "MT5ForConditionalGeneration": transformers.MT5EncoderModel, +} + + +def _delegate_standard_bnb4_to_transformers( + component_config: dict, + component_name: str, +) -> None: + """Use Transformers when it owns a standard serialized BnB4 checkpoint.""" + if uses_native_transformers_bnb4(component_config, component_name): + raise NativeComponentLoaderRequired( + f"{component_name!r} delegates serialized bitsandbytes checkpoint " + "loading to Transformers" + ) + def _configure_encoder_quantization( model_config: EncoderConfig, @@ -79,12 +106,16 @@ def _configure_encoder_quantization( # themselves; running the generic lifecycle as well would process twice. return + _delegate_standard_bnb4_to_transformers( + component_config, + component_name, + ) try: quant_config = get_quant_config( component_config, component_model_path, ) - except (KeyError, ValueError) as error: + except (KeyError, TypeError, ValueError) as error: raise ComponentCheckpointUnsupportedError( f"Cannot configure checkpoint quantization for {component_name!r}: {error}" ) from error @@ -130,6 +161,10 @@ def _resolve_and_configure_encoder_quantization( try: model_cls, _ = ModelRegistry.resolve_model_cls(architectures) except Exception as resolution_error: + _delegate_standard_bnb4_to_transformers( + component_config, + component_name, + ) try: quant_config = get_quant_config(component_config, component_model_path) except Exception as quantization_error: @@ -259,43 +294,7 @@ class TextEncoderLoader(ComponentLoader): allow_patterns_overrides: list[str] | None = None """If defined, weights will load exclusively using these patterns.""" - def load_native( - self, - component_model_path: str, - server_args: ServerArgs, - transformers_or_diffusers: str, - component_name: str | None = None, - ): - if transformers_or_diffusers != "transformers": - return super().load_native( - component_model_path, - server_args, - transformers_or_diffusers, - component_name, - ) - - encoder_idx = ( - self._extract_encoder_index(component_name or "text_encoder_2") - if component_name - else 1 if component_model_path.rstrip("/").endswith("text_encoder_2") else 0 - ) - encoder_dtype = server_args.pipeline_config.text_encoder_precisions[encoder_idx] - dtype = precision_to_dtype( - encoder_dtype, - f"text_encoder_precisions[{encoder_idx}]", - ) - transformers_model_class = self._resolve_transformers_text_encoder_class( - component_model_path, server_args - ) - return transformers_model_class.from_pretrained( - component_model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - torch_dtype=dtype, - ) - - @staticmethod - def _resolve_transformers_text_encoder_class(component_model_path, server_args): + def resolve_native_transformers_model_class(self, config: PretrainedConfig) -> type: """Resolve the concrete transformers class for a text encoder. AutoModel maps encoder-decoder model types (e.g. T5/UMT5) to full @@ -305,32 +304,12 @@ class TextEncoderLoader(ComponentLoader): full seq2seq architecture to its encoder-only counterpart. Encoders that are not encoder-decoder keep using AutoModel unchanged. """ - import transformers - from transformers import AutoConfig, AutoModel - - try: - config = AutoConfig.from_pretrained( - component_model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - ) - except Exception: - return AutoModel - if getattr(config, "is_encoder_decoder", False): - encoder_only_map = { - "T5Model": "T5EncoderModel", - "T5ForConditionalGeneration": "T5EncoderModel", - "UMT5Model": "UMT5EncoderModel", - "UMT5ForConditionalGeneration": "UMT5EncoderModel", - "MT5Model": "MT5EncoderModel", - "MT5ForConditionalGeneration": "MT5EncoderModel", - } - for arch in getattr(config, "architectures", None) or []: - encoder_arch = encoder_only_map.get(arch, arch) - transformers_model_class = getattr(transformers, encoder_arch, None) - if isinstance(transformers_model_class, type): + if config.is_encoder_decoder: + for arch in config.architectures or []: + transformers_model_class = _TRANSFORMERS_ENCODER_ONLY_CLASSES.get(arch) + if transformers_model_class is not None: return transformers_model_class - return AutoModel + return transformers.AutoModel def _prepare_weights( self, 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 30b2fd775..b5746ad7b 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 @@ -2,6 +2,8 @@ import unittest from types import SimpleNamespace from unittest import mock +import torch + from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, @@ -71,3 +73,55 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): with self._config_patch(config): self._load() self.load_native.assert_called_once() + + +class TestImageEncoderNativeLoading(unittest.TestCase): + def test_bnb4_uses_shared_transformers_path_and_image_precision(self): + component_config = SimpleNamespace( + is_encoder_decoder=False, + architectures=["CLIPVisionModelWithProjection"], + quantization_config={ + "load_in_4bit": True, + "quant_method": "bitsandbytes", + }, + ) + loaded_encoder = object() + model_class = SimpleNamespace( + from_pretrained=mock.Mock(return_value=loaded_encoder) + ) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace(image_encoder_precision="bf16"), + require_component_resident=mock.Mock(), + 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, + ): + component = loader.load_native( + "/model/image_encoder", + server_args, + "transformers", + "image_encoder", + ) + + self.assertIs(component, loaded_encoder) + server_args.require_component_resident.assert_called_once_with( + "image_encoder", + feature_name="Transformers bitsandbytes component", + ) + model_class.from_pretrained.assert_called_once_with( + "/model/image_encoder", + config=component_config, + trust_remote_code=False, + revision=None, + torch_dtype=torch.bfloat16, + ) 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 542d95b7e..1bd7245c5 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 @@ -10,11 +10,13 @@ from sglang.multimodal_gen.runtime.layers.linear import LinearBase from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, + NativeComponentLoaderRequired, ) from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( TextEncoderLoader, _configure_encoder_quantization, _process_quantized_encoder_weights, + _resolve_and_configure_encoder_quantization, ) from sglang.multimodal_gen.runtime.models.encoders.base import ( CheckpointQuantizationCapability, @@ -33,18 +35,11 @@ class TestTextEncoderClassResolution(unittest.TestCase): module is used purely as a text encoder. """ - server_args = SimpleNamespace(trust_remote_code=False, revision=None) - def _resolve(self, is_encoder_decoder, architectures): config = SimpleNamespace( is_encoder_decoder=is_encoder_decoder, architectures=architectures ) - with mock.patch.object( - transformers.AutoConfig, "from_pretrained", return_value=config - ): - return TextEncoderLoader._resolve_transformers_text_encoder_class( - "dummy/path", self.server_args - ) + return TextEncoderLoader().resolve_native_transformers_model_class(config) def test_umt5_encoder_decoder_uses_encoder_only_class(self): self.assertIs( @@ -83,16 +78,52 @@ class TestTextEncoderClassResolution(unittest.TestCase): def test_unknown_architecture_falls_back_to_automodel(self): self.assertIs(self._resolve(True, ["NotARealClass"]), transformers.AutoModel) - def test_config_load_failure_falls_back_to_automodel(self): + def test_bitsandbytes_native_load_requires_resident_encoder(self): + loaded_encoder = nn.Linear(1, 1) + transformers_model_class = SimpleNamespace( + from_pretrained=mock.Mock(return_value=loaded_encoder) + ) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace(text_encoder_precisions=["bf16"]), + require_component_resident=mock.Mock(), + revision=None, + trust_remote_code=False, + ) + component_config = { + "quantization_config": { + "load_in_4bit": True, + "quant_method": "bitsandbytes", + } + } + with mock.patch.object( - transformers.AutoConfig, - "from_pretrained", - side_effect=OSError("no config"), + TextEncoderLoader, + "resolve_native_transformers_model_class", + return_value=transformers_model_class, + ), mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.get_hf_config", + return_value=component_config, ): - cls = TextEncoderLoader._resolve_transformers_text_encoder_class( - "dummy/path", self.server_args + encoder = TextEncoderLoader().load_native( + "/model/text_encoder", + server_args, + "transformers", + "text_encoder", ) - self.assertIs(cls, transformers.AutoModel) + + self.assertIs(encoder, loaded_encoder) + server_args.require_component_resident.assert_called_once_with( + "text_encoder", + feature_name="Transformers bitsandbytes component", + ) + transformers_model_class.from_pretrained.assert_called_once_with( + "/model/text_encoder", + config=component_config, + trust_remote_code=False, + revision=None, + torch_dtype=torch.bfloat16, + ) class TestMiniMaxH3CheckpointFilter(unittest.TestCase): @@ -175,6 +206,69 @@ class TestTextEncoderQuantization(unittest.TestCase): "text_encoder", ) + def test_standard_bitsandbytes_delegates_to_transformers(self): + component_config = { + "quantization_config": { + "load_in_4bit": True, + "quant_method": "bitsandbytes", + } + } + + for architecture in ( + "T5EncoderModel", + "CLIPTextModel", + "ThirdPartyTextEncoder", + ): + with self.subTest(architecture=architecture), self.assertRaisesRegex( + NativeComponentLoaderRequired, + "delegates serialized bitsandbytes checkpoint loading to Transformers", + ): + _resolve_and_configure_encoder_quantization( + SimpleNamespace(architectures=[architecture], quant_config=None), + component_config, + "/model/text_encoder", + "text_encoder", + ) + self.get_quant_config.assert_not_called() + + def test_rejects_nonstandard_bitsandbytes_metadata_location(self): + with self.assertRaisesRegex( + ComponentCheckpointUnsupportedError, + "requires a top-level quantization_config", + ): + _configure_encoder_quantization( + SimpleNamespace(quant_config=None), + TextEncoder, + { + "compression_config": { + "load_in_4bit": True, + "quant_method": "bitsandbytes", + } + }, + "/model/text_encoder", + "text_encoder", + ) + + def test_rejects_bitsandbytes_8bit(self): + with self.assertRaisesRegex( + ComponentCheckpointUnsupportedError, + "supports only serialized BitsAndBytes 4-bit checkpoints", + ): + _resolve_and_configure_encoder_quantization( + SimpleNamespace( + architectures=["ThirdPartyTextEncoder"], quant_config=None + ), + { + "quantization_config": { + "load_in_4bit": False, + "load_in_8bit": True, + "quant_method": "bitsandbytes", + } + }, + "/model/text_encoder", + "text_encoder", + ) + def test_srt_backend_is_not_admitted_without_an_adapter(self): model_config = SimpleNamespace(quant_config=None) capability = CheckpointQuantizationCapability( @@ -207,7 +301,12 @@ class TestTextEncoderQuantization(unittest.TestCase): _configure_encoder_quantization( model_config, TextEncoder, - {}, + { + "quantization_config": { + "load_in_4bit": True, + "quant_method": "bitsandbytes", + } + }, "/model/text_encoder", "text_encoder", )