[diffusion] feat: support serialized comfy convrot int8 dits (#35994)

This commit is contained in:
Mick
2026-08-23 09:29:59 +08:00
committed by GitHub
parent 7f30d66045
commit a36c0746b9
11 changed files with 253 additions and 71 deletions
@@ -281,11 +281,13 @@ contract and low resident weight memory, but that part is slower than a fully
quantized FP8 GEMM. TP, Ulysses/Ring sequence parallelism, and quantized FP8 GEMM. TP, Ulysses/Ring sequence parallelism, and
component/layerwise offload are supported; FSDP inference is rejected. component/layerwise offload are supported; FSDP inference is rejected.
The `pruned_int8_convrot` files are detected but remain unsupported. They The `pruned_int8_convrot` files use the same override path and are detected
require online regular-Hadamard ConvRot, dynamic INT8 activation quantization, automatically. Install `comfy-kitchen`, replace the FP8 filename above with
and a matching W8A8 GEMM. SGLang fails before loading them instead of silently `minimax_h3_fl2va_pruned_int8_convrot.safetensors`, and still omit
treating their stored INT8 values as ordinary weights. A native ConvRot kernel `--quantization`. SGLang loads their serialized INT8 weights and row scales into
path should be added and benchmarked separately before these files are accepted. the fused ConvRot kernel without requantizing them. TP1/2/4 and sequence or
layerwise offload are supported; TP8 violates the checkpoint's 256-element
ConvRot group boundary, and FSDP is rejected.
### Advanced: precomputed AdaLN cache ### Advanced: precomputed AdaLN cache
@@ -907,7 +909,7 @@ sglang generate \
`--attention-backend sol_attn` with `--attention-backend sol_attn` with
`--attention-backend-config dense_backend=sage_attn,dense_steps=10` and `--attention-backend-config dense_backend=sage_attn,dense_steps=10` and
`--component-attention-backends text_encoder=torch_sdpa,transformer=sol_attn`. `--component-attention-backends text_encoder=torch_sdpa,transformer=sol_attn`.
See [Quantization](/docs/sglang-diffusion/quantization#kitchen-int8-online-quantization) See [Quantization](/docs/sglang-diffusion/quantization#kitchen-int8)
and [Attention Backends](/docs/sglang-diffusion/attention_backends#sage-then-sol-hybrid). and [Attention Backends](/docs/sglang-diffusion/attention_backends#sage-then-sol-hybrid).
<Warning> <Warning>
+21 -5
View File
@@ -50,10 +50,11 @@ directory directly as `--model-path`, but that is a compatibility path. If a
repo contains multiple candidate checkpoints, pass repo contains multiple candidate checkpoints, pass
`--transformer-weights-path` explicitly. `--transformer-weights-path` explicitly.
MiniMax-H3 auto-detects the per-layer metadata in Comfy's MiniMax-H3 is a verified example for Comfy safetensors with per-layer metadata,
`pruned_fp8_scaled` safetensors. Pass one selected FL2VA or Ref2VA file by local including `pruned_fp8_scaled` and serialized ConvRot INT8. Pass one selected
path, `owner/repo/path/file.safetensors`, or direct Hugging Face file URL; do FL2VA or Ref2VA DiT file by local path, `owner/repo/path/file.safetensors`, or
not combine it with `--quantization`. MiniMax-H3 GGUF usage is documented in direct Hugging Face file URL; do not combine it with `--quantization`. Its GGUF
usage is documented in
the [MiniMax-H3 cookbook](/cookbook/diffusion/MiniMax/MiniMax-H3#pre-quantized-gguf-transformer). the [MiniMax-H3 cookbook](/cookbook/diffusion/MiniMax/MiniMax-H3#pre-quantized-gguf-transformer).
## Quantized Component Repositories ## Quantized Component Repositories
@@ -165,6 +166,14 @@ backend.
<td>None</td> <td>None</td>
<td>CUDA; auto-detected; TP, sequence parallelism, and component/layerwise offload are supported, while FSDP is not. Checkpoint-marked <code>fc2</code> layers retain FP8 storage and use compute-dtype matmul.</td> <td>CUDA; auto-detected; TP, sequence parallelism, and component/layerwise offload are supported, while FSDP is not. Checkpoint-marked <code>fc2</code> layers retain FP8 storage and use compute-dtype matmul.</td>
</tr> </tr>
<tr>
<td><code>comfy-int8-convrot</code></td>
<td>One selected safetensors file with per-layer <code>int8_tensorwise</code> and ConvRot metadata</td>
<td><code>--transformer-weights-path</code></td>
<td>MiniMax-H3 native DiT; pruned FL2VA is E2E-verified and Ref2VA has the same validated tensor contract</td>
<td><code>comfy-kitchen</code></td>
<td>CUDA; auto-detected; uses the fused Kitchen INT8 kernel and validates weight/scale layout before model construction. TP requires every row-parallel input shard to preserve the checkpoint's ConvRot group boundary; the H3 256-group checkpoint supports TP1/2/4, not TP8. Offload is supported; FSDP is not.</td>
</tr>
<tr> <tr>
<td><code>qvg-kv</code></td> <td><code>qvg-kv</code></td>
<td>Unquantized model with runtime causal KV-cache compression</td> <td>Unquantized model with runtime causal KV-cache compression</td>
@@ -334,7 +343,14 @@ sglang generate \
``` ```
**Note:** Requires `aiter` package with MXFP4 kernel support **Note:** Requires `aiter` package with MXFP4 kernel support
### Kitchen INT8 Online Quantization ### Kitchen INT8
Serialized Comfy ConvRot INT8 DiTs are selected through
`--transformer-weights-path` and auto-detected from their per-layer markers.
They load INT8 weights and row scales directly; omit `--quantization`.
For a BF16 checkpoint, `--quantization kitchen_int8` instead performs online
quantization after loading:
`kitchen_int8` quantizes DiT linear weights online from the stock BF16 `kitchen_int8` quantizes DiT linear weights online from the stock BF16
checkpoint. Forward uses the fused `comfy_kitchen.int8_linear` op (rotation, checkpoint. Forward uses the fused `comfy_kitchen.int8_linear` op (rotation,
@@ -88,6 +88,8 @@ class ComfyFullPrecisionFp8LinearMethod(LinearMethodBase):
class ComfyFp8Config(QuantizationConfig): class ComfyFp8Config(QuantizationConfig):
"""Dispatch each Linear according to its serialized ``comfy_quant`` marker.""" """Dispatch each Linear according to its serialized ``comfy_quant`` marker."""
checkpoint_uses_native_qkv_layout = True
def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None: def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None:
super().__init__() super().__init__()
self.layer_markers = layer_markers self.layer_markers = layer_markers
@@ -67,6 +67,7 @@ class QuantizationConfig(ABC):
# for quantization frameworks with a separate quantized model provided, e.g. Nunchaku # for quantization frameworks with a separate quantized model provided, e.g. Nunchaku
quantized_model_path: str | None = None quantized_model_path: str | None = None
checkpoint_uses_native_qkv_layout: bool = False
def __init__(self): def __init__(self):
super().__init__() super().__init__()
@@ -1,11 +1,5 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
"""Config for online INT8 ConvRot quantization via comfy_kitchen. """Config for online or serialized INT8 ConvRot via comfy_kitchen."""
A no-arg ``KitchenInt8Config()`` is the only supported form: weights load in
their source dtype and are quantized in ``process_weights_after_loading``.
Registered CLI name: ``kitchen_int8``.
"""
from __future__ import annotations from __future__ import annotations
@@ -27,17 +21,14 @@ _SUPPORTED_GROUP_SIZES = (16, 64, 256)
class KitchenInt8Config(QuantizationConfig): class KitchenInt8Config(QuantizationConfig):
"""Config for online INT8 ConvRot quantization via comfy_kitchen. """Dispatch online quantization or serialized Comfy ConvRot layers."""
A no-arg ``KitchenInt8Config()`` is the only supported form: weights load in
their source dtype and are quantized in ``process_weights_after_loading``.
"""
def __init__( def __init__(
self, self,
group_size: int = 256, group_size: int = 256,
ignored_layers: list[str] | None = None, ignored_layers: list[str] | None = None,
packed_modules_mapping: dict[str, list[str]] | None = None, packed_modules_mapping: dict[str, list[str]] | None = None,
layer_markers: dict[str, dict[str, Any]] | None = None,
) -> None: ) -> None:
super().__init__() super().__init__()
if group_size not in _SUPPORTED_GROUP_SIZES: if group_size not in _SUPPORTED_GROUP_SIZES:
@@ -48,6 +39,30 @@ class KitchenInt8Config(QuantizationConfig):
self.group_size = group_size self.group_size = group_size
self.ignored_layers = ignored_layers or [] self.ignored_layers = ignored_layers or []
self.packed_modules_mapping = packed_modules_mapping or {} self.packed_modules_mapping = packed_modules_mapping or {}
self.layer_markers = layer_markers
self.is_checkpoint_int8_serialized = layer_markers is not None
self.checkpoint_uses_native_qkv_layout = self.is_checkpoint_int8_serialized
self._serialized_group_sizes: dict[str, int] = {}
if layer_markers is not None:
for prefix, marker in layer_markers.items():
if marker.get("format") != "int8_tensorwise":
raise ValueError(
f"Unsupported Comfy INT8 format for {prefix!r}: "
f"{marker.get('format')!r}"
)
if marker.get("convrot") is not True:
raise ValueError(
f"Serialized kitchen_int8 layer {prefix!r} must set "
"convrot=true"
)
marker_group_size = marker.get("convrot_groupsize")
if marker_group_size not in _SUPPORTED_GROUP_SIZES:
raise ValueError(
f"Serialized kitchen_int8 layer {prefix!r} must declare "
f"convrot_groupsize in {_SUPPORTED_GROUP_SIZES}, got "
f"{marker_group_size!r}"
)
self._serialized_group_sizes[prefix] = marker_group_size
# Which layers actually got quantized is worth stating plainly in the # Which layers actually got quantized is worth stating plainly in the
# log: a silent fallback to BF16 looks exactly like a slow kernel. # log: a silent fallback to BF16 looks exactly like a slow kernel.
self.selected: list[str] = [] self.selected: list[str] = []
@@ -89,6 +104,22 @@ class KitchenInt8Config(QuantizationConfig):
if not isinstance(layer, LinearBase): if not isinstance(layer, LinearBase):
return None return None
if self.layer_markers is not None:
marker_group_size = self._serialized_group_sizes.get(prefix)
if marker_group_size is None:
return UnquantizedLinearMethod()
if layer.input_size % marker_group_size:
raise ValueError(
f"Serialized kitchen_int8 layer {prefix!r} has input size "
f"{layer.input_size}, which is not divisible by its "
f"ConvRot group size {marker_group_size}"
)
self.selected.append(prefix)
return KitchenInt8LinearMethod(
self,
group_size=marker_group_size,
is_checkpoint_serialized=True,
)
if is_layer_skipped( if is_layer_skipped(
prefix, self.ignored_layers, fused_mapping=self.packed_modules_mapping prefix, self.ignored_layers, fused_mapping=self.packed_modules_mapping
): ):
@@ -102,7 +133,11 @@ class KitchenInt8Config(QuantizationConfig):
self.skipped.append(f"{prefix}(in={layer.input_size})") self.skipped.append(f"{prefix}(in={layer.input_size})")
return UnquantizedLinearMethod() return UnquantizedLinearMethod()
self.selected.append(prefix) self.selected.append(prefix)
return KitchenInt8LinearMethod(self) return KitchenInt8LinearMethod(
self,
group_size=self.group_size,
is_checkpoint_serialized=False,
)
def note_quantized(self, saved_bytes: int) -> None: def note_quantized(self, saved_bytes: int) -> None:
self._processed += 1 self._processed += 1
@@ -8,10 +8,9 @@ The difference is that it is a single fused op -- it takes a BF16 activation and
does the Hadamard rotation, dynamic per-row activation quantization, IMMA GEMM, does the Hadamard rotation, dynamic per-row activation quantization, IMMA GEMM,
dequantization and bias add without ever materializing the intermediates. dequantization and bias add without ever materializing the intermediates.
Quantization is data-free (group-wise Hadamard rotation + per-output-channel The online path applies data-free group-wise Hadamard rotation and per-output
absmax), so weights are quantized here after loading rather than read from a channel scaling after loading a stock BF16 checkpoint. Compatible serialized
pre-quantized checkpoint. That keeps this usable with the stock BF16 checkpoint Comfy checkpoints instead load their INT8 weights and row scales directly.
and avoids depending on any external file layout.
""" """
from __future__ import annotations from __future__ import annotations
@@ -73,10 +72,18 @@ def _load_comfy_kitchen():
class KitchenInt8LinearMethod(LinearMethodBase): class KitchenInt8LinearMethod(LinearMethodBase):
"""Quantizes BF16 weights to INT8 after load and runs the fused kernel.""" """Loads or creates ConvRot INT8 weights and runs the fused kernel."""
def __init__(self, quant_config: KitchenInt8Config) -> None: def __init__(
self,
quant_config: KitchenInt8Config,
*,
group_size: int,
is_checkpoint_serialized: bool,
) -> None:
self.quant_config = quant_config self.quant_config = quant_config
self.group_size = group_size
self.is_checkpoint_serialized = is_checkpoint_serialized
_load_comfy_kitchen() _load_comfy_kitchen()
def create_weights( def create_weights(
@@ -92,36 +99,47 @@ class KitchenInt8LinearMethod(LinearMethodBase):
# get_quant_method already screened the unsharded input size, so this # get_quant_method already screened the unsharded input size, so this
# only fires under TP > 1, where a row-parallel layer splits the very # only fires under TP > 1, where a row-parallel layer splits the very
# dimension the rotation groups over. # dimension the rotation groups over.
if input_size_per_partition % self.quant_config.group_size: if input_size_per_partition % self.group_size:
raise ValueError( raise ValueError(
f"kitchen_int8 needs input_size_per_partition " f"kitchen_int8 needs input_size_per_partition "
f"({input_size_per_partition}) divisible by group_size " f"({input_size_per_partition}) divisible by group_size "
f"{self.quant_config.group_size}" f"{self.group_size}"
) )
# Deliberately identical to UnquantizedLinearMethod: weights load as # The online path initially matches UnquantizedLinearMethod so the
# BF16 through the model's existing loaders (H3 for instance installs a # source weights load in BF16 before quantization. Serialized weights
# custom qkv loader that reorders the grouped checkpoint layout), and # allocate their final INT8 storage immediately.
# only then get replaced by their quantized form.
weight = Parameter( weight = Parameter(
torch.empty( torch.empty(
sum(output_partition_sizes), sum(output_partition_sizes),
input_size_per_partition, input_size_per_partition,
dtype=params_dtype, dtype=(torch.int8 if self.is_checkpoint_serialized else params_dtype),
), ),
requires_grad=False, requires_grad=False,
) )
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0}) set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
layer.register_parameter("weight", weight) layer.register_parameter("weight", weight)
set_weight_attrs(weight, extra_weight_attrs) set_weight_attrs(weight, extra_weight_attrs)
if self.is_checkpoint_serialized:
weight_scale = Parameter(
torch.empty(
sum(output_partition_sizes),
1,
dtype=torch.float32,
),
requires_grad=False,
)
set_weight_attrs(weight_scale, {"output_dim": 0})
set_weight_attrs(weight_scale, extra_weight_attrs)
layer.register_parameter("weight_scale", weight_scale)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None: def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
from comfy_kitchen.tensor.int8 import TensorWiseINT8Layout
weight = layer.weight.data weight = layer.weight.data
if weight.dtype == torch.int8: # already processed if self.is_checkpoint_serialized or weight.dtype == torch.int8:
return return
from comfy_kitchen.tensor.int8 import TensorWiseINT8Layout
# Quantization runs on CUDA, but the model may still be staged on CPU # Quantization runs on CUDA, but the model may still be staged on CPU
# for offload. Round-trip one layer at a time rather than relying on # for offload. Round-trip one layer at a time rather than relying on
# the loader's whole-model device move, which would not fit in VRAM. # the loader's whole-model device move, which would not fit in VRAM.
@@ -131,7 +149,7 @@ class KitchenInt8LinearMethod(LinearMethodBase):
is_weight=True, is_weight=True,
per_channel=True, per_channel=True,
convrot=True, convrot=True,
convrot_groupsize=self.quant_config.group_size, convrot_groupsize=self.group_size,
stochastic_rounding=0, stochastic_rounding=0,
) )
layer.weight = Parameter(qdata.to(home), requires_grad=False) layer.weight = Parameter(qdata.to(home), requires_grad=False)
@@ -171,7 +189,7 @@ class KitchenInt8LinearMethod(LinearMethodBase):
bias, bias,
out_code, out_code,
True, # convrot True, # convrot
self.quant_config.group_size, self.group_size,
) )
n_rows, n_out = x.shape[0], layer.weight.shape[0] n_rows, n_out = x.shape[0], layer.weight.shape[0]
@@ -280,10 +280,10 @@ class TransformerLoader(ComponentLoader):
or cpu_offload_flag or cpu_offload_flag
) )
use_fsdp = server_args.should_use_fsdp_for_component(component_name) use_fsdp = server_args.should_use_fsdp_for_component(component_name)
if quant_spec.is_comfy_fp8 and use_fsdp: if quant_spec.uses_comfy_layer_markers and use_fsdp:
raise ValueError( raise ValueError(
"MiniMax-H3 Comfy FP8 does not support FSDP inference; use TP " "Comfy quantized checkpoints do not support FSDP "
"and/or sequence parallelism instead" "inference; use TP and/or sequence parallelism instead"
) )
if quant_spec.gguf_file is not None: if quant_spec.gguf_file is not None:
@@ -308,7 +308,7 @@ class TransformerLoader(ComponentLoader):
"quant_config": quant_spec.runtime_quant_config, "quant_config": quant_spec.runtime_quant_config,
} }
checkpoint_key_filter: Callable[[str], bool] | None = ( checkpoint_key_filter: Callable[[str], bool] | None = (
comfy_quant_key_filter if quant_spec.is_comfy_fp8 else None comfy_quant_key_filter if quant_spec.uses_comfy_layer_markers else None
) )
adaln_cache_path = component_server_args.minimax_h3_adaln_cache_path adaln_cache_path = component_server_args.minimax_h3_adaln_cache_path
if adaln_cache_path is not None: if adaln_cache_path is not None:
@@ -362,7 +362,10 @@ class TransformerLoader(ComponentLoader):
local_torch_device, local_torch_device,
component_starts_on_cpu=component_starts_on_cpu, component_starts_on_cpu=component_starts_on_cpu,
runtime_quant_config=quant_spec.runtime_quant_config, runtime_quant_config=quant_spec.runtime_quant_config,
quantized_cpu_load_supported=quant_spec.gguf_file is not None, quantized_cpu_load_supported=(
quant_spec.gguf_file is not None
or quant_spec.is_serialized_kitchen_int8
),
) )
) )
direct_gpu_weight_loading = bool( direct_gpu_weight_loading = bool(
@@ -10,6 +10,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig, QuantizationConfig,
) )
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
KitchenInt8Config,
)
def comfy_quant_key_filter(name: str) -> bool: def comfy_quant_key_filter(name: str) -> bool:
@@ -23,6 +26,7 @@ def inspect_minimax_h3_safetensors(
adaln_curve_shape = None adaln_curve_shape = None
layer_markers: dict[str, dict[str, Any]] = {} layer_markers: dict[str, dict[str, Any]] = {}
checkpoint_keys: set[str] = set() checkpoint_keys: set[str] = set()
checkpoint_meta: dict[str, tuple[str, tuple[int, ...]]] = {}
fp8_weight_prefixes: set[str] = set() fp8_weight_prefixes: set[str] = set()
for path in safetensors_list: for path in safetensors_list:
@@ -44,11 +48,15 @@ def inspect_minimax_h3_safetensors(
adaln_curve_shape = shape adaln_curve_shape = shape
for key in keys: for key in keys:
if ( if key.endswith((".weight", ".weight_scale")):
key.endswith(".weight") tensor_slice = checkpoint.get_slice(key)
and checkpoint.get_slice(key).get_dtype() == "F8_E4M3" dtype = tensor_slice.get_dtype()
): checkpoint_meta[key] = (
fp8_weight_prefixes.add(key.removesuffix(".weight")) dtype,
tuple(tensor_slice.get_shape()),
)
if key.endswith(".weight") and dtype == "F8_E4M3":
fp8_weight_prefixes.add(key.removesuffix(".weight"))
if not key.endswith(".comfy_quant"): if not key.endswith(".comfy_quant"):
continue continue
try: try:
@@ -78,17 +86,39 @@ def inspect_minimax_h3_safetensors(
) )
for prefix, marker in layer_markers.items(): for prefix, marker in layer_markers.items():
if marker.get("format") != "float8_e4m3fn": marker_format = marker.get("format")
continue
required = {f"{prefix}.weight", f"{prefix}.weight_scale"} required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
if not marker.get("full_precision_matrix_mult", False): if marker_format == "float8_e4m3fn" and not marker.get(
"full_precision_matrix_mult", False
):
required.add(f"{prefix}.input_scale") required.add(f"{prefix}.input_scale")
if marker_format not in ("float8_e4m3fn", "int8_tensorwise"):
continue
missing = required - checkpoint_keys missing = required - checkpoint_keys
if missing: if missing:
raise ValueError( raise ValueError(
f"MiniMax-H3 Comfy FP8 layer {prefix!r} is missing checkpoint " f"MiniMax-H3 Comfy layer {prefix!r} is missing checkpoint "
f"tensors: {sorted(missing)}" f"tensors: {sorted(missing)}"
) )
if marker_format == "int8_tensorwise":
weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"]
scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"]
if weight_dtype != "I8" or scale_dtype != "F32":
raise ValueError(
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs I8 weights "
f"and F32 scales, got {weight_dtype} and {scale_dtype}"
)
if len(weight_shape) != 2:
raise ValueError(
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs a 2D weight, "
f"got {weight_shape}"
)
expected_scale_shape = (weight_shape[0], 1)
if scale_shape != expected_scale_shape:
raise ValueError(
f"MiniMax-H3 Comfy INT8 layer {prefix!r} needs scale shape "
f"{expected_scale_shape}, got {scale_shape}"
)
return adaln_curve_shape, layer_markers return adaln_curve_shape, layer_markers
@@ -100,13 +130,8 @@ def resolve_minimax_h3_checkpoint_quantization(
return None return None
formats = sorted({str(marker.get("format")) for marker in layer_markers.values()}) formats = sorted({str(marker.get("format")) for marker in layer_markers.values()})
if "int8_tensorwise" in formats: if formats == ["int8_tensorwise"]:
raise NotImplementedError( return KitchenInt8Config(layer_markers=layer_markers)
"MiniMax-H3 pruned_int8_convrot is not supported yet. Its "
"int8_tensorwise weights require an online regular-Hadamard ConvRot "
"and dynamic INT8 activation quantization kernel; loading them as "
"ordinary INT8/BF16 weights would produce incorrect output."
)
if formats == ["float8_e4m3fn"]: if formats == ["float8_e4m3fn"]:
return ComfyFp8Config(layer_markers) return ComfyFp8Config(layer_markers)
raise NotImplementedError( raise NotImplementedError(
@@ -18,6 +18,9 @@ from diffusers.utils import SAFE_WEIGHTS_INDEX_NAME
from torch import nn from torch import nn
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
KitchenInt8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig, NunchakuConfig,
_patch_nunchaku_scales, _patch_nunchaku_scales,
@@ -164,6 +167,17 @@ class TransformerQuantLoadSpec:
def is_comfy_fp8(self) -> bool: def is_comfy_fp8(self) -> bool:
return _get_quant_config_name(self.quant_config) == "comfy_fp8" return _get_quant_config_name(self.quant_config) == "comfy_fp8"
@property
def is_serialized_kitchen_int8(self) -> bool:
return (
isinstance(self.quant_config, KitchenInt8Config)
and self.quant_config.is_checkpoint_int8_serialized
)
@property
def uses_comfy_layer_markers(self) -> bool:
return self.is_comfy_fp8 or self.is_serialized_kitchen_int8
class _TransformerQuantAdapter: class _TransformerQuantAdapter:
def prepare(self) -> None: def prepare(self) -> None:
@@ -844,6 +858,9 @@ def _needs_device_weight_postprocess(
quant_name = _get_quant_config_name(quant_config) quant_name = _get_quant_config_name(quant_config)
if quant_name in ("modelopt_fp8", "comfy_fp8"): if quant_name in ("modelopt_fp8", "comfy_fp8"):
return True return True
if quant_name == "kitchen_int8":
assert isinstance(quant_config, KitchenInt8Config)
return not quant_config.is_checkpoint_int8_serialized
serialized_flag_by_quant_name = { serialized_flag_by_quant_name = {
"fp8": "is_checkpoint_fp8_serialized", "fp8": "is_checkpoint_fp8_serialized",
@@ -612,10 +612,13 @@ class MiniMaxH3Attention(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.qkv_proj", prefix=f"{prefix}.qkv_proj",
) )
# The reorder below translates the *safetensors* checkpoint layout. A # Official safetensors interleave Q/K/V by head. Comfy and GGUF
# GGUF checkpoint already stores qkv as [q_all, k_all, v_all], and its # checkpoints already store [q_all, k_all, v_all].
# packed parameter is `qweight`, so there is nothing to reorder. checkpoint_qkv_is_native = quant_config is not None and (
if quant_config is None or quant_config.get_name() != "gguf": quant_config.get_name() == "gguf"
or quant_config.checkpoint_uses_native_qkv_layout
)
if not checkpoint_qkv_is_native:
self._install_qkv_weight_loader(arch) self._install_qkv_weight_loader(arch)
self.q_norm = _norm(arch.attention_head_dim, eps=arch.qk_norm_eps) self.q_norm = _norm(arch.attention_head_dim, eps=arch.qk_norm_eps)
self.k_norm = _norm(arch.attention_head_dim, eps=arch.qk_norm_eps) self.k_norm = _norm(arch.attention_head_dim, eps=arch.qk_norm_eps)
@@ -47,12 +47,16 @@ sys.modules.setdefault("partial_json_parser.core.options", partial_json_parser_o
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
LinearBase, LinearBase,
ReplicatedLinear,
UnquantizedLinearMethod, UnquantizedLinearMethod,
) )
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ( from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import (
ComfyFp8Config, ComfyFp8Config,
ComfyFullPrecisionFp8LinearMethod, ComfyFullPrecisionFp8LinearMethod,
) )
from sglang.multimodal_gen.runtime.layers.quantization.configs.kitchen_int8_config import (
KitchenInt8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig, NunchakuConfig,
) )
@@ -246,11 +250,19 @@ class TestTransformerQuantHelpers(unittest.TestCase):
mock_download.reset_mock() mock_download.reset_mock()
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self): def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self):
marker = json.dumps({"format": "int8_tensorwise", "convrot": True}).encode() marker = json.dumps(
{
"format": "int8_tensorwise",
"convrot": True,
"convrot_groupsize": 256,
}
).encode()
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f: with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
save_file( save_file(
{ {
"adaln_t_table": torch.zeros((1025, 8)), "adaln_t_table": torch.zeros((1025, 8)),
"blocks.0.mlp.fc1.weight": torch.ones((2, 256), dtype=torch.int8),
"blocks.0.mlp.fc1.weight_scale": torch.ones((2, 1)),
"blocks.0.mlp.fc1.comfy_quant": torch.tensor( "blocks.0.mlp.fc1.comfy_quant": torch.tensor(
list(marker), dtype=torch.uint8 list(marker), dtype=torch.uint8
), ),
@@ -282,13 +294,60 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertEqual(layer_markers["blocks.0.mlp.fc1"], {"format": "float8_e4m3fn"}) self.assertEqual(layer_markers["blocks.0.mlp.fc1"], {"format": "float8_e4m3fn"})
def test_minimax_h3_comfy_int8_fails_before_weight_loading(self): def test_minimax_h3_comfy_int8_resolves_serialized_kitchen(self):
with self.assertRaisesRegex(NotImplementedError, "regular-Hadamard"): config = resolve_minimax_h3_checkpoint_quantization(
{
"blocks.0.mlp.fc1": {
"format": "int8_tensorwise",
"convrot": True,
"convrot_groupsize": 256,
}
}
)
self.assertIsInstance(config, KitchenInt8Config)
self.assertTrue(config.is_checkpoint_int8_serialized)
self.assertTrue(config.checkpoint_uses_native_qkv_layout)
self.assertFalse(KitchenInt8Config().checkpoint_uses_native_qkv_layout)
self.assertFalse(_needs_device_weight_postprocess(config))
@patch(
"sglang.multimodal_gen.runtime.layers.quantization.kitchen_int8."
"_load_comfy_kitchen"
)
def test_serialized_kitchen_constructs_int8_weight_and_row_scale(self, _load):
config = KitchenInt8Config(
layer_markers={
"proj": {
"format": "int8_tensorwise",
"convrot": True,
"convrot_groupsize": 256,
}
}
)
layer = ReplicatedLinear(
256,
3,
bias=False,
params_dtype=torch.bfloat16,
quant_config=config,
prefix="proj",
)
self.assertEqual(layer.weight.dtype, torch.int8)
self.assertEqual(layer.weight.shape, (3, 256))
self.assertEqual(layer.weight_scale.dtype, torch.float32)
self.assertEqual(layer.weight_scale.shape, (3, 1))
def test_serialized_kitchen_rejects_non_convrot_marker(self):
with self.assertRaisesRegex(ValueError, "convrot=true"):
resolve_minimax_h3_checkpoint_quantization( resolve_minimax_h3_checkpoint_quantization(
{ {
"blocks.0.mlp.fc1": { "blocks.0.mlp.fc1": {
"format": "int8_tensorwise", "format": "int8_tensorwise",
"convrot": True, "convrot": False,
"convrot_groupsize": 256,
} }
} }
) )
@@ -305,6 +364,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
) )
self.assertIsInstance(config, ComfyFp8Config) self.assertIsInstance(config, ComfyFp8Config)
self.assertTrue(config.checkpoint_uses_native_qkv_layout)
layer = LinearBase(input_size=1, output_size=1) layer = LinearBase(input_size=1, output_size=1)
self.assertIsInstance( self.assertIsInstance(
config.get_quant_method(layer, "blocks.0.mlp.fc2"), config.get_quant_method(layer, "blocks.0.mlp.fc2"),