diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index 81d4e06de..99e46a15e 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -1060,6 +1060,19 @@ W4A4 Qwen3-VL files use the same overlay, for example: This checkpoint keeps its unmarked embedding and vision tower in their source precision; no explicit component quantization option is needed. +The official Comfy NVFP4-AWQ encoder uses the same flagless overlay: + +```bash Overlay +--component-paths.text_encoder \ + Comfy-Org/MiniMax-H3/text_encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors +``` + +SGLang auto-detects its row-wise INT8 embedding, NVFP4 language linears, and +AWQ input pre-scales. The weights stay compressed at rest; each active linear +is dequantized for a BF16/FP16 matrix multiplication, so this path primarily +reduces resident memory rather than encoder latency. Do not add a component or +transformer quantization option. + The same component option accepts a self-describing Quanto qint8 file without an additional quantization flag: diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index bc8fee5b0..27ee163a9 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -53,12 +53,12 @@ repo contains multiple candidate checkpoints, pass `--transformer-weights-path` explicitly. MiniMax-H3 is a verified example for Comfy safetensors with per-layer metadata, -including `pruned_fp8_scaled` and serialized ConvRot INT8. Other Comfy FP8 -exports are also auto-detected: the presence of an input scale selects static -activation scaling, while its absence selects dynamic scaling. Pass one selected -FL2VA or Ref2VA DiT file by local path, `owner/repo/path/file.safetensors`, or -direct Hugging Face file URL; do not combine it with `--quantization`. Its GGUF -usage is documented in +including `pruned_fp8_scaled`, serialized ConvRot formats, and the official +NVFP4-AWQ Qwen3-VL encoder. Other Comfy FP8 exports are also auto-detected: the +presence of an input scale selects static activation scaling, while its absence +selects dynamic scaling. Pass one selected DiT or component file by local path, +`owner/repo/path/file.safetensors`, or direct Hugging Face file URL; do not +combine it with an explicit quantization option. Its GGUF usage is documented in the [MiniMax-H3 cookbook](/cookbook/diffusion/MiniMax/MiniMax-H3#pre-quantized-gguf-transformer). ## Quantized Component Repositories @@ -212,6 +212,14 @@ backend. comfy-kitchen Auto-detected; omit --quantization. Each layer dispatches to its serialized W4A4 or INT8 ConvRot kernel. CUDA requires SM75+; TP must preserve each format's quantization and ConvRot group boundaries. Offload is supported and FSDP is not. + + comfy-nvfp4-full-precision + Safetensors with serialized nvfp4 and optional row-wise int8_tensorwise layer metadata + --component-paths.text_encoder + MiniMax-H3 native Qwen3-VL encoder + None + Auto-detected; omit explicit quantization. Preserves packed storage, high-nibble-first weights, swizzled block scales, and AWQ input pre-scales. Each active NVFP4 matrix is dequantized for BF16/FP16 compute, so this is a memory path rather than a native FP4 speed path. + quanto-int8 One native encoder safetensors file with an embedded Quanto quantization map diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py new file mode 100644 index 000000000..6c8fcd0ca --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py @@ -0,0 +1,257 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Portable full-precision execution for Comfy NVFP4 checkpoints.""" + +from __future__ import annotations + +from typing import Any + +import torch +import torch.nn.functional as F +from torch import nn + +from sglang.multimodal_gen.runtime.layers.linear import ( + LinearBase, + UnquantizedLinearMethod, +) +from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( + QuantizeMethodBase, +) +from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( + ModelOptFp4Config, + ModelOptFp4LinearMethod, + _swizzled_nvfp4_scales_to_linear, +) +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs +from sglang.srt.layers.quantization.dequantization import dequantize_nvfp4 + + +def _register_parameter( + layer: nn.Module, + name: str, + data: torch.Tensor, + weight_attrs: dict[str, Any], + parallel_dims: dict[str, int] | None = None, +) -> None: + parameter = nn.Parameter(data, requires_grad=False) + if parallel_dims is not None: + set_weight_attrs(parameter, parallel_dims) + set_weight_attrs(parameter, weight_attrs) + layer.register_parameter(name, parameter) + + +class ComfyRowwiseInt8EmbeddingMethod(QuantizeMethodBase): + """Gather and dequantize only selected rows of an INT8 embedding.""" + + def create_weights( + self, + layer: 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: Any, + ) -> None: + del input_size, output_size + self.output_dtype = params_dtype + output_size_per_partition = sum(output_partition_sizes) + _register_parameter( + layer, + "weight", + torch.empty( + output_size_per_partition, + input_size_per_partition, + dtype=torch.int8, + ), + extra_weight_attrs, + {"input_dim": 1, "output_dim": 0}, + ) + _register_parameter( + layer, + "weight_scale", + torch.empty(output_size_per_partition, 1, dtype=torch.float32), + extra_weight_attrs, + {"output_dim": 0}, + ) + + def apply( + self, + layer: nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + raise NotImplementedError("Comfy INT8 embedding weights support lookup only") + + def embedding(self, layer: nn.Module, input_: torch.Tensor) -> torch.Tensor: + weight = F.embedding(input_, layer.weight).to(self.output_dtype) + scale = F.embedding(input_, layer.weight_scale).to(self.output_dtype) + return weight * scale + + +class ComfyFullPrecisionNvfp4LinearMethod(ModelOptFp4LinearMethod): + """Keep NVFP4 storage and dequantize one active Linear for its matmul.""" + + def __init__( + self, + quant_config: ComfyNvfp4Config, + *, + has_pre_quant_scale: bool, + ) -> None: + self.quant_config = quant_config + self.has_pre_quant_scale = has_pre_quant_scale + + def create_weights( + self, + layer: 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: Any, + ) -> None: + if len(output_partition_sizes) != 1: + raise ValueError( + "Comfy full_precision_matrix_mult does not support fused linears" + ) + super().create_weights( + layer, + input_size_per_partition, + output_partition_sizes, + input_size, + output_size, + params_dtype, + **extra_weight_attrs, + ) + # Comfy uses runtime activations directly for this weight-only path. + layer.register_parameter("input_scale", None) + if not self.has_pre_quant_scale: + return + _register_parameter( + layer, + "pre_quant_scale", + torch.empty(input_size_per_partition, dtype=params_dtype), + extra_weight_attrs, + {"input_dim": 0}, + ) + + def process_weights_after_loading(self, layer: nn.Module) -> None: + # The portable path consumes the serialized representation directly. + # ModelOpt's inherited hook instead prepares a Blackwell-only kernel. + return + + def apply( + self, + layer: nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + if self.has_pre_quant_scale: + x = x * layer.pre_quant_scale + weight_scale = _swizzled_nvfp4_scales_to_linear(layer.weight_scale) + weight = dequantize_nvfp4( + layer.weight, + weight_scale, + layer.weight_scale_2, + out_dtype=x.dtype, + high_nibble_first=True, + ) + return F.linear(x, weight, bias) + + +class ComfyNvfp4Config(ModelOptFp4Config): + """Dispatch full-precision Comfy NVFP4 linears and their INT8 embedding.""" + + checkpoint_uses_comfy_quantization = True + + def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None: + super().__init__( + is_checkpoint_nvfp4_serialized=True, + group_size=16, + exclude_modules=[], + checkpoint_uses_comfy_quantization=True, + ) + self.layer_markers = layer_markers + self.selected: list[str] = [] + for prefix, marker in layer_markers.items(): + marker_format = marker.get("format") + if marker_format == "int8_tensorwise" and marker.get("_is_rowwise"): + continue + if marker_format != "nvfp4": + raise ValueError( + f"Unsupported Comfy NVFP4 companion for {prefix!r}: " + f"{marker_format!r}" + ) + if marker.get("full_precision_matrix_mult") is not True: + raise ValueError( + f"Comfy NVFP4 layer {prefix!r} must request " + "full_precision_matrix_mult" + ) + + @classmethod + def get_name(cls) -> str: + return "comfy_nvfp4" + + @classmethod + def get_supported_act_dtypes(cls) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16] + + @classmethod + def get_min_capability(cls) -> int: + return 0 + + @classmethod + def get_config_filenames(cls) -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> ComfyNvfp4Config: + raise ValueError( + "comfy_nvfp4 is inferred from per-layer checkpoint metadata; " + "it is not an online quantization method" + ) + + def get_quant_method( + self, layer: 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_rowwise" + ): + raise ValueError( + f"Unsupported quantized embedding marker for {prefix!r}: {marker}" + ) + self.selected.append(prefix) + return ComfyRowwiseInt8EmbeddingMethod() + if not isinstance(layer, LinearBase): + return None + if marker is None: + return UnquantizedLinearMethod() + if marker.get("format") != "nvfp4": + raise ValueError(f"Unsupported quantized linear marker for {prefix!r}") + self.selected.append(prefix) + return ComfyFullPrecisionNvfp4LinearMethod( + self, + has_pre_quant_scale=bool(marker.get("_has_pre_quant_scale")), + ) + + 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_rowwise") + ) + + +__all__ = [ + "ComfyFullPrecisionNvfp4LinearMethod", + "ComfyNvfp4Config", + "ComfyRowwiseInt8EmbeddingMethod", +] 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 c7383ccf2..d9a3d2021 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 @@ -31,6 +31,9 @@ from sglang.multimodal_gen.runtime.layers.linear import ( UnquantizedLinearMethod, ) from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config +from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import ( + ComfyNvfp4Config, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( QuantizationConfig, ) @@ -414,7 +417,13 @@ def _require_quantized_encoder_layers( ) if isinstance( quant_config, - (ComfyFp8Config, KitchenInt8Config, KitchenW4A4Config, KitchenW4A8Config), + ( + ComfyFp8Config, + ComfyNvfp4Config, + KitchenInt8Config, + KitchenW4A4Config, + KitchenW4A8Config, + ), ): expected = set(quant_config.layer_markers) selected = set(quant_config.selected) @@ -932,7 +941,7 @@ class TextEncoderLoader(ComponentLoader): if quant_config is not None and not isinstance(quant_config, GGUFConfig): postprocess_device: torch.device | None = local_torch_device - if isinstance(quant_config, QuantoInt8Config) or ( + if isinstance(quant_config, (ComfyNvfp4Config, QuantoInt8Config)) or ( isinstance(quant_config, KitchenInt8Config) and quant_config.is_checkpoint_int8_serialized ): diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index 62e11954b..1b8ff0519 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -13,6 +13,9 @@ from sglang.multimodal_gen.runtime.layers.quantization import ( get_quantization_config, ) from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config +from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import ( + ComfyNvfp4Config, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, ) @@ -84,6 +87,7 @@ def inspect_comfy_quant_markers( if key.endswith(".weight") and tensor_slice.get_dtype() in ( "F8_E4M3", "I8", + "U8", ): marked_dtype_weight_prefixes.add(key.removesuffix(".weight")) if not key.endswith(".comfy_quant"): @@ -139,11 +143,14 @@ def inspect_comfy_quant_markers( f"{prefix}.weight_s_rel", f"{prefix}.weight_s_channel", } + if marker_format == "nvfp4": + required.add(f"{prefix}.weight_scale_2") if marker_format not in ( "float8_e4m3fn", "int8_tensorwise", "asym_w4a8_int8", "convrot_w4a4", + "nvfp4", ): continue missing = required - checkpoint_meta.keys() @@ -240,6 +247,44 @@ def inspect_comfy_quant_markers( f"and convrot_groupsize={convrot_group_size}" ) continue + if marker_format == "nvfp4": + weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"] + scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"] + scale_2_dtype, scale_2_shape = checkpoint_meta[f"{prefix}.weight_scale_2"] + if weight_dtype != "U8" or scale_dtype != "F8_E4M3": + raise ValueError( + f"Comfy NVFP4 layer {prefix!r} needs U8 packed weights and " + f"FP8 block scales, got {weight_dtype} and {scale_dtype}" + ) + if scale_2_dtype != "F32" or scale_2_shape not in ((), (1,)): + raise ValueError( + f"Comfy NVFP4 layer {prefix!r} needs a scalar F32 " + f"weight_scale_2, got {scale_2_dtype}{scale_2_shape}" + ) + if len(weight_shape) != 2: + raise ValueError( + f"Comfy NVFP4 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 // 16) + if logical_input_size % 16 or scale_shape != expected_scale_shape: + raise ValueError( + f"Comfy NVFP4 layer {prefix!r} has incompatible weight/scale " + f"shapes: {weight_shape} and {scale_shape}" + ) + pre_quant_scale_key = f"{prefix}.pre_quant_scale" + marker["_has_pre_quant_scale"] = pre_quant_scale_key in checkpoint_meta + if marker["_has_pre_quant_scale"]: + pre_scale_dtype, pre_scale_shape = checkpoint_meta[pre_quant_scale_key] + if pre_scale_dtype not in ("BF16", "F16", "F32") or ( + pre_scale_shape != (logical_input_size,) + ): + raise ValueError( + f"Comfy NVFP4 layer {prefix!r} has an incompatible " + f"pre_quant_scale: {pre_scale_dtype}{pre_scale_shape}" + ) + continue if marker_format != "int8_tensorwise": continue weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"] @@ -262,6 +307,7 @@ def inspect_comfy_quant_markers( f"Comfy INT8 layer {prefix!r} has incompatible weight/scale " f"shapes: {weight_shape} and {scale_shape}" ) + marker["_is_rowwise"] = True mapped_markers: dict[str, dict[str, Any]] = {} for prefix, marker in raw_markers.items(): @@ -292,6 +338,8 @@ def resolve_comfy_checkpoint_quantization( return KitchenW4A4Config(layer_markers) if formats == ["float8_e4m3fn"]: return ComfyFp8Config(layer_markers) + if formats in (["nvfp4"], ["int8_tensorwise", "nvfp4"]): + return ComfyNvfp4Config(layer_markers) if formats == ["mxfp8"]: return MXFP8Config( is_checkpoint_fp8_serialized=True, 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 0d9453fc9..c75bac370 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 @@ -10,6 +10,11 @@ from safetensors.torch import save_file from torch import nn from sglang.multimodal_gen.runtime.layers.linear import LinearBase +from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import ( + ComfyFullPrecisionNvfp4LinearMethod, + ComfyNvfp4Config, + ComfyRowwiseInt8EmbeddingMethod, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, ) @@ -526,6 +531,107 @@ class TestTextEncoderQuantization(unittest.TestCase): model_config.quant_config.layer_markers, ) + def test_nvfp4_awq_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.self_attn.o_proj": { + "format": "nvfp4", + "full_precision_matrix_mult": True, + }, + } + with tempfile.NamedTemporaryFile(suffix=".safetensors") as checkpoint: + save_file( + { + "model.embed_tokens.weight": torch.ones((4, 64), dtype=torch.int8), + "model.embed_tokens.weight_scale": torch.ones(4, 1), + "model.layers.0.self_attn.o_proj.weight": torch.full( + (128, 32), 0x21, dtype=torch.uint8 + ), + "model.layers.0.self_attn.o_proj.weight_scale": torch.ones( + (128, 4), dtype=torch.float8_e4m3fn + ), + "model.layers.0.self_attn.o_proj.weight_scale_2": torch.tensor(0.5), + "model.layers.0.self_attn.o_proj.pre_quant_scale": torch.ones( + 64, dtype=torch.bfloat16 + ), + }, + 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, ComfyNvfp4Config) + self.assertTrue( + model_config.quant_config.quantizes_embedding( + "model.language_model.embed_tokens" + ) + ) + marker = model_config.quant_config.layer_markers[ + "model.language_model.layers.0.self_attn.o_proj" + ] + self.assertTrue(marker["_has_pre_quant_scale"]) + + def test_nvfp4_awq_portable_linear_and_rowwise_embedding(self): + config = ComfyNvfp4Config( + { + "proj": { + "format": "nvfp4", + "full_precision_matrix_mult": True, + "_has_pre_quant_scale": True, + } + } + ) + method = ComfyFullPrecisionNvfp4LinearMethod(config, has_pre_quant_scale=True) + layer = nn.Module() + layer.weight = nn.Parameter( + torch.full((128, 32), 0x21, dtype=torch.uint8), requires_grad=False + ) + layer.weight_scale = nn.Parameter( + torch.ones((128, 4), dtype=torch.float8_e4m3fn), requires_grad=False + ) + layer.weight_scale_2 = nn.Parameter(torch.tensor(0.5), requires_grad=False) + layer.pre_quant_scale = nn.Parameter( + torch.full((64,), 2.0), requires_grad=False + ) + inputs = torch.zeros(1, 64) + inputs[:, 0::2] = 1 + + output = method.apply(layer, inputs) + + torch.testing.assert_close(output, torch.full((1, 128), 32.0)) + + embedding_method = ComfyRowwiseInt8EmbeddingMethod() + embedding = nn.Module() + embedding_method.create_weights( + embedding, + input_size_per_partition=2, + output_partition_sizes=[3], + input_size=2, + output_size=3, + params_dtype=torch.bfloat16, + ) + embedding.weight.data.copy_(torch.tensor([[1, 2], [3, 4], [5, 6]])) + embedding.weight_scale.data.copy_(torch.tensor([[0.5], [1.0], [2.0]])) + rows = embedding_method.embedding(embedding, torch.tensor([2, 0])) + torch.testing.assert_close( + rows, + torch.tensor([[10.0, 12.0], [0.5, 1.0]], dtype=torch.bfloat16), + ) + def test_gguf_maps_h3_names_and_drops_unused_language_layers(self): self.get_quant_config.return_value = None diff --git a/python/sglang/srt/layers/quantization/dequantization.py b/python/sglang/srt/layers/quantization/dequantization.py index 5c7e7b34d..6b335111c 100644 --- a/python/sglang/srt/layers/quantization/dequantization.py +++ b/python/sglang/srt/layers/quantization/dequantization.py @@ -74,9 +74,11 @@ def dequantize_nvfp4( w_s: torch.Tensor, w_s2: Optional[torch.Tensor], out_dtype: torch.dtype = torch.bfloat16, + high_nibble_first: bool = False, ) -> torch.Tensor: """NVFP4 -> ``out_dtype``. ``w_q``: uint8 [..., out, in/2] packed e2m1 - (low nibble = even idx). ``w_s``: fp8 e4m3 [..., out, in/16] per-block. + (low nibble = even idx unless ``high_nibble_first``). ``w_s``: fp8 e4m3 + [..., out, in/16] per-block. ``w_s2``: optional fp32 per-tensor scalar that multiplies the per-block scale (ModelOpt / AMD Quark NVFP4).""" device = w_q.device @@ -85,10 +87,11 @@ def dequantize_nvfp4( low = (w_q & 0xF).to(torch.int64) high = (w_q >> 4).to(torch.int64) + first, second = (high, low) if high_nibble_first else (low, high) lut = _FP4_E2M1_LUT.to(device=device, dtype=torch.float32) deq = torch.empty(*batch, out_dim, in_dim, dtype=torch.float32, device=device) - deq[..., 0::2] = lut[low] - deq[..., 1::2] = lut[high] + deq[..., 0::2] = lut[first] + deq[..., 1::2] = lut[second] scale = w_s.to(torch.float32) if w_s2 is not None: