From 8df3b9eff90a77939d313adf256d194d837c9fe2 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 24 Aug 2026 13:31:27 +0800 Subject: [PATCH] [diffusion] feat: support loading mixed w4a8 text encoders (#36037) --- .../cookbook/diffusion/MiniMax/MiniMax-H3.mdx | 11 ++ docs/docs/sglang-diffusion/quantization.mdx | 6 +- .../quantization/configs/base_config.py | 4 + .../configs/kitchen_w4a8_config.py | 38 ++++++- .../layers/quantization/kitchen_w4a8.py | 102 ++++++++++++++---- .../component_loaders/text_encoder_loader.py | 9 +- .../runtime/models/encoders/qwen3vl.py | 21 +++- .../runtime/utils/quantization_utils.py | 10 ++ .../test/unit/test_text_encoder_loader.py | 93 ++++++++++++++++ 9 files changed, 260 insertions(+), 34 deletions(-) diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index ee3ece08c..a9427b618 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -958,6 +958,17 @@ 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. +The mixed W4A8 community encoder follows the same component-path contract: + +```bash Overlay +--component-paths.text_encoder \ + Winnougan/MiniMax-H3-INT4_Convrot_ComfyUI/qwen3vl_32b_minimax_h3-w4a8_convrot.safetensors +``` + +Install `comfy-kitchen>=0.2.27` and omit `--quantization`. SGLang automatically +loads its W4A8 language linears and tensorwise INT8 embedding; the unmarked +vision tower remains BF16. + diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 6cfa5cbcd..45663f939 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -195,10 +195,10 @@ backend. comfy-w4a8-convrot Safetensors with serialized asym_w4a8_int8 layer metadata and packed weights - --transformer-weights-path - MiniMax-H3 FL2VA / Ref2VA DiTs + --transformer-weights-path or --component-paths.text_encoder + MiniMax-H3 FL2VA / Ref2VA DiTs and its native Qwen3-VL encoder comfy-kitchen>=0.2.27 - Auto-detected; omit --quantization. Requires SM80+ and validates packed weights, group/channel scales, and optional codebooks before model construction. TP must preserve ConvRot group boundaries; offload is supported and FSDP is not. + Auto-detected; omit --quantization. Requires SM80+ and validates packed weights, group/channel scales, and optional codebooks before model construction. Mixed encoder files may keep their embedding tensorwise INT8. TP must preserve ConvRot group boundaries; offload is supported and FSDP is not. qvg-kv diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py index ebb68a2e4..220177410 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/base_config.py @@ -46,6 +46,10 @@ class QuantizationConfig(SRTQuantizationConfig): """Whether a row-parallel shard preserves this format's input layout.""" return True + def quantizes_embedding(self, prefix: str) -> bool: + """Whether this checkpoint config owns the named embedding table.""" + return False + def remap_checkpoint_prefixes(self, param_names_mapping: dict) -> None: """Translate checkpoint module names to the native model namespace.""" return diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py index ea89b5cbd..38a12f667 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py @@ -16,8 +16,12 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor QuantizeMethodBase, ) from sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8 import ( + KitchenInt8EmbeddingMethod, KitchenW4A8LinearMethod, ) +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) from sglang.multimodal_gen.runtime.platforms import current_platform @@ -44,10 +48,15 @@ class KitchenW4A8Config(QuantizationConfig): self.selected: list[str] = [] for prefix, marker in layer_markers.items(): - if marker.get("format") != "asym_w4a8_int8": + marker_format = marker.get("format") + if marker_format == "int8_tensorwise" and marker.get( + "_is_tensorwise_scalar" + ): + continue + if marker_format != "asym_w4a8_int8": raise ValueError( f"Unsupported Comfy W4A8 format for {prefix!r}: " - f"{marker.get('format')!r}" + f"{marker_format!r}" ) if marker.get("convrot") is not True: raise ValueError( @@ -80,11 +89,24 @@ class KitchenW4A8Config(QuantizationConfig): def get_quant_method( self, layer: torch.nn.Module, prefix: str ) -> QuantizeMethodBase | None: + marker = self.layer_markers.get(prefix) + if isinstance(layer, VocabParallelEmbedding): + if marker is None: + return None + if marker.get("format") != "int8_tensorwise" or not marker.get( + "_is_tensorwise_scalar" + ): + raise ValueError( + f"Unsupported quantized embedding marker for {prefix!r}: {marker}" + ) + self.selected.append(prefix) + return KitchenInt8EmbeddingMethod() if not isinstance(layer, LinearBase): return None - marker = self.layer_markers.get(prefix) if marker is None: return UnquantizedLinearMethod() + if marker.get("format") != "asym_w4a8_int8": + raise ValueError(f"Unsupported quantized linear marker for {prefix!r}") group_size = int(marker.get("group_size", 16)) convrot_group_size = int(marker.get("convrot_groupsize", 256)) @@ -120,7 +142,7 @@ class KitchenW4A8Config(QuantizationConfig): self, prefix: str, input_size_per_partition: int ) -> bool: marker = self.layer_markers.get(prefix) - if marker is None: + if marker is None or marker.get("format") != "asym_w4a8_int8": return True return self._supports_input_size( input_size_per_partition, @@ -130,3 +152,11 @@ class KitchenW4A8Config(QuantizationConfig): def get_scaled_act_names(self) -> list[str]: return [] + + def quantizes_embedding(self, prefix: str) -> bool: + marker = self.layer_markers.get(prefix) + return bool( + marker is not None + and marker.get("format") == "int8_tensorwise" + and marker.get("_is_tensorwise_scalar") + ) diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py index 193c86b30..2a47826cb 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py @@ -14,6 +14,23 @@ try: except ImportError: # pragma: no cover - optional dependency w4a8_int8_linear = None +_OUTPUT_DTYPE_CODE = {torch.float32: 0, torch.float16: 1, torch.bfloat16: 2} + + +def _register_weight( + layer: torch.nn.Module, + name: str, + shape: tuple[int, ...], + dtype: torch.dtype, + weight_attrs: dict, + parallel_dims: dict[str, int] | None = None, +) -> None: + weight = Parameter(torch.empty(shape, dtype=dtype), requires_grad=False) + if parallel_dims is not None: + set_weight_attrs(weight, parallel_dims) + set_weight_attrs(weight, weight_attrs) + layer.register_parameter(name, weight) + class KitchenW4A8LinearMethod(LinearMethodBase): """Load packed INT4 weights and execute the W4A8 ConvRot kernel.""" @@ -36,21 +53,6 @@ class KitchenW4A8LinearMethod(LinearMethodBase): self.has_codebook = has_codebook self.has_correction = has_correction - @staticmethod - def _register_weight( - layer: torch.nn.Module, - name: str, - shape: tuple[int, ...], - dtype: torch.dtype, - weight_attrs: dict, - parallel_dims: dict[str, int] | None = None, - ) -> None: - weight = Parameter(torch.empty(shape, dtype=dtype), requires_grad=False) - if parallel_dims is not None: - set_weight_attrs(weight, parallel_dims) - set_weight_attrs(weight, weight_attrs) - layer.register_parameter(name, weight) - def create_weights( self, layer: torch.nn.Module, @@ -70,7 +72,7 @@ class KitchenW4A8LinearMethod(LinearMethodBase): ) output_size_per_partition = sum(output_partition_sizes) - self._register_weight( + _register_weight( layer, "weight", (output_size_per_partition, input_size_per_partition // 2), @@ -78,7 +80,7 @@ class KitchenW4A8LinearMethod(LinearMethodBase): extra_weight_attrs, {"input_dim": 1, "output_dim": 0}, ) - self._register_weight( + _register_weight( layer, "weight_s_rel", (output_size_per_partition, input_size_per_partition // self.group_size), @@ -86,7 +88,7 @@ class KitchenW4A8LinearMethod(LinearMethodBase): extra_weight_attrs, {"input_dim": 1, "output_dim": 0}, ) - self._register_weight( + _register_weight( layer, "weight_s_channel", (output_size_per_partition,), @@ -95,7 +97,7 @@ class KitchenW4A8LinearMethod(LinearMethodBase): {"output_dim": 0}, ) if self.has_codebook: - self._register_weight( + _register_weight( layer, "weight_codebook", (16,), @@ -105,7 +107,7 @@ class KitchenW4A8LinearMethod(LinearMethodBase): else: layer.register_parameter("weight_codebook", None) if self.has_correction: - self._register_weight( + _register_weight( layer, "weight_correction", ( @@ -146,4 +148,62 @@ class KitchenW4A8LinearMethod(LinearMethodBase): return output -__all__ = ["KitchenW4A8LinearMethod"] +class KitchenInt8EmbeddingMethod(LinearMethodBase): + """Gather and dequantize only the selected rows of a tensorwise INT8 table.""" + + def __init__(self) -> None: + try: + torch.ops.comfy_kitchen.dequantize_int8_embedding + except AttributeError as exc: + raise ImportError( + "Tensorwise INT8 embeddings require comfy-kitchen>=0.2.27 " + "(`pip install -U comfy-kitchen`)." + ) from exc + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ) -> None: + del input_size, output_size + self.output_dtype = params_dtype + _register_weight( + layer, + "weight", + (sum(output_partition_sizes), input_size_per_partition), + torch.int8, + extra_weight_attrs, + {"input_dim": 1, "output_dim": 0}, + ) + _register_weight( + layer, + "weight_scale", + (), + torch.float32, + extra_weight_attrs, + ) + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + raise NotImplementedError("Kitchen INT8 embedding weights support lookup only") + + def embedding(self, layer: torch.nn.Module, input_: torch.Tensor) -> torch.Tensor: + return torch.ops.comfy_kitchen.dequantize_int8_embedding( + layer.weight, + layer.weight_scale, + input_, + 0, + _OUTPUT_DTYPE_CODE[self.output_dtype], + ) + + +__all__ = ["KitchenInt8EmbeddingMethod", "KitchenW4A8LinearMethod"] 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 c690e5ea7..fe5d72a99 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 @@ -34,6 +34,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, ) +from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_w4a8_config import ( + KitchenW4A8Config, +) from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, ComponentLoader, @@ -133,13 +136,13 @@ def _get_encoder_quant_config( mapping_fn = get_param_names_mapping(mapping) def name_mapper(name: str) -> str: - mapped_name, merge_index, _ = mapping_fn(name) + mapped_name, merge_index, _ = mapping_fn(f"{name}.weight") if merge_index is not None: raise ValueError( "Comfy quantized component weights cannot use a " "stacked parameter-name mapping" ) - return mapped_name + return mapped_name.removesuffix(".weight") markers = inspect_comfy_quant_markers( [component_weights_path], @@ -308,7 +311,7 @@ def _require_quantized_encoder_layers( f"The native {type(model).__name__} implementation does not construct " f"quantized linear layers for {component_name!r}" ) - if isinstance(quant_config, (ComfyFp8Config, KitchenInt8Config)): + if isinstance(quant_config, (ComfyFp8Config, KitchenInt8Config, KitchenW4A8Config)): missing = set(quant_config.layer_markers) - set(quant_config.selected) if missing: raise ComponentCheckpointUnsupportedError( diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py index fe4f8849d..2b4ffb522 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py @@ -27,6 +27,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import ( WeightOnlyFP8Linear, WeightOnlyFP8RowParallelLinear, ) +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) 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.qwen3vl_vision import ( @@ -504,9 +507,21 @@ class Qwen3VLTextModel(nn.Module): self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.embed_tokens = nn.Embedding( - config.vocab_size, config.hidden_size, self.padding_idx - ) + embedding_prefix = add_prefix("embed_tokens", prefix) + if quant_config is not None and quant_config.quantizes_embedding( + embedding_prefix + ): + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + params_dtype=torch.get_default_dtype(), + quant_config=quant_config, + prefix=embedding_prefix, + ) + else: + self.embed_tokens = nn.Embedding( + config.vocab_size, config.hidden_size, self.padding_idx + ) self.layers = nn.ModuleList( [ Qwen3VLTextDecoderLayer( diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index 67b96c1ad..862416c6e 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -211,6 +211,14 @@ def inspect_comfy_quant_markers( continue weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"] scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"] + if weight_dtype == "I8" and scale_dtype == "F32" and scale_shape == (): + if len(weight_shape) != 2: + raise ValueError( + f"Comfy tensorwise INT8 layer {prefix!r} needs a 2D weight, " + f"got {weight_shape}" + ) + marker["_is_tensorwise_scalar"] = True + continue if weight_dtype != "I8" or scale_dtype != "F32": raise ValueError( f"Comfy INT8 layer {prefix!r} needs I8 weights and F32 scales, " @@ -243,6 +251,8 @@ def resolve_comfy_checkpoint_quantization( return KitchenInt8Config(layer_markers=layer_markers) if formats == ["asym_w4a8_int8"]: return KitchenW4A8Config(layer_markers) + if formats == ["asym_w4a8_int8", "int8_tensorwise"]: + return KitchenW4A8Config(layer_markers) if formats == ["float8_e4m3fn"]: return ComfyFp8Config(layer_markers) if formats == ["mxfp8"]: 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 60aa00c48..cc01cd334 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 @@ -13,6 +13,9 @@ from sglang.multimodal_gen.runtime.layers.linear import LinearBase from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, ) +from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_w4a8_config import ( + KitchenW4A8Config, +) from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, @@ -30,6 +33,7 @@ from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import ( MiniMaxH3ConditioningProjection, MiniMaxH3Qwen3VLEncoder, ) +from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLTextModel class TestTextEncoderClassResolution(unittest.TestCase): @@ -80,6 +84,41 @@ class TestTextEncoderClassResolution(unittest.TestCase): # e.g. CLIP/Mistral/Qwen text encoders are not encoder-decoder. self.assertIs(self._resolve(False, ["CLIPTextModel"]), transformers.AutoModel) + def test_qwen_text_model_constructs_checkpoint_owned_embedding(self): + config = SimpleNamespace( + pad_token_id=0, + vocab_size=64, + hidden_size=256, + num_hidden_layers=0, + rms_norm_eps=1e-6, + ) + quant_config = mock.Mock() + quant_config.quantizes_embedding.return_value = True + replacement = nn.Embedding(64, 256) + with mock.patch( + "sglang.multimodal_gen.runtime.models.encoders.qwen3vl." + "VocabParallelEmbedding", + return_value=replacement, + ) as embedding_cls: + model = Qwen3VLTextModel( + config, + quant_config=quant_config, + use_tensor_parallel=True, + prefix="model.language_model", + ) + + self.assertIs(model.embed_tokens, replacement) + quant_config.quantizes_embedding.assert_called_once_with( + "model.language_model.embed_tokens" + ) + embedding_cls.assert_called_once_with( + 64, + 256, + params_dtype=torch.get_default_dtype(), + quant_config=quant_config, + prefix="model.language_model.embed_tokens", + ) + def test_unknown_architecture_falls_back_to_automodel(self): self.assertIs(self._resolve(True, ["NotARealClass"]), transformers.AutoModel) @@ -358,6 +397,60 @@ class TestTextEncoderQuantization(unittest.TestCase): {"model.language_model.layers.0.self_attn.q_proj"}, ) + def test_mixed_w4a8_weight_file_maps_embedding_and_linear_markers(self): + self.get_quant_config.return_value = None + layers = { + "model.embed_tokens": {"format": "int8_tensorwise"}, + "model.layers.0.mlp.down_proj": { + "format": "asym_w4a8_int8", + "convrot": True, + "group_size": 16, + "convrot_groupsize": 256, + }, + } + with tempfile.NamedTemporaryFile(suffix=".safetensors") as checkpoint: + save_file( + { + "model.embed_tokens.weight": torch.ones((4, 256), dtype=torch.int8), + "model.embed_tokens.weight_scale": torch.tensor(0.25), + "model.layers.0.mlp.down_proj.weight": torch.ones( + (2, 128), dtype=torch.int8 + ), + "model.layers.0.mlp.down_proj.weight_s_rel": torch.ones( + (2, 16), dtype=torch.float8_e4m3fn + ), + "model.layers.0.mlp.down_proj.weight_s_channel": torch.ones(2), + "model.layers.0.mlp.down_proj.weight_codebook": torch.ones(16), + }, + checkpoint.name, + metadata={"_quantization_metadata": json.dumps({"layers": layers})}, + ) + model_config = SimpleNamespace(quant_config=None) + with mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "text_encoder_loader.get_quant_config_from_safetensors_metadata", + return_value=None, + ): + _configure_encoder_quantization( + model_config, + MiniMaxH3Qwen3VLEncoder, + {}, + "/model/text_encoder", + checkpoint.name, + "text_encoder", + ) + + self.assertIsInstance(model_config.quant_config, KitchenW4A8Config) + self.assertTrue( + model_config.quant_config.quantizes_embedding( + "model.language_model.embed_tokens" + ) + ) + self.assertIn( + "model.language_model.layers.0.mlp.down_proj", + model_config.quant_config.layer_markers, + ) + def test_encoder_must_use_native_loader(self): model_config = SimpleNamespace(quant_config=None) with self.assertRaisesRegex(