[diffusion] feat: load serialized bnb4 components with transformers (#35945)

This commit is contained in:
Mick
2026-08-22 15:33:53 +08:00
committed by GitHub
parent 90354326c7
commit b391ef171f
5 changed files with 281 additions and 84 deletions
@@ -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 not enabled by default, and is rejected by MiniMax-H3's strict
`quality="high"` deployment contract. `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 ## Validated ModelOpt Checkpoints
This section is the canonical support matrix for the thirteen published This section is the canonical support matrix for the thirteen published
@@ -10,9 +10,15 @@ from abc import ABC
from typing import Any, Type from typing import Any, Type
import torch import torch
import transformers
from diffusers import AutoModel from diffusers import AutoModel
from torch import nn 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.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.layers.attention.selector import ( 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.""" """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): def _load_auto_tokenizer_with_roberta_processing_compat(*args, **kwargs):
from tokenizers import processors from tokenizers import processors
@@ -301,14 +338,18 @@ class ComponentLoader(ABC):
load_kwargs["torch_dtype"] = precision load_kwargs["torch_dtype"] = precision
if transformers_or_diffusers == "transformers": if transformers_or_diffusers == "transformers":
from transformers import AutoModel
config = get_hf_config( config = get_hf_config(
component_model_path, component_model_path,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, 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, component_model_path,
config=config, config=config,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
@@ -330,6 +371,9 @@ class ComponentLoader(ABC):
else: else:
raise ValueError(f"Unsupported library: {transformers_or_diffusers}") raise ValueError(f"Unsupported library: {transformers_or_diffusers}")
def resolve_native_transformers_model_class(self, config: PretrainedConfig) -> type:
return transformers.AutoModel
def load_customized( def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str self, component_model_path: str, server_args: ServerArgs, component_name: str
): ):
@@ -7,7 +7,9 @@ from itertools import chain
from typing import cast from typing import cast
import torch import torch
import transformers
from torch import nn from torch import nn
from transformers import PretrainedConfig
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from sglang.multimodal_gen.configs.models import EncoderConfig 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 ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentCheckpointUnsupportedError, ComponentCheckpointUnsupportedError,
ComponentLoader, ComponentLoader,
NativeComponentLoaderRequired,
uses_native_transformers_bnb4,
) )
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
set_default_torch_dtype, set_default_torch_dtype,
@@ -58,13 +62,36 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
load_dict, load_dict,
) )
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger 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.runtime.utils.quantization_utils import get_quant_config
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
from sglang.srt.environ import envs from sglang.srt.environ import envs
logger = init_logger(__name__) 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( def _configure_encoder_quantization(
model_config: EncoderConfig, model_config: EncoderConfig,
@@ -79,12 +106,16 @@ def _configure_encoder_quantization(
# themselves; running the generic lifecycle as well would process twice. # themselves; running the generic lifecycle as well would process twice.
return return
_delegate_standard_bnb4_to_transformers(
component_config,
component_name,
)
try: try:
quant_config = get_quant_config( quant_config = get_quant_config(
component_config, component_config,
component_model_path, component_model_path,
) )
except (KeyError, ValueError) as error: except (KeyError, TypeError, ValueError) as error:
raise ComponentCheckpointUnsupportedError( raise ComponentCheckpointUnsupportedError(
f"Cannot configure checkpoint quantization for {component_name!r}: {error}" f"Cannot configure checkpoint quantization for {component_name!r}: {error}"
) from error ) from error
@@ -130,6 +161,10 @@ def _resolve_and_configure_encoder_quantization(
try: try:
model_cls, _ = ModelRegistry.resolve_model_cls(architectures) model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
except Exception as resolution_error: except Exception as resolution_error:
_delegate_standard_bnb4_to_transformers(
component_config,
component_name,
)
try: try:
quant_config = get_quant_config(component_config, component_model_path) quant_config = get_quant_config(component_config, component_model_path)
except Exception as quantization_error: except Exception as quantization_error:
@@ -259,43 +294,7 @@ class TextEncoderLoader(ComponentLoader):
allow_patterns_overrides: list[str] | None = None allow_patterns_overrides: list[str] | None = None
"""If defined, weights will load exclusively using these patterns.""" """If defined, weights will load exclusively using these patterns."""
def load_native( def resolve_native_transformers_model_class(self, config: PretrainedConfig) -> type:
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):
"""Resolve the concrete transformers class for a text encoder. """Resolve the concrete transformers class for a text encoder.
AutoModel maps encoder-decoder model types (e.g. T5/UMT5) to full 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 full seq2seq architecture to its encoder-only counterpart. Encoders that
are not encoder-decoder keep using AutoModel unchanged. are not encoder-decoder keep using AutoModel unchanged.
""" """
import transformers if config.is_encoder_decoder:
from transformers import AutoConfig, AutoModel for arch in config.architectures or []:
transformers_model_class = _TRANSFORMERS_ENCODER_ONLY_CLASSES.get(arch)
try: if transformers_model_class is not None:
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):
return transformers_model_class return transformers_model_class
return AutoModel return transformers.AutoModel
def _prepare_weights( def _prepare_weights(
self, self,
@@ -2,6 +2,8 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest import mock from unittest import mock
import torch
from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentCheckpointUnsupportedError, ComponentCheckpointUnsupportedError,
@@ -71,3 +73,55 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase):
with self._config_patch(config): with self._config_patch(config):
self._load() self._load()
self.load_native.assert_called_once() 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,
)
@@ -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.layers.quantization.fp8 import Fp8Config
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentCheckpointUnsupportedError, ComponentCheckpointUnsupportedError,
NativeComponentLoaderRequired,
) )
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
TextEncoderLoader, TextEncoderLoader,
_configure_encoder_quantization, _configure_encoder_quantization,
_process_quantized_encoder_weights, _process_quantized_encoder_weights,
_resolve_and_configure_encoder_quantization,
) )
from sglang.multimodal_gen.runtime.models.encoders.base import ( from sglang.multimodal_gen.runtime.models.encoders.base import (
CheckpointQuantizationCapability, CheckpointQuantizationCapability,
@@ -33,18 +35,11 @@ class TestTextEncoderClassResolution(unittest.TestCase):
module is used purely as a text encoder. module is used purely as a text encoder.
""" """
server_args = SimpleNamespace(trust_remote_code=False, revision=None)
def _resolve(self, is_encoder_decoder, architectures): def _resolve(self, is_encoder_decoder, architectures):
config = SimpleNamespace( config = SimpleNamespace(
is_encoder_decoder=is_encoder_decoder, architectures=architectures is_encoder_decoder=is_encoder_decoder, architectures=architectures
) )
with mock.patch.object( return TextEncoderLoader().resolve_native_transformers_model_class(config)
transformers.AutoConfig, "from_pretrained", return_value=config
):
return TextEncoderLoader._resolve_transformers_text_encoder_class(
"dummy/path", self.server_args
)
def test_umt5_encoder_decoder_uses_encoder_only_class(self): def test_umt5_encoder_decoder_uses_encoder_only_class(self):
self.assertIs( self.assertIs(
@@ -83,16 +78,52 @@ class TestTextEncoderClassResolution(unittest.TestCase):
def test_unknown_architecture_falls_back_to_automodel(self): def test_unknown_architecture_falls_back_to_automodel(self):
self.assertIs(self._resolve(True, ["NotARealClass"]), transformers.AutoModel) 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( with mock.patch.object(
transformers.AutoConfig, TextEncoderLoader,
"from_pretrained", "resolve_native_transformers_model_class",
side_effect=OSError("no config"), 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( encoder = TextEncoderLoader().load_native(
"dummy/path", self.server_args "/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): class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
@@ -175,6 +206,69 @@ class TestTextEncoderQuantization(unittest.TestCase):
"text_encoder", "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): def test_srt_backend_is_not_admitted_without_an_adapter(self):
model_config = SimpleNamespace(quant_config=None) model_config = SimpleNamespace(quant_config=None)
capability = CheckpointQuantizationCapability( capability = CheckpointQuantizationCapability(
@@ -207,7 +301,12 @@ class TestTextEncoderQuantization(unittest.TestCase):
_configure_encoder_quantization( _configure_encoder_quantization(
model_config, model_config,
TextEncoder, TextEncoder,
{}, {
"quantization_config": {
"load_in_4bit": True,
"quant_method": "bitsandbytes",
}
},
"/model/text_encoder", "/model/text_encoder",
"text_encoder", "text_encoder",
) )