[diffusion] feat: support mixed INT8 embeddings and Comfy NVFP4 encoders for minimax-h3 (#38506)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
@@ -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 | [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 | [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. |
|
| 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 | [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 | [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 | [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. |
|
| 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
|
ConvRot checkpoint, for example, must satisfy both the Turbo sampling contract
|
||||||
and the ConvRot storage/backend 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`,
|
For H3, the registered component names are `transformer`, `text_encoder`,
|
||||||
`video_vae`, and `audio_vae`. The shorter `--transformer-weights-path` and
|
`video_vae`, and `audio_vae`. The shorter `--transformer-weights-path` and
|
||||||
`--text-encoder-path` aliases remain supported. `conditioning_projection` is an
|
`--text-encoder-path` aliases remain supported. `conditioning_projection` is an
|
||||||
|
|||||||
@@ -210,12 +210,12 @@ backend.
|
|||||||
<td>Auto-detected; omit <code>--quantization</code>. 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.</td>
|
<td>Auto-detected; omit <code>--quantization</code>. 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.</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>comfy-nvfp4-full-precision</code></td>
|
<td><code>comfy-nvfp4</code></td>
|
||||||
<td>Safetensors with serialized <code>nvfp4</code> and optional row-wise <code>int8_tensorwise</code> layer metadata</td>
|
<td>Safetensors with serialized <code>nvfp4</code> and optional scalar/row-wise <code>int8_tensorwise</code> layer metadata</td>
|
||||||
<td><code>--component-weights-paths.text_encoder</code></td>
|
<td><code>--component-weights-paths.text_encoder</code></td>
|
||||||
<td>MiniMax-H3 native Qwen3-VL encoder</td>
|
<td>MiniMax-H3 native Qwen3-VL encoder</td>
|
||||||
<td>None</td>
|
<td><code>comfy-kitchen</code> for NVFP4 matmul</td>
|
||||||
<td>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.</td>
|
<td>Auto-detected; omit explicit quantization. Preserves high-nibble-first weights, swizzled block scales, and optional AWQ input pre-scales. Layers declaring <code>full_precision_matrix_mult</code> 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.</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>quanto-int8</code></td>
|
<td><code>quanto-int8</code></td>
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# 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
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -13,6 +13,10 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
LinearBase,
|
LinearBase,
|
||||||
UnquantizedLinearMethod,
|
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 (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
QuantizeMethodBase,
|
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 (
|
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs
|
from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs
|
||||||
from sglang.srt.layers.quantization.dequantization import dequantize_nvfp4
|
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(
|
def _register_parameter(
|
||||||
layer: nn.Module,
|
layer: nn.Module,
|
||||||
@@ -42,55 +52,6 @@ def _register_parameter(
|
|||||||
layer.register_parameter(name, 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):
|
class ComfyFullPrecisionNvfp4LinearMethod(ModelOptFp4LinearMethod):
|
||||||
"""Keep NVFP4 storage and dequantize one active Linear for its matmul."""
|
"""Keep NVFP4 storage and dequantize one active Linear for its matmul."""
|
||||||
|
|
||||||
@@ -162,8 +123,51 @@ class ComfyFullPrecisionNvfp4LinearMethod(ModelOptFp4LinearMethod):
|
|||||||
return F.linear(x, weight, bias)
|
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):
|
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
|
checkpoint_uses_comfy_quantization = True
|
||||||
|
|
||||||
@@ -178,18 +182,15 @@ class ComfyNvfp4Config(ModelOptFp4Config):
|
|||||||
self.selected: list[str] = []
|
self.selected: list[str] = []
|
||||||
for prefix, marker in layer_markers.items():
|
for prefix, marker in layer_markers.items():
|
||||||
marker_format = marker.get("format")
|
marker_format = marker.get("format")
|
||||||
if marker_format == "int8_tensorwise" and marker.get("_is_rowwise"):
|
if is_comfy_int8_embedding(marker):
|
||||||
continue
|
continue
|
||||||
if marker_format != "nvfp4":
|
if marker_format != "nvfp4":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported Comfy NVFP4 companion for {prefix!r}: "
|
f"Unsupported Comfy NVFP4 companion for {prefix!r}: "
|
||||||
f"{marker_format!r}"
|
f"{marker_format!r}"
|
||||||
)
|
)
|
||||||
if marker.get("full_precision_matrix_mult") is not True:
|
if marker.get("convrot", False):
|
||||||
raise ValueError(
|
raise ValueError(f"Rotated NVFP4 weights are not supported: {prefix!r}")
|
||||||
f"Comfy NVFP4 layer {prefix!r} must request "
|
|
||||||
"full_precision_matrix_mult"
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_name(cls) -> str:
|
def get_name(cls) -> str:
|
||||||
@@ -221,14 +222,14 @@ class ComfyNvfp4Config(ModelOptFp4Config):
|
|||||||
if isinstance(layer, VocabParallelEmbedding):
|
if isinstance(layer, VocabParallelEmbedding):
|
||||||
if marker is None:
|
if marker is None:
|
||||||
return None
|
return None
|
||||||
if marker.get("format") != "int8_tensorwise" or not marker.get(
|
if not is_comfy_int8_embedding(marker):
|
||||||
"_is_rowwise"
|
|
||||||
):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported quantized embedding marker for {prefix!r}: {marker}"
|
f"Unsupported quantized embedding marker for {prefix!r}: {marker}"
|
||||||
)
|
)
|
||||||
self.selected.append(prefix)
|
self.selected.append(prefix)
|
||||||
return ComfyRowwiseInt8EmbeddingMethod()
|
return ComfyInt8EmbeddingMethod(
|
||||||
|
tensorwise=bool(marker.get("_is_tensorwise_scalar"))
|
||||||
|
)
|
||||||
if not isinstance(layer, LinearBase):
|
if not isinstance(layer, LinearBase):
|
||||||
return None
|
return None
|
||||||
if marker is None:
|
if marker is None:
|
||||||
@@ -236,22 +237,22 @@ class ComfyNvfp4Config(ModelOptFp4Config):
|
|||||||
if marker.get("format") != "nvfp4":
|
if marker.get("format") != "nvfp4":
|
||||||
raise ValueError(f"Unsupported quantized linear marker for {prefix!r}")
|
raise ValueError(f"Unsupported quantized linear marker for {prefix!r}")
|
||||||
self.selected.append(prefix)
|
self.selected.append(prefix)
|
||||||
return ComfyFullPrecisionNvfp4LinearMethod(
|
method = (
|
||||||
|
ComfyFullPrecisionNvfp4LinearMethod
|
||||||
|
if marker.get("full_precision_matrix_mult", False)
|
||||||
|
else ComfyNvfp4LinearMethod
|
||||||
|
)
|
||||||
|
return method(
|
||||||
self,
|
self,
|
||||||
has_pre_quant_scale=bool(marker.get("_has_pre_quant_scale")),
|
has_pre_quant_scale=bool(marker.get("_has_pre_quant_scale")),
|
||||||
)
|
)
|
||||||
|
|
||||||
def quantizes_embedding(self, prefix: str) -> bool:
|
def quantizes_embedding(self, prefix: str) -> bool:
|
||||||
marker = self.layer_markers.get(prefix)
|
return is_comfy_int8_embedding(self.layer_markers.get(prefix))
|
||||||
return bool(
|
|
||||||
marker is not None
|
|
||||||
and marker.get("format") == "int8_tensorwise"
|
|
||||||
and marker.get("_is_rowwise")
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ComfyFullPrecisionNvfp4LinearMethod",
|
"ComfyFullPrecisionNvfp4LinearMethod",
|
||||||
"ComfyNvfp4Config",
|
"ComfyNvfp4Config",
|
||||||
"ComfyRowwiseInt8EmbeddingMethod",
|
"ComfyNvfp4LinearMethod",
|
||||||
]
|
]
|
||||||
|
|||||||
+21
@@ -8,10 +8,17 @@ from typing import Any
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
|
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 (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
QuantizeMethodBase,
|
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.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
||||||
|
|
||||||
@@ -45,6 +52,8 @@ class KitchenInt8Config(QuantizationConfig):
|
|||||||
self._serialized_group_sizes: dict[str, int] = {}
|
self._serialized_group_sizes: dict[str, int] = {}
|
||||||
if layer_markers is not None:
|
if layer_markers is not None:
|
||||||
for prefix, marker in layer_markers.items():
|
for prefix, marker in layer_markers.items():
|
||||||
|
if is_comfy_int8_embedding(marker):
|
||||||
|
continue
|
||||||
if marker.get("format") != "int8_tensorwise":
|
if marker.get("format") != "int8_tensorwise":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported Comfy INT8 format for {prefix!r}: "
|
f"Unsupported Comfy INT8 format for {prefix!r}: "
|
||||||
@@ -102,6 +111,13 @@ class KitchenInt8Config(QuantizationConfig):
|
|||||||
KitchenInt8LinearMethod,
|
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):
|
if not isinstance(layer, LinearBase):
|
||||||
return None
|
return None
|
||||||
if self.layer_markers is not None:
|
if self.layer_markers is not None:
|
||||||
@@ -166,3 +182,8 @@ class KitchenInt8Config(QuantizationConfig):
|
|||||||
return True
|
return True
|
||||||
group_size = marker_group_size
|
group_size = marker_group_size
|
||||||
return input_size_per_partition % group_size == 0
|
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)
|
||||||
|
)
|
||||||
|
|||||||
+10
-14
@@ -11,12 +11,15 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
LinearBase,
|
LinearBase,
|
||||||
UnquantizedLinearMethod,
|
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 (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
QuantizeMethodBase,
|
QuantizeMethodBase,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8 import (
|
from sglang.multimodal_gen.runtime.layers.quantization.kitchen_w4a8 import (
|
||||||
KitchenInt8EmbeddingMethod,
|
|
||||||
KitchenW4A8LinearMethod,
|
KitchenW4A8LinearMethod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
||||||
@@ -49,9 +52,7 @@ class KitchenW4A8Config(QuantizationConfig):
|
|||||||
|
|
||||||
for prefix, marker in layer_markers.items():
|
for prefix, marker in layer_markers.items():
|
||||||
marker_format = marker.get("format")
|
marker_format = marker.get("format")
|
||||||
if marker_format == "int8_tensorwise" and marker.get(
|
if is_comfy_int8_embedding(marker):
|
||||||
"_is_tensorwise_scalar"
|
|
||||||
):
|
|
||||||
continue
|
continue
|
||||||
if marker_format != "asym_w4a8_int8":
|
if marker_format != "asym_w4a8_int8":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -92,14 +93,14 @@ class KitchenW4A8Config(QuantizationConfig):
|
|||||||
if isinstance(layer, VocabParallelEmbedding):
|
if isinstance(layer, VocabParallelEmbedding):
|
||||||
if marker is None:
|
if marker is None:
|
||||||
return None
|
return None
|
||||||
if marker.get("format") != "int8_tensorwise" or not marker.get(
|
if not is_comfy_int8_embedding(marker):
|
||||||
"_is_tensorwise_scalar"
|
|
||||||
):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported quantized embedding marker for {prefix!r}: {marker}"
|
f"Unsupported quantized embedding marker for {prefix!r}: {marker}"
|
||||||
)
|
)
|
||||||
self.selected.append(prefix)
|
self.selected.append(prefix)
|
||||||
return KitchenInt8EmbeddingMethod()
|
return ComfyInt8EmbeddingMethod(
|
||||||
|
tensorwise=bool(marker.get("_is_tensorwise_scalar"))
|
||||||
|
)
|
||||||
if not isinstance(layer, LinearBase):
|
if not isinstance(layer, LinearBase):
|
||||||
return None
|
return None
|
||||||
if marker is None:
|
if marker is None:
|
||||||
@@ -153,9 +154,4 @@ class KitchenW4A8Config(QuantizationConfig):
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
def quantizes_embedding(self, prefix: str) -> bool:
|
def quantizes_embedding(self, prefix: str) -> bool:
|
||||||
marker = self.layer_markers.get(prefix)
|
return is_comfy_int8_embedding(self.layer_markers.get(prefix))
|
||||||
return bool(
|
|
||||||
marker is not None
|
|
||||||
and marker.get("format") == "int8_tensorwise"
|
|
||||||
and marker.get("_is_tensorwise_scalar")
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -14,8 +14,6 @@ try:
|
|||||||
except ImportError: # pragma: no cover - optional dependency
|
except ImportError: # pragma: no cover - optional dependency
|
||||||
w4a8_int8_linear = None
|
w4a8_int8_linear = None
|
||||||
|
|
||||||
_OUTPUT_DTYPE_CODE = {torch.float32: 0, torch.float16: 1, torch.bfloat16: 2}
|
|
||||||
|
|
||||||
|
|
||||||
def _register_weight(
|
def _register_weight(
|
||||||
layer: torch.nn.Module,
|
layer: torch.nn.Module,
|
||||||
@@ -148,62 +146,4 @@ class KitchenW4A8LinearMethod(LinearMethodBase):
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
class KitchenInt8EmbeddingMethod(LinearMethodBase):
|
__all__ = ["KitchenW4A8LinearMethod"]
|
||||||
"""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"]
|
|
||||||
|
|||||||
@@ -103,6 +103,16 @@ def process_model_weights_after_loading(
|
|||||||
return processed_layers
|
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(
|
def inspect_comfy_quant_markers(
|
||||||
safetensors_list: list[str],
|
safetensors_list: list[str],
|
||||||
param_name_mapper: Callable[[str], str] | None = None,
|
param_name_mapper: Callable[[str], str] | None = None,
|
||||||
@@ -140,12 +150,7 @@ def inspect_comfy_quant_markers(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Comfy quantization metadata for {prefix!r} must be an object"
|
f"Comfy quantization metadata for {prefix!r} must be an object"
|
||||||
)
|
)
|
||||||
previous = raw_markers.get(prefix)
|
_merge_comfy_quant_marker(raw_markers, prefix, marker)
|
||||||
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():
|
for key in checkpoint.keys():
|
||||||
tensor_slice = checkpoint.get_slice(key)
|
tensor_slice = checkpoint.get_slice(key)
|
||||||
checkpoint_meta[key] = (
|
checkpoint_meta[key] = (
|
||||||
@@ -171,12 +176,7 @@ def inspect_comfy_quant_markers(
|
|||||||
f"Comfy quantization marker {key!r} must contain a JSON object"
|
f"Comfy quantization marker {key!r} must contain a JSON object"
|
||||||
)
|
)
|
||||||
prefix = key.removesuffix(".comfy_quant")
|
prefix = key.removesuffix(".comfy_quant")
|
||||||
previous = raw_markers.get(prefix)
|
_merge_comfy_quant_marker(raw_markers, prefix, marker)
|
||||||
if previous is not None and previous != marker:
|
|
||||||
raise ValueError(
|
|
||||||
f"Conflicting Comfy quantization markers for {prefix!r}"
|
|
||||||
)
|
|
||||||
raw_markers[prefix] = marker
|
|
||||||
|
|
||||||
if global_quant_formats == {"mxfp8"}:
|
if global_quant_formats == {"mxfp8"}:
|
||||||
for prefix in marked_dtype_weight_prefixes:
|
for prefix in marked_dtype_weight_prefixes:
|
||||||
|
|||||||
@@ -13,10 +13,13 @@ from torch import nn
|
|||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config
|
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.linear import LinearBase
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_int8 import (
|
||||||
|
ComfyInt8EmbeddingMethod,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import (
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_nvfp4 import (
|
||||||
ComfyFullPrecisionNvfp4LinearMethod,
|
ComfyFullPrecisionNvfp4LinearMethod,
|
||||||
ComfyNvfp4Config,
|
ComfyNvfp4Config,
|
||||||
ComfyRowwiseInt8EmbeddingMethod,
|
ComfyNvfp4LinearMethod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
|
||||||
KitchenInt8Config,
|
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.fp8 import Fp8Config
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.gguf import GGUFConfig
|
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 (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentCheckpointUnsupportedError,
|
ComponentCheckpointUnsupportedError,
|
||||||
NativeComponentLoaderRequired,
|
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.models.encoders.qwen3vl import Qwen3VLTextModel
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||||
|
inspect_comfy_quant_markers,
|
||||||
process_model_weights_after_loading,
|
process_model_weights_after_loading,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.linear import LinearBase as SrtLinearBase
|
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("missing", [False, True])
|
||||||
@pytest.mark.parametrize("competing_index", [False, True])
|
@pytest.mark.parametrize("competing_index", [False, True])
|
||||||
def test_native_encoder_restoration_checks_checkpoint_before_fallback(
|
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))
|
torch.testing.assert_close(output, torch.full((1, 128), 32.0))
|
||||||
|
|
||||||
embedding_method = ComfyRowwiseInt8EmbeddingMethod()
|
embedding_method = ComfyInt8EmbeddingMethod()
|
||||||
embedding = nn.Module()
|
embedding = nn.Module()
|
||||||
embedding_method.create_weights(
|
embedding_method.create_weights(
|
||||||
embedding,
|
embedding,
|
||||||
|
|||||||
Reference in New Issue
Block a user