From d55f1c28e223226575d40d4449266630eaf2a277 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 18 Aug 2026 09:10:54 +0800 Subject: [PATCH] [diffusion] feat: load quantized H3 text encoder checkpoints (#34986) Co-authored-by: Yiqi Yang --- .../cookbook/diffusion/MiniMax/MiniMax-H3.mdx | 23 +++ docs/docs/sglang-diffusion/api/cli.mdx | 6 + docs/docs/sglang-diffusion/quantization.mdx | 30 +++- .../configs/pipeline_configs/minimax_h3.py | 8 + .../component_loaders/text_encoder_loader.py | 137 ++++++++++++++++++ .../runtime/models/encoders/base.py | 6 + .../runtime/models/encoders/ideogram.py | 1 + .../models/encoders/minimax_h3_qwen3vl.py | 7 +- .../runtime/models/encoders/qwen3vl.py | 9 +- .../test/unit/test_text_encoder_loader.py | 129 +++++++++++++++++ 10 files changed, 351 insertions(+), 5 deletions(-) diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index 8097a3313..d9c432aa7 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -738,6 +738,29 @@ The picker exposes this option only on the B200 and B300 topologies used for real H3 validation runs. +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. + diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index 103a67496..78e46b81e 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -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-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}` compresses completed KV-cache chunks independently of transformer weight quantization. It is lossy and disabled by default. diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 20339f6b9..94449742c 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -2,11 +2,12 @@ title: "Quantization" tag: "approx" 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 -the base model and the quantized transformer override separate. +SGLang-Diffusion supports quantized transformer checkpoints and selected +quantized native text-encoder checkpoints. Transformer and text-encoder +precision are resolved independently. ## 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 - `--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_`) +- `--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 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. +## 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 This section is the canonical support matrix for the thirteen published diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py index afb29a65d..565f1872f 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py @@ -110,6 +110,12 @@ class MiniMaxH3PipelineConfig(PipelineConfig): else type(current_platform).__name__ ) 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 = { "attention_backend": attention_backend, "backend": self._server_arg_value(server_args.backend), @@ -123,6 +129,7 @@ class MiniMaxH3PipelineConfig(PipelineConfig): "num_gpus": server_args.num_gpus, "performance_mode": server_args.performance_mode, "quantization": server_args.quantization, + "text_encoder_quantization": text_encoder_quantization, "regional_compile": server_args.regional_compile, "ring_degree": server_args.ring_degree, "sp_degree": server_args.sp_degree, @@ -144,6 +151,7 @@ class MiniMaxH3PipelineConfig(PipelineConfig): "num_gpus": 4, "performance_mode": "speed", "quantization": None, + "text_encoder_quantization": None, "regional_compile": False, "ring_degree": 1, "sp_degree": 4, diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 40f4f3609..d79ce8cd2 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -3,6 +3,7 @@ import glob import os import re from collections.abc import Callable, Generator, Iterable +from itertools import chain from typing import cast import torch @@ -19,6 +20,10 @@ from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed.parallel_state import ( 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 ( 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.precision import precision_to_dtype +from sglang.multimodal_gen.runtime.utils.quantization_utils import get_quant_config from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.srt.environ import envs logger = init_logger(__name__) +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): """Loader for text encoders.""" @@ -319,6 +414,12 @@ class TextEncoderLoader(ComponentLoader): model_cls, _ = ModelRegistry.resolve_model_cls( 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 finalize_encoder_folding( encoder_config, @@ -376,6 +477,31 @@ class TextEncoderLoader(ComponentLoader): component_name: str = "text_encoder", ): 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(): 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 current_platform.is_mps(): model = model.to(local_torch_device) diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/base.py b/python/sglang/multimodal_gen/runtime/models/encoders/base.py index 481ff8b77..0d8c34f73 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/base.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/base.py @@ -174,6 +174,12 @@ class TextEncoder( # Qwen2_5_VLCausalLMOutputWithPast). Off by default so a new encoder is # replicated rather than silently broken; flip it once dp is verified there. 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 layer_names = [ "layers", diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py b/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py index 2ec421af4..95b3a9f94 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/ideogram.py @@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLTextMod class IdeogramQwen3VLTextEncoder(TextEncoder): """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) layer_names = ["language_model.layers"] diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py index 9b08546ac..edb51f98d 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/minimax_h3_qwen3vl.py @@ -44,6 +44,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder): layer_names = [*TextEncoder.layer_names, "model.visual.blocks"] supports_dp_encode = True + supported_checkpoint_quantization_methods = frozenset({"fp8"}) @staticmethod def should_materialize_checkpoint_weight(name: str) -> bool: @@ -61,7 +62,11 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder): "MiniMax H3 Qwen3-VL config must be trimmed to " 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. self.model.language_model.norm = nn.Identity() self.image_token_id = int(arch.image_token_id) diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py index daff67a52..720ad375c 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py @@ -662,11 +662,18 @@ class Qwen3VLModel(nn.Module): config: Qwen3VLConfig _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__() self.visual = Qwen3VLVisionTransformer(config.vision_config) self.language_model = Qwen3VLTextModel( config.text_config, + quant_config=quant_config, use_tensor_parallel=use_tensor_parallel, ) self.rope_deltas = None # cache rope_deltas here diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py index caae1c033..ac8871b78 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py @@ -2,11 +2,18 @@ import unittest from types import SimpleNamespace from unittest import mock +import torch 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 ( 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 ( 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__": unittest.main()