[diffusion] feat: load quantized H3 text encoder checkpoints (#34986)
Co-authored-by: Yiqi Yang <yangyiqi8787@gmail.com>
This commit is contained in:
@@ -738,6 +738,29 @@ The picker exposes this option only on the B200 and B300 topologies used for
|
|||||||
real H3 validation runs.
|
real H3 validation runs.
|
||||||
</Warning>
|
</Warning>
|
||||||
|
|
||||||
|
The Qwen3-VL text encoder can be replaced independently of the DiT. To reduce
|
||||||
|
its resident memory, point the text-encoder component at the serialized FP8
|
||||||
|
checkpoint used in validation:
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
sglang serve \
|
||||||
|
--model-path MiniMaxAI/MiniMax-H3 \
|
||||||
|
--model-variant fl2va \
|
||||||
|
--component-paths.text_encoder Qwen/Qwen3-VL-32B-Instruct-FP8 \
|
||||||
|
--num-gpus 4 \
|
||||||
|
--tp-size 2 \
|
||||||
|
--ulysses-degree 2 \
|
||||||
|
--performance-mode speed \
|
||||||
|
--port 30010
|
||||||
|
```
|
||||||
|
|
||||||
|
`--text-encoder-path` is accepted as a shorter alias. No separate quantization
|
||||||
|
flag is required: SGLang reads the checkpoint's `quantization_config` and
|
||||||
|
fails closed if the native encoder does not support that format. The language
|
||||||
|
linear layers use FP8 while embeddings, normalization, and the vision tower
|
||||||
|
remain BF16. This is an approximate serve-time choice and is incompatible with
|
||||||
|
the strict `quality="high"` deployment contract.
|
||||||
|
|
||||||
</Tab>
|
</Tab>
|
||||||
|
|
||||||
</Tabs>
|
</Tabs>
|
||||||
|
|||||||
@@ -129,6 +129,12 @@ For quantized transformer checkpoints, prefer:
|
|||||||
- `--quantization` for online quantization (apply quantization to unquantized models at load time, activations are quantized dynamically)
|
- `--quantization` for online quantization (apply quantization to unquantized models at load time, activations are quantized dynamically)
|
||||||
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
||||||
|
|
||||||
|
Component checkpoint paths are selected separately, so changing DiT precision
|
||||||
|
never silently changes prompt embeddings. For a native text encoder:
|
||||||
|
|
||||||
|
- `--component-paths.text_encoder {MODEL}` replaces the text-encoder checkpoint; `--text-encoder-path {MODEL}` is its shorter alias
|
||||||
|
- Quantization metadata is auto-detected from that checkpoint. Each native encoder must explicitly support the serialized format; this is not blanket quantization support for every component, and unsupported combinations fail before weight loading.
|
||||||
|
|
||||||
For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}`
|
For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}`
|
||||||
compresses completed KV-cache chunks independently of transformer weight
|
compresses completed KV-cache chunks independently of transformer weight
|
||||||
quantization. It is lossy and disabled by default.
|
quantization. It is lossy and disabled by default.
|
||||||
|
|||||||
@@ -2,11 +2,12 @@
|
|||||||
title: "Quantization"
|
title: "Quantization"
|
||||||
tag: "approx"
|
tag: "approx"
|
||||||
metatags:
|
metatags:
|
||||||
description: "Configure transformer weight quantization and Quant-VideoGen causal KV-cache quantization in SGLang-Diffusion."
|
description: "Configure transformer and native text-encoder checkpoint quantization, plus Quant-VideoGen causal KV-cache quantization in SGLang-Diffusion."
|
||||||
---
|
---
|
||||||
|
|
||||||
SGLang-Diffusion supports quantized transformer checkpoints. In most cases, keep
|
SGLang-Diffusion supports quantized transformer checkpoints and selected
|
||||||
the base model and the quantized transformer override separate.
|
quantized native text-encoder checkpoints. Transformer and text-encoder
|
||||||
|
precision are resolved independently.
|
||||||
|
|
||||||
## Quick Reference
|
## Quick Reference
|
||||||
|
|
||||||
@@ -17,6 +18,8 @@ Use these paths:
|
|||||||
- `--transformer-weights-path`: quantized transformer weights provided as a single safetensors file, a sharded safetensors directory, a local path, or a Hugging Face repo ID
|
- `--transformer-weights-path`: quantized transformer weights provided as a single safetensors file, a sharded safetensors directory, a local path, or a Hugging Face repo ID
|
||||||
- `--quantization`: apply online quantization to unquantized models at load time (activations are quantized dynamically)
|
- `--quantization`: apply online quantization to unquantized models at load time (activations are quantized dynamically)
|
||||||
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
||||||
|
- `--component-paths.text_encoder`: replace a native text encoder with a checkpoint whose `quantization_config` is auto-detected
|
||||||
|
- `--text-encoder-path`: shorter alias for `--component-paths.text_encoder`
|
||||||
- `--kv-cache-quant`: compress completed causal KV-cache chunks for supported realtime models
|
- `--kv-cache-quant`: compress completed causal KV-cache chunks for supported realtime models
|
||||||
|
|
||||||
Recommended example for pre-quantized checkpoints:
|
Recommended example for pre-quantized checkpoints:
|
||||||
@@ -280,6 +283,27 @@ sglang generate \
|
|||||||
|
|
||||||
Each pattern is matched against the full layer prefix (e.g. `layers.0.attention.to_q`). A layer is skipped and left unquantizd if its prefix contains any of the given patterns.
|
Each pattern is matched against the full layer prefix (e.g. `layers.0.attention.to_q`). A layer is skipped and left unquantizd if its prefix contains any of the given patterns.
|
||||||
|
|
||||||
|
## MiniMax-H3 Text Encoder FP8
|
||||||
|
|
||||||
|
MiniMax-H3 can load a serialized FP8 checkpoint for the language linear layers
|
||||||
|
in its native Qwen3-VL text encoder independently of the DiT. Embeddings,
|
||||||
|
normalization layers, and the Qwen vision tower remain in BF16.
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
sglang serve \
|
||||||
|
--model-path MiniMaxAI/MiniMax-H3 \
|
||||||
|
--model-variant fl2va \
|
||||||
|
--component-paths.text_encoder Qwen/Qwen3-VL-32B-Instruct-FP8 \
|
||||||
|
--num-gpus 4 \
|
||||||
|
--port 30010
|
||||||
|
```
|
||||||
|
|
||||||
|
`--text-encoder-path` is accepted as a shorter alias. No quantization flag is
|
||||||
|
needed: SGLang detects the checkpoint metadata and only enables formats that
|
||||||
|
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.
|
||||||
|
|
||||||
## 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
|
||||||
|
|||||||
@@ -110,6 +110,12 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
|||||||
else type(current_platform).__name__
|
else type(current_platform).__name__
|
||||||
)
|
)
|
||||||
model_variant = str(server_args.model_variant or "fl2va").lower()
|
model_variant = str(server_args.model_variant or "fl2va").lower()
|
||||||
|
resolved_quant_config = self.text_encoder_configs[0].quant_config
|
||||||
|
text_encoder_quantization = (
|
||||||
|
resolved_quant_config.get_name()
|
||||||
|
if resolved_quant_config is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
actual = {
|
actual = {
|
||||||
"attention_backend": attention_backend,
|
"attention_backend": attention_backend,
|
||||||
"backend": self._server_arg_value(server_args.backend),
|
"backend": self._server_arg_value(server_args.backend),
|
||||||
@@ -123,6 +129,7 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
|||||||
"num_gpus": server_args.num_gpus,
|
"num_gpus": server_args.num_gpus,
|
||||||
"performance_mode": server_args.performance_mode,
|
"performance_mode": server_args.performance_mode,
|
||||||
"quantization": server_args.quantization,
|
"quantization": server_args.quantization,
|
||||||
|
"text_encoder_quantization": text_encoder_quantization,
|
||||||
"regional_compile": server_args.regional_compile,
|
"regional_compile": server_args.regional_compile,
|
||||||
"ring_degree": server_args.ring_degree,
|
"ring_degree": server_args.ring_degree,
|
||||||
"sp_degree": server_args.sp_degree,
|
"sp_degree": server_args.sp_degree,
|
||||||
@@ -144,6 +151,7 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
|||||||
"num_gpus": 4,
|
"num_gpus": 4,
|
||||||
"performance_mode": "speed",
|
"performance_mode": "speed",
|
||||||
"quantization": None,
|
"quantization": None,
|
||||||
|
"text_encoder_quantization": None,
|
||||||
"regional_compile": False,
|
"regional_compile": False,
|
||||||
"ring_degree": 1,
|
"ring_degree": 1,
|
||||||
"sp_degree": 4,
|
"sp_degree": 4,
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import glob
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from collections.abc import Callable, Generator, Iterable
|
from collections.abc import Callable, Generator, Iterable
|
||||||
|
from itertools import chain
|
||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -19,6 +20,10 @@ from sglang.multimodal_gen.runtime.distributed import (
|
|||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
use_tensor_parallel_group,
|
use_tensor_parallel_group,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
LinearBase,
|
||||||
|
UnquantizedLinearMethod,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentLoader,
|
ComponentLoader,
|
||||||
)
|
)
|
||||||
@@ -48,12 +53,102 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
)
|
)
|
||||||
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.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.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__)
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_text_encoder_quantization(
|
||||||
|
model_config: EncoderConfig,
|
||||||
|
model_cls: type[nn.Module],
|
||||||
|
component_config: dict,
|
||||||
|
component_model_path: str,
|
||||||
|
) -> None:
|
||||||
|
if getattr(model_cls, "manages_checkpoint_quantization", False):
|
||||||
|
# Preserve model-owned formats such as Ideogram's bitsandbytes state.
|
||||||
|
# Those models parse metadata, construct layers, and attach quant states
|
||||||
|
# themselves; running the generic lifecycle as well would process twice.
|
||||||
|
return
|
||||||
|
|
||||||
|
quant_config = get_quant_config(
|
||||||
|
component_config,
|
||||||
|
component_model_path,
|
||||||
|
)
|
||||||
|
model_config.quant_config = quant_config
|
||||||
|
if quant_config is None:
|
||||||
|
return
|
||||||
|
if not issubclass(model_cls, TextEncoder):
|
||||||
|
raise ValueError(
|
||||||
|
"A quantized text-encoder checkpoint requires an in-tree native "
|
||||||
|
"TextEncoder; "
|
||||||
|
f"got {model_cls.__name__}"
|
||||||
|
)
|
||||||
|
quant_method = quant_config.get_name()
|
||||||
|
supported_methods = model_cls.supported_checkpoint_quantization_methods
|
||||||
|
if quant_method not in supported_methods:
|
||||||
|
raise ValueError(
|
||||||
|
f"{model_cls.__name__} does not support text-encoder checkpoints "
|
||||||
|
f"quantized with {quant_method!r}; supported methods: "
|
||||||
|
f"{sorted(supported_methods)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
||||||
|
"""Return the device of a module's own tensors.
|
||||||
|
|
||||||
|
Quantized linear layers are expected to keep their parameters and buffers
|
||||||
|
together. Failing explicitly is safer than staging only part of a layer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
devices = {
|
||||||
|
tensor.device
|
||||||
|
for tensor in chain(
|
||||||
|
module.parameters(recurse=False),
|
||||||
|
module.buffers(recurse=False),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if len(devices) > 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"Cannot stage {type(module).__name__} with tensors on multiple "
|
||||||
|
f"devices: {sorted(map(str, devices))}"
|
||||||
|
)
|
||||||
|
return next(iter(devices), None)
|
||||||
|
|
||||||
|
|
||||||
|
def _process_quantized_text_encoder_weights(
|
||||||
|
model: nn.Module,
|
||||||
|
process_device: torch.device,
|
||||||
|
) -> int:
|
||||||
|
processed_layers = 0
|
||||||
|
for module in model.modules():
|
||||||
|
if not isinstance(module, LinearBase):
|
||||||
|
continue
|
||||||
|
quant_method = module.quant_method
|
||||||
|
if quant_method is None or isinstance(quant_method, UnquantizedLinearMethod):
|
||||||
|
continue
|
||||||
|
|
||||||
|
origin_device = _module_tensor_device(module)
|
||||||
|
should_stage = origin_device is not None and origin_device != process_device
|
||||||
|
if should_stage:
|
||||||
|
module.to(process_device)
|
||||||
|
try:
|
||||||
|
quant_method.process_weights_after_loading(module)
|
||||||
|
processed_layers += 1
|
||||||
|
finally:
|
||||||
|
# Post-load methods may replace parameters or register buffers. Move
|
||||||
|
# the complete layer back so component residency remains authoritative.
|
||||||
|
if should_stage:
|
||||||
|
module.to(origin_device)
|
||||||
|
if processed_layers == 0:
|
||||||
|
raise ValueError(
|
||||||
|
"The text-encoder checkpoint declares quantization, but the model "
|
||||||
|
"did not construct any quantized linear layers"
|
||||||
|
)
|
||||||
|
return processed_layers
|
||||||
|
|
||||||
|
|
||||||
class TextEncoderLoader(ComponentLoader):
|
class TextEncoderLoader(ComponentLoader):
|
||||||
"""Loader for text encoders."""
|
"""Loader for text encoders."""
|
||||||
|
|
||||||
@@ -319,6 +414,12 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
model_cls, _ = ModelRegistry.resolve_model_cls(
|
model_cls, _ = ModelRegistry.resolve_model_cls(
|
||||||
getattr(encoder_config, "architectures", [])
|
getattr(encoder_config, "architectures", [])
|
||||||
)
|
)
|
||||||
|
_configure_text_encoder_quantization(
|
||||||
|
encoder_config,
|
||||||
|
model_cls,
|
||||||
|
model_config,
|
||||||
|
component_model_path,
|
||||||
|
)
|
||||||
# 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,
|
||||||
@@ -376,6 +477,31 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
component_name: str = "text_encoder",
|
component_name: str = "text_encoder",
|
||||||
):
|
):
|
||||||
local_torch_device = get_local_torch_device()
|
local_torch_device = get_local_torch_device()
|
||||||
|
quant_config = model_config.quant_config
|
||||||
|
param_dtype = PRECISION_TO_TYPE[dtype]
|
||||||
|
if quant_config is not None:
|
||||||
|
if param_dtype not in quant_config.get_supported_act_dtypes():
|
||||||
|
raise ValueError(
|
||||||
|
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
||||||
|
f"does not support activation dtype {param_dtype}"
|
||||||
|
)
|
||||||
|
if current_platform.is_mps():
|
||||||
|
raise ValueError(
|
||||||
|
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
||||||
|
"is not supported on MPS"
|
||||||
|
)
|
||||||
|
if current_platform.is_cuda():
|
||||||
|
capability = current_platform.get_device_capability()
|
||||||
|
if (
|
||||||
|
capability is not None
|
||||||
|
and capability.to_int() < quant_config.get_min_capability()
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"Text-encoder quantization method {quant_config.get_name()!r} "
|
||||||
|
"requires CUDA compute capability "
|
||||||
|
f">= {quant_config.get_min_capability() / 10:.1f}; got "
|
||||||
|
f"{capability.to_int() / 10:.1f}"
|
||||||
|
)
|
||||||
|
|
||||||
if not current_platform.is_cpu():
|
if not current_platform.is_cpu():
|
||||||
component_starts_on_cpu = (
|
component_starts_on_cpu = (
|
||||||
@@ -439,6 +565,17 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if quant_config is not None:
|
||||||
|
processed_layers = _process_quantized_text_encoder_weights(
|
||||||
|
model,
|
||||||
|
local_torch_device,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Processed %d %s text-encoder linear layers",
|
||||||
|
processed_layers,
|
||||||
|
quant_config.get_name(),
|
||||||
|
)
|
||||||
|
|
||||||
if component_starts_on_cpu:
|
if component_starts_on_cpu:
|
||||||
if current_platform.is_mps():
|
if current_platform.is_mps():
|
||||||
model = model.to(local_torch_device)
|
model = model.to(local_torch_device)
|
||||||
|
|||||||
@@ -174,6 +174,12 @@ class TextEncoder(
|
|||||||
# Qwen2_5_VLCausalLMOutputWithPast). Off by default so a new encoder is
|
# Qwen2_5_VLCausalLMOutputWithPast). Off by default so a new encoder is
|
||||||
# replicated rather than silently broken; flip it once dp is verified there.
|
# replicated rather than silently broken; flip it once dp is verified there.
|
||||||
supports_dp_encode = False
|
supports_dp_encode = False
|
||||||
|
# Quantized checkpoints are opt-in because an encoder must construct
|
||||||
|
# quantized linears and load the checkpoint's auxiliary scale parameters.
|
||||||
|
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
|
||||||
|
# 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
|
||||||
layerwise_offload_dit_group_enabled = False
|
layerwise_offload_dit_group_enabled = False
|
||||||
layer_names = [
|
layer_names = [
|
||||||
"layers",
|
"layers",
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLTextMod
|
|||||||
class IdeogramQwen3VLTextEncoder(TextEncoder):
|
class IdeogramQwen3VLTextEncoder(TextEncoder):
|
||||||
"""Language-only Qwen3-VL text encoder stored inside Ideogram checkpoints."""
|
"""Language-only Qwen3-VL text encoder stored inside Ideogram checkpoints."""
|
||||||
|
|
||||||
|
manages_checkpoint_quantization = True
|
||||||
_activation_layers = (0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 35)
|
_activation_layers = (0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 35)
|
||||||
layer_names = ["language_model.layers"]
|
layer_names = ["language_model.layers"]
|
||||||
|
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
||||||
|
|
||||||
supports_dp_encode = True
|
supports_dp_encode = True
|
||||||
|
supported_checkpoint_quantization_methods = frozenset({"fp8"})
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def should_materialize_checkpoint_weight(name: str) -> bool:
|
def should_materialize_checkpoint_weight(name: str) -> bool:
|
||||||
@@ -61,7 +62,11 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
"MiniMax H3 Qwen3-VL config must be trimmed to "
|
"MiniMax H3 Qwen3-VL config must be trimmed to "
|
||||||
f"{selected_layer} language layers before construction"
|
f"{selected_layer} language layers before construction"
|
||||||
)
|
)
|
||||||
self.model = Qwen3VLModel(arch, use_tensor_parallel=True)
|
self.model = Qwen3VLModel(
|
||||||
|
arch,
|
||||||
|
quant_config=config.quant_config,
|
||||||
|
use_tensor_parallel=True,
|
||||||
|
)
|
||||||
# H3 consumes the unnormalized output immediately after layer 49.
|
# H3 consumes the unnormalized output immediately after layer 49.
|
||||||
self.model.language_model.norm = nn.Identity()
|
self.model.language_model.norm = nn.Identity()
|
||||||
self.image_token_id = int(arch.image_token_id)
|
self.image_token_id = int(arch.image_token_id)
|
||||||
|
|||||||
@@ -662,11 +662,18 @@ class Qwen3VLModel(nn.Module):
|
|||||||
config: Qwen3VLConfig
|
config: Qwen3VLConfig
|
||||||
_no_split_modules = ["Qwen3VLTextDecoderLayer", "Qwen3VLVisionBlock"]
|
_no_split_modules = ["Qwen3VLTextDecoderLayer", "Qwen3VLVisionBlock"]
|
||||||
|
|
||||||
def __init__(self, config, *, use_tensor_parallel: bool = False):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config,
|
||||||
|
*,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
use_tensor_parallel: bool = False,
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.visual = Qwen3VLVisionTransformer(config.vision_config)
|
self.visual = Qwen3VLVisionTransformer(config.vision_config)
|
||||||
self.language_model = Qwen3VLTextModel(
|
self.language_model = Qwen3VLTextModel(
|
||||||
config.text_config,
|
config.text_config,
|
||||||
|
quant_config=quant_config,
|
||||||
use_tensor_parallel=use_tensor_parallel,
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
)
|
)
|
||||||
self.rope_deltas = None # cache rope_deltas here
|
self.rope_deltas = None # cache rope_deltas here
|
||||||
|
|||||||
@@ -2,11 +2,18 @@ import unittest
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
|
import torch
|
||||||
import transformers
|
import transformers
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
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.text_encoder_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
|
||||||
TextEncoderLoader,
|
TextEncoderLoader,
|
||||||
|
_configure_text_encoder_quantization,
|
||||||
|
_process_quantized_text_encoder_weights,
|
||||||
)
|
)
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
@@ -102,5 +109,127 @@ class TestMiniMaxH3CheckpointFilter(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTextEncoderQuantization(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
serialized = Fp8Config(
|
||||||
|
is_checkpoint_fp8_serialized=True,
|
||||||
|
activation_scheme="dynamic",
|
||||||
|
weight_block_size=[128, 128],
|
||||||
|
)
|
||||||
|
self.quant_config_patcher = mock.patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"text_encoder_loader.get_quant_config",
|
||||||
|
return_value=serialized,
|
||||||
|
)
|
||||||
|
self.get_quant_config = self.quant_config_patcher.start()
|
||||||
|
self.addCleanup(self.quant_config_patcher.stop)
|
||||||
|
self.serialized = serialized
|
||||||
|
|
||||||
|
def test_serialized_fp8_checkpoint_configures_h3_encoder(self):
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
_configure_text_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
MiniMaxH3Qwen3VLEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
)
|
||||||
|
self.assertIs(model_config.quant_config, self.serialized)
|
||||||
|
|
||||||
|
def test_encoder_class_must_opt_in(self):
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
with self.assertRaisesRegex(ValueError, "does not support"):
|
||||||
|
_configure_text_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
TextEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_model_managed_quantization_bypasses_generic_lifecycle(self):
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
with mock.patch.object(
|
||||||
|
TextEncoder,
|
||||||
|
"manages_checkpoint_quantization",
|
||||||
|
True,
|
||||||
|
):
|
||||||
|
_configure_text_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
TextEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(model_config.quant_config)
|
||||||
|
self.get_quant_config.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
class _RecordingQuantMethod:
|
||||||
|
def __init__(self, *, error: Exception | None = None):
|
||||||
|
self.error = error
|
||||||
|
self.devices = []
|
||||||
|
|
||||||
|
def process_weights_after_loading(self, layer):
|
||||||
|
self.devices.append(layer.weight.device)
|
||||||
|
if self.error is not None:
|
||||||
|
raise self.error
|
||||||
|
|
||||||
|
|
||||||
|
class _QuantizedLinear(LinearBase):
|
||||||
|
def __init__(self, quant_method):
|
||||||
|
nn.Module.__init__(self)
|
||||||
|
self.weight = nn.Parameter(torch.empty(2, 2), requires_grad=False)
|
||||||
|
self.quant_method = quant_method
|
||||||
|
|
||||||
|
|
||||||
|
class _QuantizedEncoder(nn.Module):
|
||||||
|
def __init__(self, quant_method):
|
||||||
|
super().__init__()
|
||||||
|
self.quantized = _QuantizedLinear(quant_method)
|
||||||
|
self.unquantized = nn.Linear(2, 2, bias=False)
|
||||||
|
|
||||||
|
|
||||||
|
class TestQuantizedTextEncoderPostprocess(unittest.TestCase):
|
||||||
|
def test_processes_quantized_layers_without_moving_the_model(self):
|
||||||
|
quant_method = _RecordingQuantMethod()
|
||||||
|
model = _QuantizedEncoder(quant_method)
|
||||||
|
|
||||||
|
processed = _process_quantized_text_encoder_weights(
|
||||||
|
model,
|
||||||
|
torch.device("cpu"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(processed, 1)
|
||||||
|
self.assertEqual(quant_method.devices, [torch.device("cpu")])
|
||||||
|
self.assertEqual(model.unquantized.weight.device, torch.device("cpu"))
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
def test_stages_only_the_quantized_layer_and_restores_it(self):
|
||||||
|
quant_method = _RecordingQuantMethod()
|
||||||
|
model = _QuantizedEncoder(quant_method)
|
||||||
|
|
||||||
|
processed = _process_quantized_text_encoder_weights(
|
||||||
|
model,
|
||||||
|
torch.device("cuda", torch.cuda.current_device()),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(processed, 1)
|
||||||
|
self.assertEqual(quant_method.devices[0].type, "cuda")
|
||||||
|
self.assertEqual(model.quantized.weight.device, torch.device("cpu"))
|
||||||
|
self.assertEqual(model.unquantized.weight.device, torch.device("cpu"))
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
def test_restores_staged_layer_when_postprocess_fails(self):
|
||||||
|
model = _QuantizedEncoder(_RecordingQuantMethod(error=RuntimeError("boom")))
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "boom"):
|
||||||
|
_process_quantized_text_encoder_weights(
|
||||||
|
model,
|
||||||
|
torch.device("cuda", torch.cuda.current_device()),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(model.quantized.weight.device, torch.device("cpu"))
|
||||||
|
self.assertEqual(model.unquantized.weight.device, torch.device("cpu"))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user