From 2d84de5e6918c81741efa831c415077a4de6b973 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 24 Aug 2026 10:32:21 +0800 Subject: [PATCH] [diffusion] feat: support loading serialized comfy w4a8 checkpoints (#36036) --- .../cookbook/diffusion/MiniMax/MiniMax-H3.mdx | 8 + docs/docs/sglang-diffusion/quantization.mdx | 8 + .../configs/kitchen_w4a8_config.py | 132 ++++++++++++++++ .../layers/quantization/kitchen_w4a8.py | 149 ++++++++++++++++++ .../component_loaders/transformer_loader.py | 1 + .../runtime/loader/transformer_load_utils.py | 8 + .../runtime/utils/quantization_utils.py | 101 +++++++++++- .../test/unit/test_transformer_quant.py | 75 +++++++++ 8 files changed, 481 insertions(+), 1 deletion(-) create mode 100644 python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py create mode 100644 python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index 321df36a9..ee3ece08c 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -298,6 +298,14 @@ layers in their original dtype, and reuses SRT's MXFP8 dense kernels. The selected SRT backend must support MXFP8 on the target GPU; FSDP is rejected for this mixed per-layer layout. +W4A8 ConvRot DiT files use the same flagless flow. With +`comfy-kitchen>=0.2.27`, pass a file such as +`starsfriday/MiniMax-H3-w4a8/minimax_h3_fl2va_pruned_w4a8_mixed.safetensors` +to `--transformer-weights-path`; SGLang reads the serialized per-layer metadata +and packed INT4 tensors automatically. Do not add `--quantization`. TP remains +subject to each row-parallel shard preserving the checkpoint's ConvRot group +boundary, and FSDP is rejected. + ### Advanced: precomputed AdaLN cache The [model card](https://huggingface.co/MiniMaxAI/MiniMax-H3) notes that about diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 642987a1d..0d891df01 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -184,6 +184,14 @@ backend. SRT's platform MXFP8 backend Serialized metadata is auto-detected. NVIDIA and ROCm reuse SRT's dense MXFP8 kernels; Ascend keeps its native online path. Mixed per-layer checkpoints do not support FSDP. + + comfy-w4a8-convrot + Safetensors with serialized asym_w4a8_int8 layer metadata and packed weights + --transformer-weights-path + MiniMax-H3 FL2VA / Ref2VA DiTs + 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. + qvg-kv Unquantized model with runtime causal KV-cache compression 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 new file mode 100644 index 000000000..ea89b5cbd --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_w4a8_config.py @@ -0,0 +1,132 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Config for serialized Comfy Kitchen W4A8 ConvRot weights.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from sglang.multimodal_gen.runtime.layers.linear import ( + LinearBase, + UnquantizedLinearMethod, +) +from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8 import ( + KitchenW4A8LinearMethod, +) +from sglang.multimodal_gen.runtime.platforms import current_platform + + +class KitchenW4A8Config(QuantizationConfig): + """Dispatch each linear from its serialized ``asym_w4a8_int8`` marker.""" + + def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None: + super().__init__() + if current_platform.is_mps(): + raise ValueError("Serialized W4A8 checkpoints are not supported on MPS") + if current_platform.is_cuda(): + capability = current_platform.get_device_capability() + if ( + capability is not None + and capability.to_int() < self.get_min_capability() + ): + raise ValueError( + "Serialized W4A8 checkpoints require CUDA compute capability " + f">= {self.get_min_capability() / 10:.1f}; got " + f"{capability.to_int() / 10:.1f}" + ) + self.layer_markers = layer_markers + self.checkpoint_uses_native_qkv_layout = True + self.selected: list[str] = [] + + for prefix, marker in layer_markers.items(): + if marker.get("format") != "asym_w4a8_int8": + raise ValueError( + f"Unsupported Comfy W4A8 format for {prefix!r}: " + f"{marker.get('format')!r}" + ) + if marker.get("convrot") is not True: + raise ValueError( + f"Serialized W4A8 layer {prefix!r} must set convrot=true" + ) + + @classmethod + def get_name(cls) -> str: + return "kitchen_w4a8" + + @classmethod + def get_supported_act_dtypes(cls) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16] + + @classmethod + def get_min_capability(cls) -> int: + return 80 + + @classmethod + def get_config_filenames(cls) -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> KitchenW4A8Config: + raise ValueError( + "kitchen_w4a8 is inferred from per-layer checkpoint metadata; " + "it is not an online quantization method" + ) + + def get_quant_method( + self, layer: torch.nn.Module, prefix: str + ) -> QuantizeMethodBase | None: + if not isinstance(layer, LinearBase): + return None + marker = self.layer_markers.get(prefix) + if marker is None: + return UnquantizedLinearMethod() + + group_size = int(marker.get("group_size", 16)) + convrot_group_size = int(marker.get("convrot_groupsize", 256)) + if not self._supports_input_size( + layer.input_size, group_size, convrot_group_size + ): + raise ValueError( + f"Serialized W4A8 layer {prefix!r} has input size " + f"{layer.input_size}, incompatible with group_size={group_size} " + f"and convrot_groupsize={convrot_group_size}" + ) + self.selected.append(prefix) + return KitchenW4A8LinearMethod( + group_size=group_size, + convrot_group_size=convrot_group_size, + has_codebook=bool(marker.get("_has_codebook")), + has_correction=bool(marker.get("_has_correction")), + ) + + @staticmethod + def _supports_input_size( + input_size: int, group_size: int, convrot_group_size: int + ) -> bool: + return ( + group_size >= 4 + and (16 % group_size == 0 or group_size % 16 == 0) + and input_size % 16 == 0 + and input_size % group_size == 0 + and input_size % convrot_group_size == 0 + ) + + def supports_input_partition( + self, prefix: str, input_size_per_partition: int + ) -> bool: + marker = self.layer_markers.get(prefix) + if marker is None: + return True + return self._supports_input_size( + input_size_per_partition, + int(marker.get("group_size", 16)), + int(marker.get("convrot_groupsize", 256)), + ) + + def get_scaled_act_names(self) -> list[str]: + return [] diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py new file mode 100644 index 000000000..193c86b30 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py @@ -0,0 +1,149 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Serialized grouped W4A8 linear backed by Comfy Kitchen.""" + +from __future__ import annotations + +import torch +from torch.nn.parameter import Parameter + +from sglang.multimodal_gen.runtime.layers.linear import LinearMethodBase +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs + +try: + from comfy_kitchen import w4a8_int8_linear +except ImportError: # pragma: no cover - optional dependency + w4a8_int8_linear = None + + +class KitchenW4A8LinearMethod(LinearMethodBase): + """Load packed INT4 weights and execute the W4A8 ConvRot kernel.""" + + def __init__( + self, + *, + group_size: int, + convrot_group_size: int, + has_codebook: bool, + has_correction: bool, + ) -> None: + if w4a8_int8_linear is None: + raise ImportError( + "W4A8 checkpoints require comfy-kitchen>=0.2.27 " + "(`pip install -U comfy-kitchen`)." + ) + self.group_size = group_size + self.convrot_group_size = convrot_group_size + 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, + 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, params_dtype + if input_size_per_partition % self.convrot_group_size: + raise ValueError( + "W4A8 needs input_size_per_partition " + f"({input_size_per_partition}) divisible by ConvRot group size " + f"{self.convrot_group_size}" + ) + + output_size_per_partition = sum(output_partition_sizes) + self._register_weight( + layer, + "weight", + (output_size_per_partition, input_size_per_partition // 2), + torch.int8, + extra_weight_attrs, + {"input_dim": 1, "output_dim": 0}, + ) + self._register_weight( + layer, + "weight_s_rel", + (output_size_per_partition, input_size_per_partition // self.group_size), + torch.float8_e4m3fn, + extra_weight_attrs, + {"input_dim": 1, "output_dim": 0}, + ) + self._register_weight( + layer, + "weight_s_channel", + (output_size_per_partition,), + torch.float32, + extra_weight_attrs, + {"output_dim": 0}, + ) + if self.has_codebook: + self._register_weight( + layer, + "weight_codebook", + (16,), + torch.float32, + extra_weight_attrs, + ) + else: + layer.register_parameter("weight_codebook", None) + if self.has_correction: + self._register_weight( + layer, + "weight_correction", + ( + input_size_per_partition // self.group_size, + output_size_per_partition, + ), + torch.float32, + extra_weight_attrs, + {"input_dim": 0, "output_dim": 1}, + ) + else: + layer.register_parameter("weight_correction", None) + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + original_shape = x.shape + if x.dim() != 2: + x = x.reshape(-1, original_shape[-1]) + assert w4a8_int8_linear is not None + output = w4a8_int8_linear( + x.contiguous(), + layer.weight, + layer.weight_s_rel, + layer.weight_s_channel, + codebook=layer.weight_codebook, + correction=layer.weight_correction, + bias=bias, + group_size=self.group_size, + convrot_groupsize=self.convrot_group_size, + out_dtype=x.dtype, + ) + if len(original_shape) != 2: + output = output.reshape(*original_shape[:-1], output.shape[-1]) + return output + + +__all__ = ["KitchenW4A8LinearMethod"] diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 55d4ffee9..96e9e5b3f 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -370,6 +370,7 @@ class TransformerLoader(ComponentLoader): quantized_cpu_load_supported=( quant_spec.gguf_file is not None or quant_spec.is_serialized_kitchen_int8 + or quant_spec.is_serialized_kitchen_w4a8 ), ) ) diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 6e8d65482..b5b7e0cb1 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -21,6 +21,9 @@ from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig 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.configs.nunchaku_config import ( NunchakuConfig, _patch_nunchaku_scales, @@ -173,11 +176,16 @@ class TransformerQuantLoadSpec: and self.quant_config.is_checkpoint_int8_serialized ) + @property + def is_serialized_kitchen_w4a8(self) -> bool: + return isinstance(self.quant_config, KitchenW4A8Config) + @property def uses_comfy_layer_markers(self) -> bool: return ( self.is_comfy_fp8 or self.is_serialized_kitchen_int8 + or self.is_serialized_kitchen_w4a8 or ( _get_quant_config_name(self.quant_config) == "mxfp8" and self.quant_config.layer_markers is not None diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index f1ba647bc..67b96c1ad 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -16,6 +16,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8 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.mxfp8 import MXFP8Config from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo @@ -41,6 +44,34 @@ def inspect_comfy_quant_markers( metadata = checkpoint.metadata() or {} if quant_format := metadata.get("quant_format"): global_quant_formats.add(quant_format.lower()) + serialized_metadata = metadata.get("_quantization_metadata") + if serialized_metadata is not None: + try: + metadata_config = json.loads(serialized_metadata) + except json.JSONDecodeError as exc: + raise ValueError( + f"Invalid _quantization_metadata in {path}" + ) from exc + if not isinstance(metadata_config, dict): + raise ValueError( + f"_quantization_metadata in {path} must contain an object" + ) + metadata_layers = metadata_config.get("layers") + if not isinstance(metadata_layers, dict): + raise ValueError( + f"_quantization_metadata in {path} must contain a layers object" + ) + for prefix, marker in metadata_layers.items(): + if not isinstance(marker, dict): + raise ValueError( + f"Comfy quantization metadata for {prefix!r} must be an object" + ) + previous = raw_markers.get(prefix) + if previous is not None and previous != marker: + raise ValueError( + f"Conflicting Comfy quantization markers for {prefix!r}" + ) + raw_markers[prefix] = marker for key in checkpoint.keys(): tensor_slice = checkpoint.get_slice(key) checkpoint_meta[key] = ( @@ -99,7 +130,17 @@ def inspect_comfy_quant_markers( for prefix, marker in raw_markers.items(): marker_format = marker.get("format") required = {f"{prefix}.weight", f"{prefix}.weight_scale"} - if marker_format not in ("float8_e4m3fn", "int8_tensorwise"): + if marker_format == "asym_w4a8_int8": + required = { + f"{prefix}.weight", + f"{prefix}.weight_s_rel", + f"{prefix}.weight_s_channel", + } + if marker_format not in ( + "float8_e4m3fn", + "int8_tensorwise", + "asym_w4a8_int8", + ): continue missing = required - checkpoint_meta.keys() if missing: @@ -112,6 +153,62 @@ def inspect_comfy_quant_markers( "static" if f"{prefix}.input_scale" in checkpoint_meta else "dynamic" ) continue + if marker_format == "asym_w4a8_int8": + weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"] + scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_s_rel"] + channel_dtype, channel_shape = checkpoint_meta[f"{prefix}.weight_s_channel"] + group_size = int(marker.get("group_size", 16)) + if group_size < 4: + raise ValueError( + f"Comfy W4A8 layer {prefix!r} has invalid group_size={group_size}" + ) + if weight_dtype != "I8" or scale_dtype != "F8_E4M3": + raise ValueError( + f"Comfy W4A8 layer {prefix!r} needs I8 weights and FP8 " + f"group scales, got {weight_dtype} and {scale_dtype}" + ) + if channel_dtype != "F32": + raise ValueError( + f"Comfy W4A8 layer {prefix!r} needs F32 channel scales, " + f"got {channel_dtype}" + ) + if len(weight_shape) != 2: + raise ValueError( + f"Comfy W4A8 layer {prefix!r} needs a 2D packed weight, " + f"got {weight_shape}" + ) + logical_input_size = weight_shape[1] * 2 + expected_scale_shape = (weight_shape[0], logical_input_size // group_size) + if scale_shape != expected_scale_shape or channel_shape != ( + weight_shape[0], + ): + raise ValueError( + f"Comfy W4A8 layer {prefix!r} has incompatible weight/scale " + f"shapes: {weight_shape}, {scale_shape}, and {channel_shape}" + ) + codebook_key = f"{prefix}.weight_codebook" + correction_key = f"{prefix}.weight_correction" + marker["_has_codebook"] = codebook_key in checkpoint_meta + marker["_has_correction"] = correction_key in checkpoint_meta + if marker["_has_codebook"] and checkpoint_meta[codebook_key] != ( + "F32", + (16,), + ): + raise ValueError( + f"Comfy W4A8 layer {prefix!r} needs an F32[16] codebook" + ) + expected_correction = ( + logical_input_size // group_size, + weight_shape[0], + ) + if marker["_has_correction"] and checkpoint_meta[correction_key] != ( + "F32", + expected_correction, + ): + raise ValueError( + f"Comfy W4A8 layer {prefix!r} has an incompatible correction tensor" + ) + continue weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"] scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"] if weight_dtype != "I8" or scale_dtype != "F32": @@ -144,6 +241,8 @@ def resolve_comfy_checkpoint_quantization( formats = sorted({str(marker.get("format")) for marker in layer_markers.values()}) if formats == ["int8_tensorwise"]: return KitchenInt8Config(layer_markers=layer_markers) + if formats == ["asym_w4a8_int8"]: + 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_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index 48810e0e3..74cde2416 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -57,6 +57,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ( 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.configs.nunchaku_config import ( NunchakuConfig, ) @@ -346,6 +349,78 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertTrue(config.supports_input_partition("blocks.0.mlp.fc1", 6400)) self.assertFalse(config.supports_input_partition("blocks.0.mlp.fc1", 3200)) + def test_minimax_h3_w4a8_metadata_resolves_serialized_kitchen(self): + metadata = { + "_quantization_metadata": json.dumps( + { + "layers": { + "blocks.0.mlp.fc1": { + "format": "asym_w4a8_int8", + "convrot": True, + "group_size": 16, + "convrot_groupsize": 256, + } + } + } + ) + } + with tempfile.NamedTemporaryFile(suffix=".safetensors") as checkpoint: + save_file( + { + "blocks.0.mlp.fc1.weight": torch.ones((2, 128), dtype=torch.int8), + "blocks.0.mlp.fc1.weight_s_rel": torch.ones( + (2, 16), dtype=torch.float8_e4m3fn + ), + "blocks.0.mlp.fc1.weight_s_channel": torch.ones(2), + "blocks.0.mlp.fc1.weight_codebook": torch.ones(16), + }, + checkpoint.name, + metadata=metadata, + ) + + _, markers = inspect_minimax_h3_safetensors([checkpoint.name]) + + config = resolve_minimax_h3_checkpoint_quantization(markers) + self.assertIsInstance(config, KitchenW4A8Config) + self.assertTrue(markers["blocks.0.mlp.fc1"]["_has_codebook"]) + self.assertTrue(config.supports_input_partition("blocks.0.mlp.fc1", 256)) + self.assertFalse(config.supports_input_partition("blocks.0.mlp.fc1", 128)) + + @patch( + "sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8." + "w4a8_int8_linear", + new=object(), + ) + def test_serialized_w4a8_constructs_packed_weights_and_scales(self): + config = KitchenW4A8Config( + { + "proj": { + "format": "asym_w4a8_int8", + "convrot": True, + "group_size": 16, + "convrot_groupsize": 256, + "_has_codebook": True, + "_has_correction": False, + } + } + ) + layer = ReplicatedLinear( + 256, + 3, + bias=False, + params_dtype=torch.bfloat16, + quant_config=config, + prefix="proj", + ) + + self.assertEqual(layer.weight.shape, (3, 128)) + self.assertEqual(layer.weight.dtype, torch.int8) + self.assertEqual(layer.weight_s_rel.shape, (3, 16)) + self.assertEqual(layer.weight_s_rel.dtype, torch.float8_e4m3fn) + self.assertEqual(layer.weight_s_channel.shape, (3,)) + self.assertEqual(layer.weight_codebook.shape, (16,)) + self.assertIsNone(layer.weight_correction) + @patch( "sglang.multimodal_gen.runtime.layers.quantization.kitchen_int8." "_load_comfy_kitchen"