[diffusion] feat: load serialized bnb4 components with transformers (#35945)
This commit is contained in:
@@ -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
|
||||||
):
|
):
|
||||||
|
|||||||
+43
-64
@@ -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(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,
|
||||||
)
|
)
|
||||||
self.assertIs(cls, transformers.AutoModel)
|
|
||||||
|
|
||||||
|
|
||||||
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",
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user