diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index 3c5cefaa8..fb9fdfe4f 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -116,8 +116,10 @@ may be either a normal style adapter or a timestep-distilled Turbo adapter. | DiT | [NVFP4, optionally mixed with INT8 or FP8](https://huggingface.co/Abiray/Minimax-H3-nvfp4-INT4-INT8-Convrot) | `--component-weights-paths.transformer OWNER/REPO/path/FILE.safetensors` | Auto-detected; NVFP4 execution requires NVIDIA compute capability 10.0+. | | DiT | [AutoRound W4A16 component](https://huggingface.co/Ar4ikov/MiniMax-H3-transformer-W4A16-RTN) | `--component-paths.transformer Ar4ikov/MiniMax-H3-transformer-W4A16-RTN` | Self-describing Diffusers component; SGLang reuses the SRT GPTQ/Marlin backend. The linked export is FL2VA. | | DiT | GGUF, full or AdaLN-pruned ([full](https://huggingface.co/leejet/MiniMax-H3-GGUF), [pruned](https://huggingface.co/unsloth/MiniMax-H3-GGUF)) | `--component-weights-paths.transformer OWNER/REPO/FILE.gguf` | CUDA capacity path; aligned TP and layerwise offload are supported, FSDP and LoRA are not. | +| Text encoder | Architecture-compatible Qwen3-VL BF16 finetune ([Heretic example](https://huggingface.co/llmfan46/Qwen3-VL-32B-Instruct-ultra-uncensored-heretic)) | `--component-paths.text_encoder OWNER/REPO` | Reuses the native H3 extractor, including its vision tower and layer-50 hidden-state selection. Finetuning changes conditioning, not the sampling schedule. | | Text encoder | [Serialized FP8 component](https://huggingface.co/Qwen/Qwen3-VL-32B-Instruct-FP8) | `--component-paths.text_encoder Qwen/Qwen3-VL-32B-Instruct-FP8` | Only eligible language-model linears use FP8; embeddings, norms, and the vision tower keep their declared precision. | -| Text encoder | ConvRot INT8, W4A8, or W4A4 safetensors ([INT8](https://huggingface.co/Comfy-Org/MiniMax-H3/tree/main/text_encoders), [W4A8](https://huggingface.co/Winnougan/MiniMax-H3-INT4_Convrot_ComfyUI), [W4A4](https://huggingface.co/Merserk/MiniMax-H3-INT4-ConvRot)) | `--component-weights-paths.text_encoder OWNER/REPO/path/FILE.safetensors` | Auto-detected; requires `comfy-kitchen`. Unmarked vision and embedding tensors keep their declared precision. | +| Text encoder | ConvRot INT8, W4A8, or W4A4 safetensors ([INT8](https://huggingface.co/Comfy-Org/MiniMax-H3/tree/main/text_encoders), [Heretic INT8](https://huggingface.co/ethanfel/Qwen3-VL-32B-Ultra-Heretic-H3-ComfyUI-INT8-ConvRot), [W4A8](https://huggingface.co/Winnougan/MiniMax-H3-INT4_Convrot_ComfyUI), [W4A4](https://huggingface.co/Merserk/MiniMax-H3-INT4-ConvRot)) | `--component-weights-paths.text_encoder OWNER/REPO/path/FILE.safetensors` | Auto-detected; requires `comfy-kitchen`. INT8/W4A8 files may include a scalar-scale INT8 embedding. Unmarked tensors retain their declared precision. | +| Text encoder | Comfy NVFP4 with dynamic activation quantization ([Heretic example](https://huggingface.co/sakamakismile/Qwen3-VL-32B-Heretic-MiniMax-H3-NVFP4)) | `--component-weights-paths.text_encoder OWNER/REPO/path/FILE.safetensors` | Requires `comfy-kitchen` and NVIDIA compute capability 10.0+. Per-layer metadata selects native NVFP4 matmul; scalar/row-wise INT8 embeddings remain packed. | | Text encoder | [NVFP4-AWQ](https://huggingface.co/Comfy-Org/MiniMax-H3/tree/main/text_encoders) or [Quanto qint8](https://huggingface.co/DeepBeepMeep/MiniMax-H3/tree/main/Qwen3-VL-32B-Instruct) safetensors | `--component-weights-paths.text_encoder OWNER/REPO/path/FILE.safetensors` | Memory-oriented formats: compressed storage is restored, then each active matrix uses BF16/FP16 compute. | | Text encoder | [GGUF Qwen3-VL](https://huggingface.co/DeepBeepMeep/MiniMax-H3/tree/main/Qwen3-VL-32B-Instruct) | `--component-weights-paths.text_encoder OWNER/REPO/FILE.gguf` | CUDA capacity path with encoder TP/layerwise support; encoder FSDP is not supported. | | Text encoder | [Compact Qwen3-VL 4B/8B + ClipProj](https://huggingface.co/NicoLab28/ClipProj-MiniMax-H3) | `--component-paths.text_encoder ENCODER_REPO --component-paths.conditioning_projection PROJECTION.safetensors` | Approximate conditioning replacement. A separate weight-only override may quantize the selected small encoder. | @@ -128,6 +130,11 @@ The rows compose rather than enumerate every cross-product. A Turbo-merged INT8 ConvRot checkpoint, for example, must satisfy both the Turbo sampling contract and the ConvRot storage/backend contract. +H3's Qwen3-VL text encoder also handles image understanding; there is no separate +`image_encoder` component. Select the conditioning checkpoint, not an optional +generation-tail file for prompt rewriting. The `uncensored` or `Heretic` label +describes a weight modification, not a separate loader or a guarantee of output quality. + For H3, the registered component names are `transformer`, `text_encoder`, `video_vae`, and `audio_vae`. The shorter `--transformer-weights-path` and `--text-encoder-path` aliases remain supported. `conditioning_projection` is an diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index 18116f4cb..4ada585a9 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -210,12 +210,12 @@ backend. 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 + comfy-nvfp4 + Safetensors with serialized nvfp4 and optional scalar/row-wise int8_tensorwise layer metadata --component-weights-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. + comfy-kitchen for NVFP4 matmul + Auto-detected; omit explicit quantization. Preserves high-nibble-first weights, swizzled block scales, and optional AWQ input pre-scales. Layers declaring full_precision_matrix_mult use BF16/FP16 compute with packed storage. Other NVFP4 layers use dynamic activation quantization and NVFP4 matmul on NVIDIA SM100+. Companion INT8 embeddings support scalar or per-row scales without expanding the full table. quanto-int8 diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_int8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_int8.py new file mode 100644 index 000000000..ead2e080c --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_int8.py @@ -0,0 +1,94 @@ +# SPDX-License-Identifier: Apache-2.0 +"""INT8 embedding lookup shared by serialized Comfy quantization formats.""" + +from typing import Any + +import torch +import torch.nn.functional as F +from torch import nn + +from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( + QuantizeMethodBase, +) +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs + +try: + from comfy_kitchen import registry + + _cuda_embedding = ( + torch.ops.comfy_kitchen.dequantize_int8_embedding + if registry.is_available("cuda") + else None + ) +except (ImportError, AttributeError): + _cuda_embedding = None + +_OUTPUT_DTYPE_CODE = {torch.float32: 0, torch.float16: 1, torch.bfloat16: 2} + + +def is_comfy_int8_embedding(marker: dict[str, Any] | None) -> bool: + return bool( + marker is not None + and marker.get("format") == "int8_tensorwise" + and not marker.get("convrot", False) + and (marker.get("_is_rowwise") or marker.get("_is_tensorwise_scalar")) + ) + + +class ComfyInt8EmbeddingMethod(QuantizeMethodBase): + """Keep the table packed and dequantize only the requested rows.""" + + def __init__(self, *, tensorwise: bool = False) -> None: + self.tensorwise = tensorwise + + 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: + self.output_dtype = params_dtype + rows = sum(output_partition_sizes) + for name, shape, dtype, dims in ( + ( + "weight", + (rows, input_size_per_partition), + torch.int8, + {"input_dim": 1, "output_dim": 0}, + ), + ( + "weight_scale", + () if self.tensorwise else (rows, 1), + torch.float32, + {} if self.tensorwise else {"output_dim": 0}, + ), + ): + parameter = nn.Parameter( + torch.empty(shape, dtype=dtype), requires_grad=False + ) + set_weight_attrs(parameter, extra_weight_attrs) + set_weight_attrs(parameter, dims) + layer.register_parameter(name, parameter) + + def apply(self, layer: nn.Module, x: torch.Tensor, bias=None) -> torch.Tensor: + raise NotImplementedError("Comfy INT8 embeddings support lookup only") + + def embedding(self, layer: nn.Module, input_: torch.Tensor) -> torch.Tensor: + if self.tensorwise and layer.weight.is_cuda and _cuda_embedding is not None: + return _cuda_embedding( + layer.weight, + layer.weight_scale, + input_, + 0, + _OUTPUT_DTYPE_CODE[self.output_dtype], + ) + weight = F.embedding(input_, layer.weight) + if self.tensorwise: + # scalar-scale exports multiply in FP32 before rounding to the activation dtype + return (weight.float() * layer.weight_scale).to(self.output_dtype) + scale = F.embedding(input_, layer.weight_scale).to(self.output_dtype) + return weight.to(self.output_dtype) * scale diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py index 6c8fcd0ca..2948253da 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/comfy_nvfp4.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Portable full-precision execution for Comfy NVFP4 checkpoints.""" +"""Native execution of serialized Comfy NVFP4 checkpoints.""" from __future__ import annotations @@ -13,6 +13,10 @@ from sglang.multimodal_gen.runtime.layers.linear import ( LinearBase, UnquantizedLinearMethod, ) +from sglang.multimodal_gen.runtime.layers.quantization.comfy_int8 import ( + ComfyInt8EmbeddingMethod, + is_comfy_int8_embedding, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( QuantizeMethodBase, ) @@ -24,9 +28,15 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) +from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs from sglang.srt.layers.quantization.dequantization import dequantize_nvfp4 +try: + from comfy_kitchen import quantize_nvfp4, scaled_mm_nvfp4 +except ImportError: + quantize_nvfp4 = scaled_mm_nvfp4 = None + def _register_parameter( layer: nn.Module, @@ -42,55 +52,6 @@ def _register_parameter( 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.""" @@ -162,8 +123,51 @@ class ComfyFullPrecisionNvfp4LinearMethod(ModelOptFp4LinearMethod): return F.linear(x, weight, bias) +class ComfyNvfp4LinearMethod(ComfyFullPrecisionNvfp4LinearMethod): + """Execute serialized Comfy NVFP4 with dynamic activation quantization.""" + + def __init__(self, quant_config: ComfyNvfp4Config, *, has_pre_quant_scale: bool): + super().__init__(quant_config, has_pre_quant_scale=has_pre_quant_scale) + capability = current_platform.get_device_capability() + if ( + not current_platform.is_cuda() + or capability is None + or capability.to_int() < 100 + ): + raise ValueError( + "Comfy NVFP4 matmul requires NVIDIA compute capability 10.0+" + ) + if quantize_nvfp4 is None or scaled_mm_nvfp4 is None: + raise ImportError("Comfy NVFP4 matmul requires comfy-kitchen") + + def apply(self, layer: nn.Module, x: torch.Tensor, bias=None) -> torch.Tensor: + shape = x.shape + x = x.reshape(-1, shape[-1]) + if self.has_pre_quant_scale: + x = x * layer.pre_quant_scale + scale = ( + (x.abs().amax() / (448 * 6)) + .float() + .clamp_min(torch.finfo(torch.float32).tiny) + ) + packed, block_scale = quantize_nvfp4(x.contiguous(), scale, pad_16x=True) + output = scaled_mm_nvfp4( + packed, + layer.weight, + tensor_scale_a=scale, + tensor_scale_b=layer.weight_scale_2, + block_scale_a=block_scale, + block_scale_b=layer.weight_scale, + bias=bias, + out_dtype=x.dtype, + ) + return output[: x.shape[0], : layer.output_size_per_partition].reshape( + *shape[:-1], layer.output_size_per_partition + ) + + class ComfyNvfp4Config(ModelOptFp4Config): - """Dispatch full-precision Comfy NVFP4 linears and their INT8 embedding.""" + """Honor each NVFP4 layer's matmul policy and its INT8 embedding companion.""" checkpoint_uses_comfy_quantization = True @@ -178,18 +182,15 @@ class ComfyNvfp4Config(ModelOptFp4Config): 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"): + if is_comfy_int8_embedding(marker): 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" - ) + if marker.get("convrot", False): + raise ValueError(f"Rotated NVFP4 weights are not supported: {prefix!r}") @classmethod def get_name(cls) -> str: @@ -221,14 +222,14 @@ class ComfyNvfp4Config(ModelOptFp4Config): if isinstance(layer, VocabParallelEmbedding): if marker is None: return None - if marker.get("format") != "int8_tensorwise" or not marker.get( - "_is_rowwise" - ): + if not is_comfy_int8_embedding(marker): raise ValueError( f"Unsupported quantized embedding marker for {prefix!r}: {marker}" ) self.selected.append(prefix) - return ComfyRowwiseInt8EmbeddingMethod() + return ComfyInt8EmbeddingMethod( + tensorwise=bool(marker.get("_is_tensorwise_scalar")) + ) if not isinstance(layer, LinearBase): return None if marker is None: @@ -236,22 +237,22 @@ class ComfyNvfp4Config(ModelOptFp4Config): if marker.get("format") != "nvfp4": raise ValueError(f"Unsupported quantized linear marker for {prefix!r}") self.selected.append(prefix) - return ComfyFullPrecisionNvfp4LinearMethod( + method = ( + ComfyFullPrecisionNvfp4LinearMethod + if marker.get("full_precision_matrix_mult", False) + else ComfyNvfp4LinearMethod + ) + return method( 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") - ) + return is_comfy_int8_embedding(self.layer_markers.get(prefix)) __all__ = [ "ComfyFullPrecisionNvfp4LinearMethod", "ComfyNvfp4Config", - "ComfyRowwiseInt8EmbeddingMethod", + "ComfyNvfp4LinearMethod", ] diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_int8_config.py b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_int8_config.py index 4a515d94f..8f86297eb 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_int8_config.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/configs/kitchen_int8_config.py @@ -8,10 +8,17 @@ from typing import Any import torch from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod +from sglang.multimodal_gen.runtime.layers.quantization.comfy_int8 import ( + ComfyInt8EmbeddingMethod, + is_comfy_int8_embedding, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( QuantizationConfig, QuantizeMethodBase, ) +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.srt.layers.quantization.utils import is_layer_skipped @@ -45,6 +52,8 @@ class KitchenInt8Config(QuantizationConfig): self._serialized_group_sizes: dict[str, int] = {} if layer_markers is not None: for prefix, marker in layer_markers.items(): + if is_comfy_int8_embedding(marker): + continue if marker.get("format") != "int8_tensorwise": raise ValueError( f"Unsupported Comfy INT8 format for {prefix!r}: " @@ -102,6 +111,13 @@ class KitchenInt8Config(QuantizationConfig): KitchenInt8LinearMethod, ) + if isinstance(layer, VocabParallelEmbedding) and self.quantizes_embedding( + prefix + ): + self.selected.append(prefix) + return ComfyInt8EmbeddingMethod( + tensorwise=bool(self.layer_markers[prefix].get("_is_tensorwise_scalar")) + ) if not isinstance(layer, LinearBase): return None if self.layer_markers is not None: @@ -166,3 +182,8 @@ class KitchenInt8Config(QuantizationConfig): return True group_size = marker_group_size return input_size_per_partition % group_size == 0 + + def quantizes_embedding(self, prefix: str) -> bool: + return self.layer_markers is not None and is_comfy_int8_embedding( + self.layer_markers.get(prefix) + ) 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 f80d538ba..ed6dc4172 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 @@ -11,12 +11,15 @@ from sglang.multimodal_gen.runtime.layers.linear import ( LinearBase, UnquantizedLinearMethod, ) +from sglang.multimodal_gen.runtime.layers.quantization.comfy_int8 import ( + ComfyInt8EmbeddingMethod, + is_comfy_int8_embedding, +) from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( QuantizationConfig, QuantizeMethodBase, ) from sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8 import ( - KitchenInt8EmbeddingMethod, KitchenW4A8LinearMethod, ) from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( @@ -49,9 +52,7 @@ class KitchenW4A8Config(QuantizationConfig): for prefix, marker in layer_markers.items(): marker_format = marker.get("format") - if marker_format == "int8_tensorwise" and marker.get( - "_is_tensorwise_scalar" - ): + if is_comfy_int8_embedding(marker): continue if marker_format != "asym_w4a8_int8": raise ValueError( @@ -92,14 +93,14 @@ class KitchenW4A8Config(QuantizationConfig): if isinstance(layer, VocabParallelEmbedding): if marker is None: return None - if marker.get("format") != "int8_tensorwise" or not marker.get( - "_is_tensorwise_scalar" - ): + if not is_comfy_int8_embedding(marker): raise ValueError( f"Unsupported quantized embedding marker for {prefix!r}: {marker}" ) self.selected.append(prefix) - return KitchenInt8EmbeddingMethod() + return ComfyInt8EmbeddingMethod( + tensorwise=bool(marker.get("_is_tensorwise_scalar")) + ) if not isinstance(layer, LinearBase): return None if marker is None: @@ -153,9 +154,4 @@ class KitchenW4A8Config(QuantizationConfig): 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") - ) + return is_comfy_int8_embedding(self.layer_markers.get(prefix)) 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 2a47826cb..3cb3f5c0e 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/kitchen_w4a8.py @@ -14,8 +14,6 @@ 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, @@ -148,62 +146,4 @@ class KitchenW4A8LinearMethod(LinearMethodBase): return output -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"] +__all__ = ["KitchenW4A8LinearMethod"] diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index d0813d8ea..29d79d49b 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -103,6 +103,16 @@ def process_model_weights_after_loading( return processed_layers +def _merge_comfy_quant_marker( + markers: dict[str, dict[str, Any]], prefix: str, marker: dict[str, Any] +) -> None: + # Headers may summarize a layer whose tensor marker includes more fields. + previous = markers.setdefault(prefix, {}) + if any(key in previous and previous[key] != value for key, value in marker.items()): + raise ValueError(f"Conflicting Comfy quantization markers for {prefix!r}") + previous.update(marker) + + def inspect_comfy_quant_markers( safetensors_list: list[str], param_name_mapper: Callable[[str], str] | None = None, @@ -140,12 +150,7 @@ def inspect_comfy_quant_markers( 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 + _merge_comfy_quant_marker(raw_markers, prefix, marker) for key in checkpoint.keys(): tensor_slice = checkpoint.get_slice(key) checkpoint_meta[key] = ( @@ -171,12 +176,7 @@ def inspect_comfy_quant_markers( f"Comfy quantization marker {key!r} must contain a JSON object" ) prefix = key.removesuffix(".comfy_quant") - 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 + _merge_comfy_quant_marker(raw_markers, prefix, marker) if global_quant_formats == {"mxfp8"}: for prefix in marked_dtype_weight_prefixes: 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 36ea9f522..b4d784f6a 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,10 +13,13 @@ from torch import nn from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config from sglang.multimodal_gen.runtime.layers.linear import LinearBase +from sglang.multimodal_gen.runtime.layers.quantization.comfy_int8 import ( + ComfyInt8EmbeddingMethod, +) from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import ( ComfyFullPrecisionNvfp4LinearMethod, ComfyNvfp4Config, - ComfyRowwiseInt8EmbeddingMethod, + ComfyNvfp4LinearMethod, ) from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import ( KitchenInt8Config, @@ -29,6 +32,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_w4a8_conf ) from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config from sglang.multimodal_gen.runtime.layers.quantization.gguf import GGUFConfig +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, NativeComponentLoaderRequired, @@ -55,11 +61,169 @@ from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import ( from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLTextModel from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.quantization_utils import ( + inspect_comfy_quant_markers, process_model_weights_after_loading, ) from sglang.srt.layers.linear import LinearBase as SrtLinearBase +@pytest.mark.parametrize("backend", ["int8", "nvfp4"]) +@pytest.mark.parametrize("tensorwise", [False, True]) +@pytest.mark.parametrize("tp_size", [1, 2]) +def test_comfy_embedding_checkpoint_lookup(tmp_path, backend, tensorwise, tp_size): + embed = "model.embed_tokens" + linear = "model.layers.0.self_attn.q_proj" + weights = ( + torch.arange(7 * 256).reshape(7, 256).remainder(251).sub(125).to(torch.int8) + ) + scale = ( + torch.tensor(0.0137) + if tensorwise + else torch.arange(1, 8).float().reshape(7, 1) / 113 + ) + marker = ( + {"format": "nvfp4", "full_precision_matrix_mult": True} + if backend == "nvfp4" + else {"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": 256} + ) + tensors = { + f"{embed}.weight": weights, + f"{embed}.weight_scale": scale, + f"{linear}.weight": torch.ones( + (128, 128 if backend == "nvfp4" else 256), + dtype=torch.uint8 if backend == "nvfp4" else torch.int8, + ), + f"{linear}.weight_scale": torch.ones( + (128, 16 if backend == "nvfp4" else 1), + dtype=torch.float8_e4m3fn if backend == "nvfp4" else torch.float32, + ), + } + if backend == "nvfp4": + tensors[f"{linear}.weight_scale_2"] = torch.tensor(0.5) + else: + tensors[f"{linear}.comfy_quant"] = torch.tensor( + list(json.dumps({**marker, "per_row": True}).encode()), dtype=torch.uint8 + ) + checkpoint = tmp_path / "encoder.safetensors" + save_file( + tensors, + checkpoint, + metadata={ + "_quantization_metadata": json.dumps( + {"layers": {embed: {"format": "int8_tensorwise"}, linear: marker}} + ) + }, + ) + config = _get_encoder_quant_config( + {}, str(tmp_path), str(checkpoint), MiniMaxH3Qwen3VLEncoder + ) + prefix = "model.language_model.embed_tokens" + assert config.quantizes_embedding(prefix) + for rank in range(tp_size): + embedding = VocabParallelEmbedding( + 7, + 256, + params_dtype=torch.bfloat16, + padding_size=8, + quant_config=config, + prefix=prefix, + tp_group=SimpleNamespace(world_size=tp_size, rank_in_group=rank), + ) + embedding.weight_loader(embedding.weight, weights) + embedding.weight_loader(embedding.weight_scale, scale) + start = embedding.shard_indices.org_vocab_start_index + count = embedding.shard_indices.org_vocab_end_index - start + indices = torch.arange(count) + actual = embedding.quant_method.embedding(embedding, indices) + if tensorwise: + expected = (weights[start : start + count].float() * scale).bfloat16() + assert embedding.weight_scale.shape == () + else: + expected = ( + weights[start : start + count].bfloat16() + * scale[start : start + count].bfloat16() + ) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + assert embedding.weight.dtype == torch.int8 + assert torch.count_nonzero(embedding.weight[count:]) == 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_comfy_nvfp4_dynamic_matmul(dtype): + if torch.cuda.get_device_capability()[0] < 10: + pytest.skip("requires NVFP4 tensor cores") + ck = pytest.importorskip("comfy_kitchen") + layout = pytest.importorskip("comfy_kitchen.tensor.nvfp4").TensorCoreNVFP4Layout + torch.manual_seed(42) + x = torch.randn(33, 256, device="cuda", dtype=dtype) + w = torch.randn(128, 256, device="cuda", dtype=dtype) + weight_scale = w.abs().amax().float() / (448 * 6) + packed, scales = ck.quantize_nvfp4(w, weight_scale) + layer = nn.Module() + layer.weight = nn.Parameter(packed, requires_grad=False) + layer.weight_scale = nn.Parameter(scales, requires_grad=False) + layer.weight_scale_2 = nn.Parameter(weight_scale, requires_grad=False) + layer.output_size_per_partition = 128 + config = ComfyNvfp4Config({"proj": {"format": "nvfp4"}}) + method = ComfyNvfp4LinearMethod(config, has_pre_quant_scale=False) + actual = method.apply(layer, x) + x_packed, x_params = layout.quantize(x) + expected = ck.scaled_mm_nvfp4( + x_packed, + packed, + tensor_scale_a=x_params.scale, + tensor_scale_b=weight_scale, + block_scale_a=x_params.block_scale, + block_scale_b=scales, + out_dtype=x.dtype, + )[:33] + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + assert torch.isfinite(actual).all() + zero = method.apply(layer, torch.zeros_like(x)) + assert torch.isfinite(zero).all() + assert torch.count_nonzero(zero) == 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_comfy_scalar_embedding_matches_kitchen_kernel(): + pytest.importorskip("comfy_kitchen") + method = ComfyInt8EmbeddingMethod(tensorwise=True) + layer = nn.Module() + method.create_weights(layer, 256, [128], 256, 128, torch.bfloat16) + layer.to("cuda") + layer.weight.data.copy_( + torch.arange(128 * 256, device="cuda") + .reshape(128, 256) + .remainder(251) + .sub(125) + .to(torch.int8) + ) + layer.weight_scale.data.fill_(0.0137) + indices = torch.tensor([0, 63, 127, 0], device="cuda") + expected = (layer.weight[indices].float() * layer.weight_scale).bfloat16() + torch.testing.assert_close( + method.embedding(layer, indices), expected, rtol=0, atol=0 + ) + + +def test_comfy_marker_conflicting_values_rejected(tmp_path): + checkpoint = tmp_path / "encoder.safetensors" + marker = {"format": "int8_tensorwise", "convrot": True} + save_file( + { + "layer.comfy_quant": torch.tensor( + list(json.dumps({**marker, "convrot": False}).encode()), + dtype=torch.uint8, + ) + }, + checkpoint, + metadata={"_quantization_metadata": json.dumps({"layers": {"layer": marker}})}, + ) + with pytest.raises(ValueError, match="Conflicting Comfy quantization markers"): + inspect_comfy_quant_markers([str(checkpoint)]) + + @pytest.mark.parametrize("missing", [False, True]) @pytest.mark.parametrize("competing_index", [False, True]) def test_native_encoder_restoration_checks_checkpoint_before_fallback( @@ -706,7 +870,7 @@ class TestTextEncoderQuantization(unittest.TestCase): torch.testing.assert_close(output, torch.full((1, 128), 32.0)) - embedding_method = ComfyRowwiseInt8EmbeddingMethod() + embedding_method = ComfyInt8EmbeddingMethod() embedding = nn.Module() embedding_method.create_weights( embedding,