[diffusion] refactor: gate native encoder quantized checkpoints (#35183)
Co-authored-by: Yiqi Yang <yangyiqi8787@gmail.com>
This commit is contained in:
@@ -42,6 +42,10 @@ from sglang.multimodal_gen.runtime.utils.precision import resolve_component_prec
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ComponentCheckpointUnsupportedError(ValueError):
|
||||||
|
"""A component checkpoint is unsupported and must not use native fallback."""
|
||||||
|
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
@@ -191,7 +195,7 @@ class ComponentLoader(ABC):
|
|||||||
component_attn_name,
|
component_attn_name,
|
||||||
)
|
)
|
||||||
source = "sgl-diffusion"
|
source = "sgl-diffusion"
|
||||||
except ComponentResidencyError:
|
except (ComponentCheckpointUnsupportedError, ComponentResidencyError):
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
if self.should_raise_customized_load_error(server_args, component_name):
|
if self.should_raise_customized_load_error(server_args, component_name):
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
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,
|
||||||
|
_resolve_and_configure_encoder_quantization,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import finalize_encoder_folding
|
from sglang.multimodal_gen.runtime.models.encoders.base import finalize_encoder_folding
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
@@ -34,6 +35,12 @@ class ImageEncoderLoader(TextEncoderLoader):
|
|||||||
|
|
||||||
encoder_config = server_args.pipeline_config.image_encoder_config
|
encoder_config = server_args.pipeline_config.image_encoder_config
|
||||||
encoder_config.update_model_arch(model_config)
|
encoder_config.update_model_arch(model_config)
|
||||||
|
_resolve_and_configure_encoder_quantization(
|
||||||
|
encoder_config,
|
||||||
|
model_config,
|
||||||
|
component_model_path,
|
||||||
|
component_name,
|
||||||
|
)
|
||||||
# real dims are populated now; resolve fold vs replicate
|
# real dims are populated now; resolve fold vs replicate
|
||||||
finalize_encoder_folding(
|
finalize_encoder_folding(
|
||||||
encoder_config,
|
encoder_config,
|
||||||
|
|||||||
+86
-30
@@ -26,6 +26,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
UnquantizedLinearMethod,
|
UnquantizedLinearMethod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
|
ComponentCheckpointUnsupportedError,
|
||||||
ComponentLoader,
|
ComponentLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
@@ -61,11 +62,12 @@ from sglang.srt.environ import envs
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _configure_text_encoder_quantization(
|
def _configure_encoder_quantization(
|
||||||
model_config: EncoderConfig,
|
model_config: EncoderConfig,
|
||||||
model_cls: type[nn.Module],
|
model_cls: type[nn.Module],
|
||||||
component_config: dict,
|
component_config: dict,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
|
component_name: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
if getattr(model_cls, "manages_checkpoint_quantization", False):
|
if getattr(model_cls, "manages_checkpoint_quantization", False):
|
||||||
# Preserve model-owned formats such as Ideogram's bitsandbytes state.
|
# Preserve model-owned formats such as Ideogram's bitsandbytes state.
|
||||||
@@ -73,27 +75,79 @@ def _configure_text_encoder_quantization(
|
|||||||
# themselves; running the generic lifecycle as well would process twice.
|
# themselves; running the generic lifecycle as well would process twice.
|
||||||
return
|
return
|
||||||
|
|
||||||
quant_config = get_quant_config(
|
try:
|
||||||
component_config,
|
quant_config = get_quant_config(
|
||||||
component_model_path,
|
component_config,
|
||||||
)
|
component_model_path,
|
||||||
|
)
|
||||||
|
except (KeyError, ValueError) as error:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Cannot configure checkpoint quantization for {component_name!r}: {error}"
|
||||||
|
) from error
|
||||||
model_config.quant_config = quant_config
|
model_config.quant_config = quant_config
|
||||||
if quant_config is None:
|
if quant_config is None:
|
||||||
return
|
return
|
||||||
if not issubclass(model_cls, TextEncoder):
|
if not issubclass(model_cls, EncoderTensorParallelMixin):
|
||||||
raise ValueError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
"A quantized text-encoder checkpoint requires an in-tree native "
|
f"A quantized {component_name!r} checkpoint requires an in-tree "
|
||||||
"TextEncoder; "
|
"native encoder; "
|
||||||
f"got {model_cls.__name__}"
|
f"got {model_cls.__name__}"
|
||||||
)
|
)
|
||||||
quant_method = quant_config.get_name()
|
|
||||||
supported_methods = model_cls.supported_checkpoint_quantization_methods
|
capability = model_cls.checkpoint_quantization_capability
|
||||||
if quant_method not in supported_methods:
|
if capability is None:
|
||||||
raise ValueError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"{model_cls.__name__} does not support text-encoder checkpoints "
|
f"{model_cls.__name__} does not support quantized checkpoints for "
|
||||||
f"quantized with {quant_method!r}; supported methods: "
|
f"{component_name!r}: no checkpoint quantization capability is declared"
|
||||||
f"{sorted(supported_methods)}"
|
|
||||||
)
|
)
|
||||||
|
if capability.backend != "diffusion":
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"{model_cls.__name__} declares the {capability.backend!r} checkpoint "
|
||||||
|
f"quantization backend for {component_name!r}, but the native encoder "
|
||||||
|
"loader currently supports only the 'diffusion' backend"
|
||||||
|
)
|
||||||
|
|
||||||
|
quant_method = quant_config.get_name()
|
||||||
|
if quant_method not in capability.methods:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"{model_cls.__name__} does not support {component_name!r} checkpoints "
|
||||||
|
f"quantized with {quant_method!r}; supported methods for the "
|
||||||
|
f"{capability.backend!r} backend: {sorted(capability.methods)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_and_configure_encoder_quantization(
|
||||||
|
model_config: EncoderConfig,
|
||||||
|
component_config: dict,
|
||||||
|
component_model_path: str,
|
||||||
|
component_name: str,
|
||||||
|
) -> type[nn.Module]:
|
||||||
|
architectures = getattr(model_config, "architectures", [])
|
||||||
|
try:
|
||||||
|
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||||
|
except Exception as resolution_error:
|
||||||
|
try:
|
||||||
|
quant_config = get_quant_config(component_config, component_model_path)
|
||||||
|
except Exception as quantization_error:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Cannot parse checkpoint quantization for {component_name!r}: "
|
||||||
|
f"{quantization_error}"
|
||||||
|
) from quantization_error
|
||||||
|
if quant_config is None:
|
||||||
|
raise
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"A quantized {component_name!r} checkpoint requires an in-tree "
|
||||||
|
f"native encoder; unsupported architectures: {architectures}"
|
||||||
|
) from resolution_error
|
||||||
|
|
||||||
|
_configure_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
model_cls,
|
||||||
|
component_config,
|
||||||
|
component_model_path,
|
||||||
|
component_name,
|
||||||
|
)
|
||||||
|
return model_cls
|
||||||
|
|
||||||
|
|
||||||
def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
||||||
@@ -118,9 +172,10 @@ def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
|||||||
return next(iter(devices), None)
|
return next(iter(devices), None)
|
||||||
|
|
||||||
|
|
||||||
def _process_quantized_text_encoder_weights(
|
def _process_quantized_encoder_weights(
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
process_device: torch.device,
|
process_device: torch.device,
|
||||||
|
component_name: str,
|
||||||
) -> int:
|
) -> int:
|
||||||
processed_layers = 0
|
processed_layers = 0
|
||||||
for module in model.modules():
|
for module in model.modules():
|
||||||
@@ -144,8 +199,8 @@ def _process_quantized_text_encoder_weights(
|
|||||||
module.to(origin_device)
|
module.to(origin_device)
|
||||||
if processed_layers == 0:
|
if processed_layers == 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"The text-encoder checkpoint declares quantization, but the model "
|
f"The {component_name!r} checkpoint declares quantization, but the "
|
||||||
"did not construct any quantized linear layers"
|
"model did not construct any quantized linear layers"
|
||||||
)
|
)
|
||||||
return processed_layers
|
return processed_layers
|
||||||
|
|
||||||
@@ -412,14 +467,11 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
if post_diffusers_config_update is not None:
|
if post_diffusers_config_update is not None:
|
||||||
post_diffusers_config_update()
|
post_diffusers_config_update()
|
||||||
model_cls, _ = ModelRegistry.resolve_model_cls(
|
model_cls = _resolve_and_configure_encoder_quantization(
|
||||||
getattr(encoder_config, "architectures", [])
|
|
||||||
)
|
|
||||||
_configure_text_encoder_quantization(
|
|
||||||
encoder_config,
|
encoder_config,
|
||||||
model_cls,
|
|
||||||
model_config,
|
model_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
component_name,
|
||||||
)
|
)
|
||||||
encoder_dp_group = get_encoder_data_parallel_group()
|
encoder_dp_group = get_encoder_data_parallel_group()
|
||||||
prefer_dp = (
|
prefer_dp = (
|
||||||
@@ -485,13 +537,14 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
if quant_config is not None:
|
if quant_config is not None:
|
||||||
if param_dtype not in quant_config.get_supported_act_dtypes():
|
if param_dtype not in quant_config.get_supported_act_dtypes():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
f"{component_name!r} quantization method "
|
||||||
|
f"{quant_config.get_name()!r} "
|
||||||
f"does not support activation dtype {param_dtype}"
|
f"does not support activation dtype {param_dtype}"
|
||||||
)
|
)
|
||||||
if current_platform.is_mps():
|
if current_platform.is_mps():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
f"{component_name!r} quantization method "
|
||||||
"is not supported on MPS"
|
f"{quant_config.get_name()!r} is not supported on MPS"
|
||||||
)
|
)
|
||||||
if current_platform.is_cuda():
|
if current_platform.is_cuda():
|
||||||
capability = current_platform.get_device_capability()
|
capability = current_platform.get_device_capability()
|
||||||
@@ -500,7 +553,8 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
and capability.to_int() < quant_config.get_min_capability()
|
and capability.to_int() < quant_config.get_min_capability()
|
||||||
):
|
):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
f"{component_name!r} quantization method "
|
||||||
|
f"{quant_config.get_name()!r} "
|
||||||
"requires CUDA compute capability "
|
"requires CUDA compute capability "
|
||||||
f">= {quant_config.get_min_capability() / 10:.1f}; got "
|
f">= {quant_config.get_min_capability() / 10:.1f}; got "
|
||||||
f"{capability.to_int() / 10:.1f}"
|
f"{capability.to_int() / 10:.1f}"
|
||||||
@@ -575,14 +629,16 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if quant_config is not None:
|
if quant_config is not None:
|
||||||
processed_layers = _process_quantized_text_encoder_weights(
|
processed_layers = _process_quantized_encoder_weights(
|
||||||
model,
|
model,
|
||||||
local_torch_device,
|
local_torch_device,
|
||||||
|
component_name,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Processed %d %s text-encoder linear layers",
|
"Processed %d %s linear layers for %s",
|
||||||
processed_layers,
|
processed_layers,
|
||||||
quant_config.get_name(),
|
quant_config.get_name(),
|
||||||
|
component_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
if component_starts_on_cpu:
|
if component_starts_on_cpu:
|
||||||
|
|||||||
@@ -2,7 +2,8 @@
|
|||||||
|
|
||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import field
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
@@ -154,10 +155,22 @@ def finalize_encoder_folding(
|
|||||||
config.parallel_folding_mode = None
|
config.parallel_folding_mode = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CheckpointQuantizationCapability:
|
||||||
|
"""Quantized-checkpoint contract implemented by a native encoder."""
|
||||||
|
|
||||||
|
backend: Literal["diffusion", "srt"]
|
||||||
|
methods: frozenset[str]
|
||||||
|
|
||||||
|
|
||||||
class EncoderTensorParallelMixin:
|
class EncoderTensorParallelMixin:
|
||||||
"""Keep an encoder on the TP group that was used to build its shards."""
|
"""Keep an encoder on the TP group that was used to build its shards."""
|
||||||
|
|
||||||
_encoder_tp_group: GroupCoordinator | None = None
|
_encoder_tp_group: GroupCoordinator | None = None
|
||||||
|
checkpoint_quantization_capability: CheckpointQuantizationCapability | None = None
|
||||||
|
# Some encoders own checkpoint quantization end to end because their weight
|
||||||
|
# states or sharding contract cannot use the generic loader lifecycle.
|
||||||
|
manages_checkpoint_quantization = False
|
||||||
|
|
||||||
def bind_encoder_tp_group(self, tp_group: GroupCoordinator) -> None:
|
def bind_encoder_tp_group(self, tp_group: GroupCoordinator) -> None:
|
||||||
self._encoder_tp_group = tp_group
|
self._encoder_tp_group = tp_group
|
||||||
|
|||||||
@@ -17,7 +17,10 @@ from sglang.multimodal_gen.configs.models.encoders.minimax_h3_qwen3vl import (
|
|||||||
)
|
)
|
||||||
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.loader.weight_utils import default_weight_loader
|
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
from sglang.multimodal_gen.runtime.models.encoders.base import (
|
||||||
|
CheckpointQuantizationCapability,
|
||||||
|
TextEncoder,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLModel
|
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLModel
|
||||||
|
|
||||||
MINIMAX_H3_QWEN3VL_HIDDEN_DIM = 5120
|
MINIMAX_H3_QWEN3VL_HIDDEN_DIM = 5120
|
||||||
@@ -41,11 +44,15 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
eight otherwise-idle ranks during encoding.
|
eight otherwise-idle ranks during encoding.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
supports_dp_encode = True
|
|
||||||
# The inherited text-layer list covers Qwen's language stack; reference
|
# The inherited text-layer list covers Qwen's language stack; reference
|
||||||
# modes also execute the embedded visual tower.
|
# modes also execute the embedded visual tower.
|
||||||
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
||||||
supported_checkpoint_quantization_methods = frozenset({"fp8"})
|
|
||||||
|
supports_dp_encode = True
|
||||||
|
checkpoint_quantization_capability = CheckpointQuantizationCapability(
|
||||||
|
backend="diffusion",
|
||||||
|
methods=frozenset({"fp8"}),
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def should_materialize_checkpoint_weight(name: str) -> bool:
|
def should_materialize_checkpoint_weight(name: str) -> bool:
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
|
ComponentCheckpointUnsupportedError,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.image_encoder_loader import (
|
||||||
|
ImageEncoderLoader,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestImageEncoderQuantizationAdmission(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.loader = ImageEncoderLoader()
|
||||||
|
load_native_patcher = mock.patch.object(
|
||||||
|
self.loader, "load_native", return_value=object()
|
||||||
|
)
|
||||||
|
self.load_native = load_native_patcher.start()
|
||||||
|
self.addCleanup(load_native_patcher.stop)
|
||||||
|
self.server_args = SimpleNamespace(
|
||||||
|
pipeline_config=SimpleNamespace(
|
||||||
|
image_encoder_config=CLIPVisionConfig(),
|
||||||
|
image_encoder_precision="bf16",
|
||||||
|
native_only_components=(),
|
||||||
|
),
|
||||||
|
encoder_parallel="replicate",
|
||||||
|
resolve_component_attention_backend=lambda _name: (None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _component_config(self, architecture, *, quantized):
|
||||||
|
config = {"architectures": [architecture]}
|
||||||
|
if quantized:
|
||||||
|
config["quantization_config"] = {
|
||||||
|
"quant_method": "fp8",
|
||||||
|
"activation_scheme": "dynamic",
|
||||||
|
}
|
||||||
|
return config
|
||||||
|
|
||||||
|
def _config_patch(self, config):
|
||||||
|
return mock.patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"image_encoder_loader.get_diffusers_component_config",
|
||||||
|
return_value=config,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _load(self):
|
||||||
|
return self.loader.load(
|
||||||
|
"/model/image_encoder", self.server_args, "image_encoder", "transformers"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_quantized_clip_checkpoint_is_not_silently_enabled(self):
|
||||||
|
config = self._component_config("CLIPVisionModelWithProjection", quantized=True)
|
||||||
|
with self._config_patch(config), self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError,
|
||||||
|
"CLIPVisionModel.*image_encoder.*no checkpoint quantization capability",
|
||||||
|
):
|
||||||
|
self.loader.load_customized("/model/image_encoder", self.server_args)
|
||||||
|
|
||||||
|
def test_unknown_quantized_architecture_does_not_fall_back(self):
|
||||||
|
config = self._component_config("UnknownVisionModel", quantized=True)
|
||||||
|
with self._config_patch(config), self.assertRaises(
|
||||||
|
ComponentCheckpointUnsupportedError
|
||||||
|
):
|
||||||
|
self._load()
|
||||||
|
self.load_native.assert_not_called()
|
||||||
|
|
||||||
|
def test_unknown_unquantized_architecture_keeps_native_fallback(self):
|
||||||
|
config = self._component_config("UnknownVisionModel", quantized=False)
|
||||||
|
with self._config_patch(config):
|
||||||
|
self._load()
|
||||||
|
self.load_native.assert_called_once()
|
||||||
@@ -8,12 +8,18 @@ from torch import nn
|
|||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import LinearBase
|
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 (
|
||||||
|
ComponentCheckpointUnsupportedError,
|
||||||
|
)
|
||||||
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_text_encoder_quantization,
|
_configure_encoder_quantization,
|
||||||
_process_quantized_text_encoder_weights,
|
_process_quantized_encoder_weights,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.models.encoders.base import (
|
||||||
|
CheckpointQuantizationCapability,
|
||||||
|
TextEncoder,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import (
|
from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import (
|
||||||
MiniMaxH3Qwen3VLEncoder,
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
)
|
)
|
||||||
@@ -147,22 +153,48 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
|
|
||||||
def test_serialized_fp8_checkpoint_configures_h3_encoder(self):
|
def test_serialized_fp8_checkpoint_configures_h3_encoder(self):
|
||||||
model_config = SimpleNamespace(quant_config=None)
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
_configure_text_encoder_quantization(
|
_configure_encoder_quantization(
|
||||||
model_config,
|
model_config,
|
||||||
MiniMaxH3Qwen3VLEncoder,
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
{},
|
{},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
self.assertIs(model_config.quant_config, self.serialized)
|
self.assertIs(model_config.quant_config, self.serialized)
|
||||||
|
|
||||||
def test_encoder_class_must_opt_in(self):
|
def test_encoder_class_must_opt_in(self):
|
||||||
model_config = SimpleNamespace(quant_config=None)
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
with self.assertRaisesRegex(ValueError, "does not support"):
|
with self.assertRaisesRegex(
|
||||||
_configure_text_encoder_quantization(
|
ComponentCheckpointUnsupportedError, "does not support"
|
||||||
|
):
|
||||||
|
_configure_encoder_quantization(
|
||||||
model_config,
|
model_config,
|
||||||
TextEncoder,
|
TextEncoder,
|
||||||
{},
|
{},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"text_encoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_srt_backend_is_not_admitted_without_an_adapter(self):
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
capability = CheckpointQuantizationCapability(
|
||||||
|
backend="srt",
|
||||||
|
methods=frozenset({"fp8"}),
|
||||||
|
)
|
||||||
|
with mock.patch.object(
|
||||||
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
|
"checkpoint_quantization_capability",
|
||||||
|
capability,
|
||||||
|
), self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError,
|
||||||
|
"'srt'.*only the 'diffusion' backend",
|
||||||
|
):
|
||||||
|
_configure_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_model_managed_quantization_bypasses_generic_lifecycle(self):
|
def test_model_managed_quantization_bypasses_generic_lifecycle(self):
|
||||||
@@ -172,11 +204,12 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
"manages_checkpoint_quantization",
|
"manages_checkpoint_quantization",
|
||||||
True,
|
True,
|
||||||
):
|
):
|
||||||
_configure_text_encoder_quantization(
|
_configure_encoder_quantization(
|
||||||
model_config,
|
model_config,
|
||||||
TextEncoder,
|
TextEncoder,
|
||||||
{},
|
{},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertIsNone(model_config.quant_config)
|
self.assertIsNone(model_config.quant_config)
|
||||||
@@ -213,9 +246,10 @@ class TestQuantizedTextEncoderPostprocess(unittest.TestCase):
|
|||||||
quant_method = _RecordingQuantMethod()
|
quant_method = _RecordingQuantMethod()
|
||||||
model = _QuantizedEncoder(quant_method)
|
model = _QuantizedEncoder(quant_method)
|
||||||
|
|
||||||
processed = _process_quantized_text_encoder_weights(
|
processed = _process_quantized_encoder_weights(
|
||||||
model,
|
model,
|
||||||
torch.device("cpu"),
|
torch.device("cpu"),
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(processed, 1)
|
self.assertEqual(processed, 1)
|
||||||
@@ -227,9 +261,10 @@ class TestQuantizedTextEncoderPostprocess(unittest.TestCase):
|
|||||||
quant_method = _RecordingQuantMethod()
|
quant_method = _RecordingQuantMethod()
|
||||||
model = _QuantizedEncoder(quant_method)
|
model = _QuantizedEncoder(quant_method)
|
||||||
|
|
||||||
processed = _process_quantized_text_encoder_weights(
|
processed = _process_quantized_encoder_weights(
|
||||||
model,
|
model,
|
||||||
torch.device("cuda", torch.cuda.current_device()),
|
torch.device("cuda", torch.cuda.current_device()),
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(processed, 1)
|
self.assertEqual(processed, 1)
|
||||||
@@ -242,9 +277,10 @@ class TestQuantizedTextEncoderPostprocess(unittest.TestCase):
|
|||||||
model = _QuantizedEncoder(_RecordingQuantMethod(error=RuntimeError("boom")))
|
model = _QuantizedEncoder(_RecordingQuantMethod(error=RuntimeError("boom")))
|
||||||
|
|
||||||
with self.assertRaisesRegex(RuntimeError, "boom"):
|
with self.assertRaisesRegex(RuntimeError, "boom"):
|
||||||
_process_quantized_text_encoder_weights(
|
_process_quantized_encoder_weights(
|
||||||
model,
|
model,
|
||||||
torch.device("cuda", torch.cuda.current_device()),
|
torch.device("cuda", torch.cuda.current_device()),
|
||||||
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(model.quantized.weight.device, torch.device("cpu"))
|
self.assertEqual(model.quantized.weight.device, torch.device("cpu"))
|
||||||
|
|||||||
Reference in New Issue
Block a user