[diffusion] feat: support serialized comfy convrot int8 dits (#35994)
This commit is contained in:
@@ -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>
|
||||||
|
|||||||
@@ -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__()
|
||||||
|
|||||||
+48
-13
@@ -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"),
|
||||||
|
|||||||
Reference in New Issue
Block a user