[diffusion] quant: support pruned safetensors checkpoints for minimax-h3 (#35418)
This commit is contained in:
@@ -248,6 +248,59 @@ server environment.
|
|||||||
|
|
||||||
For MiniMax-H3, `--performance-mode speed` deliberately keeps the DiT eager. The current `torch.compile` path changes the model's numerical output, so it is not enabled implicitly by any recommended lossless preset. An explicit `--enable-torch-compile true` remains available for controlled experiments, but it should not be used to generate consistency ground truth.
|
For MiniMax-H3, `--performance-mode speed` deliberately keeps the DiT eager. The current `torch.compile` path changes the model's numerical output, so it is not enabled implicitly by any recommended lossless preset. An explicit `--enable-torch-compile true` remains available for controlled experiments, but it should not be used to generate consistency ground truth.
|
||||||
|
|
||||||
|
### AdaLN-pruned safetensors transformers
|
||||||
|
|
||||||
|
[Comfy-Org/MiniMax-H3](https://huggingface.co/Comfy-Org/MiniMax-H3/tree/main/diffusion_models)
|
||||||
|
publishes smaller DiT-only checkpoints that replace the original AdaLN branches
|
||||||
|
with an interpolated curve table. Select one file explicitly; the base model
|
||||||
|
still supplies the text encoder and VAEs.
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
sglang serve \
|
||||||
|
--model-path MiniMaxAI/MiniMax-H3 \
|
||||||
|
--model-variant fl2va \
|
||||||
|
--transformer-weights-path \
|
||||||
|
Comfy-Org/MiniMax-H3/diffusion_models/minimax_h3_fl2va_pruned_bf16.safetensors \
|
||||||
|
--num-gpus 4 \
|
||||||
|
--tp-size 2 \
|
||||||
|
--ulysses-degree 2 \
|
||||||
|
--performance-mode speed \
|
||||||
|
--port 30010
|
||||||
|
```
|
||||||
|
|
||||||
|
The pruned checkpoint is approximate and is therefore rejected by
|
||||||
|
`quality="high"`, which remains limited to the audited official BF16 DiT.
|
||||||
|
|
||||||
|
The Comfy `pruned_fp8_scaled` FL2VA and Ref2VA files are also supported. Their
|
||||||
|
per-layer markers are detected automatically; do not add `--quantization`:
|
||||||
|
|
||||||
|
```bash Command
|
||||||
|
sglang serve \
|
||||||
|
--model-path MiniMaxAI/MiniMax-H3 \
|
||||||
|
--model-variant fl2va \
|
||||||
|
--transformer-weights-path \
|
||||||
|
Comfy-Org/MiniMax-H3/diffusion_models/minimax_h3_fl2va_pruned_fp8_scaled.safetensors \
|
||||||
|
--num-gpus 4 \
|
||||||
|
--tp-size 2 \
|
||||||
|
--ulysses-degree 2 \
|
||||||
|
--performance-mode speed \
|
||||||
|
--port 30010
|
||||||
|
```
|
||||||
|
|
||||||
|
SGLang uses its native static-activation FP8 linear path for attention and
|
||||||
|
`fc1`. The checkpoint marks `fc2` for full-precision matrix multiplication, so
|
||||||
|
SGLang retains its FP8 storage but materializes and scales one compute-dtype
|
||||||
|
`fc2` matrix for each call. This preserves the checkpoint's mixed execution
|
||||||
|
contract and low resident weight memory, but that part is slower than a fully
|
||||||
|
quantized FP8 GEMM. TP, Ulysses/Ring sequence parallelism, and
|
||||||
|
component/layerwise offload are supported; FSDP inference is rejected.
|
||||||
|
|
||||||
|
The `pruned_int8_convrot` files are detected but remain unsupported. They
|
||||||
|
require online regular-Hadamard ConvRot, dynamic INT8 activation quantization,
|
||||||
|
and a matching W8A8 GEMM. SGLang fails before loading them instead of silently
|
||||||
|
treating their stored INT8 values as ordinary weights. A native ConvRot kernel
|
||||||
|
path should be added and benchmarked separately before these files are accepted.
|
||||||
|
|
||||||
### Advanced: precomputed AdaLN cache
|
### Advanced: precomputed AdaLN cache
|
||||||
|
|
||||||
The [model card](https://huggingface.co/MiniMaxAI/MiniMax-H3) notes that about
|
The [model card](https://huggingface.co/MiniMaxAI/MiniMax-H3) notes that about
|
||||||
|
|||||||
@@ -15,7 +15,9 @@ Use these paths:
|
|||||||
|
|
||||||
- `--model-path`: the base or original model
|
- `--model-path`: the base or original model
|
||||||
- `--transformer-path`: a quantized transformers-style transformer component directory that already contains its own `config.json`
|
- `--transformer-path`: a quantized transformers-style transformer component directory that already contains its own `config.json`
|
||||||
- `--transformer-weights-path`: quantized transformer weights provided as a single safetensors file, a sharded safetensors directory, a local path, or a Hugging Face repo ID
|
- `--transformer-weights-path`: replacement transformer weights in safetensors
|
||||||
|
format (file, directory, or Hub repository/file) or a supported GGUF file
|
||||||
|
(local or Hub)
|
||||||
- `--quantization`: apply online quantization to unquantized models at load time (activations are quantized dynamically)
|
- `--quantization`: apply online quantization to unquantized models at load time (activations are quantized dynamically)
|
||||||
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
- `--quantization-ignored-layers` layer name patterns to keep unquantized (e.g. `attention.to_`)
|
||||||
- `--component-paths.text_encoder`: replace a native text encoder with a checkpoint whose `quantization_config` is auto-detected
|
- `--component-paths.text_encoder`: replace a native text encoder with a checkpoint whose `quantization_config` is auto-detected
|
||||||
@@ -45,6 +47,12 @@ 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
|
||||||
|
`pruned_fp8_scaled` safetensors. Pass one selected FL2VA or Ref2VA file by local
|
||||||
|
path, `owner/repo/path/file.safetensors`, or direct Hugging Face file URL; do
|
||||||
|
not combine it with `--quantization`. MiniMax-H3 GGUF usage is documented in
|
||||||
|
the [MiniMax-H3 cookbook](/cookbook/diffusion/MiniMax/MiniMax-H3#pre-quantized-gguf-transformer).
|
||||||
|
|
||||||
## Quant Families
|
## Quant Families
|
||||||
|
|
||||||
Here, `quant_family` means a checkpoint and loading family with shared CLI
|
Here, `quant_family` means a checkpoint and loading family with shared CLI
|
||||||
@@ -111,6 +119,22 @@ backend.
|
|||||||
<td>None</td>
|
<td>None</td>
|
||||||
<td>Mixed override repos keep the base model separate; full Qwen Image exports can be loaded directly as <code>--model-path</code>; raw exports such as <code>black-forest-labs/FLUX.2-dev-NVFP4</code> still use the weights-path flow</td>
|
<td>Mixed override repos keep the base model separate; full Qwen Image exports can be loaded directly as <code>--model-path</code>; raw exports such as <code>black-forest-labs/FLUX.2-dev-NVFP4</code> still use the weights-path flow</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><code>gguf</code></td>
|
||||||
|
<td>One selected GGUF DiT file</td>
|
||||||
|
<td><code>--transformer-weights-path</code></td>
|
||||||
|
<td>MiniMax-H3 original or pruned FL2VA / Ref2VA DiTs</td>
|
||||||
|
<td>None</td>
|
||||||
|
<td>CUDA only; auto-detected; supports standard and K-quant GGML types; FSDP and the separate Qwen3-VL text-encoder GGUF files are not supported</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><code>comfy-fp8</code></td>
|
||||||
|
<td>One selected Comfy safetensors file with per-layer <code>comfy_quant</code> metadata</td>
|
||||||
|
<td><code>--transformer-weights-path</code></td>
|
||||||
|
<td>MiniMax-H3 pruned FL2VA / Ref2VA DiTs</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>
|
||||||
|
</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>
|
||||||
|
|||||||
@@ -0,0 +1,145 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Comfy per-layer FP8 checkpoint support."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
LinearBase,
|
||||||
|
LinearMethodBase,
|
||||||
|
UnquantizedLinearMethod,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
QuantizeMethodBase,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import (
|
||||||
|
Fp8Config,
|
||||||
|
Fp8LinearMethod,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.models.parameter import (
|
||||||
|
ModelWeightParameter,
|
||||||
|
PerTensorScaleParameter,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ComfyFullPrecisionFp8LinearMethod(LinearMethodBase):
|
||||||
|
"""Keep FP8 storage but honor Comfy's full-precision matmul marker."""
|
||||||
|
|
||||||
|
def create_weights(
|
||||||
|
self,
|
||||||
|
layer: nn.Module,
|
||||||
|
input_size_per_partition: int,
|
||||||
|
output_partition_sizes: list[int],
|
||||||
|
input_size: int,
|
||||||
|
output_size: int,
|
||||||
|
params_dtype: torch.dtype,
|
||||||
|
**extra_weight_attrs: Any,
|
||||||
|
) -> None:
|
||||||
|
if len(output_partition_sizes) != 1:
|
||||||
|
raise ValueError(
|
||||||
|
"Comfy full_precision_matrix_mult does not support fused linears"
|
||||||
|
)
|
||||||
|
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||||
|
layer.logical_widths = output_partition_sizes
|
||||||
|
layer.input_size_per_partition = input_size_per_partition
|
||||||
|
layer.output_size_per_partition = output_partition_sizes[0]
|
||||||
|
layer.orig_dtype = params_dtype
|
||||||
|
weight = ModelWeightParameter(
|
||||||
|
data=torch.empty(
|
||||||
|
output_partition_sizes[0],
|
||||||
|
input_size_per_partition,
|
||||||
|
dtype=torch.float8_e4m3fn,
|
||||||
|
),
|
||||||
|
input_dim=1,
|
||||||
|
output_dim=0,
|
||||||
|
weight_loader=weight_loader,
|
||||||
|
)
|
||||||
|
layer.register_parameter("weight", weight)
|
||||||
|
weight_scale = PerTensorScaleParameter(
|
||||||
|
data=torch.empty(1, dtype=torch.float32),
|
||||||
|
weight_loader=weight_loader,
|
||||||
|
)
|
||||||
|
layer.register_parameter("weight_scale", weight_scale)
|
||||||
|
|
||||||
|
def process_weights_after_loading(self, layer: nn.Module) -> None:
|
||||||
|
layer.weight = nn.Parameter(layer.weight.data, requires_grad=False)
|
||||||
|
layer.weight_scale = nn.Parameter(layer.weight_scale.data, requires_grad=False)
|
||||||
|
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
layer: nn.Module,
|
||||||
|
x: torch.Tensor,
|
||||||
|
bias: torch.Tensor | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# Comfy's marker disables quantized GEMM for this layer. Materialize one
|
||||||
|
# compute-dtype matrix per call so the checkpoint's FP8 residency benefit
|
||||||
|
# is retained instead of permanently expanding every fc2 weight to BF16.
|
||||||
|
weight = layer.weight.to(dtype=x.dtype)
|
||||||
|
weight.mul_(layer.weight_scale[0].to(dtype=x.dtype))
|
||||||
|
return F.linear(x, weight, bias)
|
||||||
|
|
||||||
|
|
||||||
|
class ComfyFp8Config(QuantizationConfig):
|
||||||
|
"""Dispatch each Linear according to its serialized ``comfy_quant`` marker."""
|
||||||
|
|
||||||
|
def __init__(self, layer_markers: dict[str, dict[str, Any]]) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.layer_markers = layer_markers
|
||||||
|
self._fp8_config = Fp8Config(
|
||||||
|
is_checkpoint_fp8_serialized=True,
|
||||||
|
activation_scheme="static",
|
||||||
|
)
|
||||||
|
|
||||||
|
unsupported = {
|
||||||
|
prefix: marker.get("format")
|
||||||
|
for prefix, marker in layer_markers.items()
|
||||||
|
if marker.get("format") != "float8_e4m3fn"
|
||||||
|
}
|
||||||
|
if unsupported:
|
||||||
|
raise ValueError(f"Unsupported Comfy FP8 layer formats: {unsupported}")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_name(cls) -> str:
|
||||||
|
return "comfy_fp8"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
|
||||||
|
return Fp8Config.get_supported_act_dtypes()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_min_capability(cls) -> int:
|
||||||
|
return Fp8Config.get_min_capability()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_config_filenames() -> list[str]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config(cls, config: dict[str, Any]) -> ComfyFp8Config:
|
||||||
|
raise ValueError(
|
||||||
|
"ComfyFp8Config must be constructed from safetensors layer markers"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_quant_method(
|
||||||
|
self, layer: nn.Module, prefix: str
|
||||||
|
) -> QuantizeMethodBase | None:
|
||||||
|
if not isinstance(layer, LinearBase):
|
||||||
|
return None
|
||||||
|
marker = self.layer_markers.get(prefix)
|
||||||
|
if marker is None:
|
||||||
|
return UnquantizedLinearMethod()
|
||||||
|
if marker.get("full_precision_matrix_mult", False):
|
||||||
|
return ComfyFullPrecisionFp8LinearMethod()
|
||||||
|
return Fp8LinearMethod(self._fp8_config)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ComfyFp8Config",
|
||||||
|
"ComfyFullPrecisionFp8LinearMethod",
|
||||||
|
]
|
||||||
+47
-3
@@ -17,6 +17,12 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model
|
from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model
|
||||||
from sglang.multimodal_gen.runtime.loader.gguf_weights import gguf_weights_iterator
|
from sglang.multimodal_gen.runtime.loader.gguf_weights import gguf_weights_iterator
|
||||||
|
from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import (
|
||||||
|
comfy_quant_key_filter,
|
||||||
|
inspect_minimax_h3_safetensors,
|
||||||
|
resolve_minimax_h3_checkpoint_quantization,
|
||||||
|
validate_minimax_h3_checkpoint_variant,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
||||||
TransformerQuantLoadSpec,
|
TransformerQuantLoadSpec,
|
||||||
resolve_transformer_gguf_to_load,
|
resolve_transformer_gguf_to_load,
|
||||||
@@ -214,6 +220,35 @@ class TransformerLoader(ComponentLoader):
|
|||||||
cls_name = config.pop("_class_name")
|
cls_name = config.pop("_class_name")
|
||||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||||
|
|
||||||
|
checkpoint_quant_config = None
|
||||||
|
if cls_name == "MiniMaxH3DiTModel":
|
||||||
|
selected_variant = str(component_server_args.model_variant or "fl2va")
|
||||||
|
if gguf_file is not None:
|
||||||
|
validate_minimax_h3_checkpoint_variant([gguf_file], selected_variant)
|
||||||
|
elif component_server_args.transformer_weights_path is not None:
|
||||||
|
validate_minimax_h3_checkpoint_variant(
|
||||||
|
safetensors_list, selected_variant
|
||||||
|
)
|
||||||
|
adaln_curve_shape, layer_markers = inspect_minimax_h3_safetensors(
|
||||||
|
safetensors_list
|
||||||
|
)
|
||||||
|
checkpoint_quant_config = resolve_minimax_h3_checkpoint_quantization(
|
||||||
|
layer_markers
|
||||||
|
)
|
||||||
|
if adaln_curve_shape is not None:
|
||||||
|
(
|
||||||
|
dit_config.arch_config.adaln_curve_grid,
|
||||||
|
dit_config.arch_config.time_embed_dim,
|
||||||
|
) = adaln_curve_shape
|
||||||
|
if (
|
||||||
|
component_server_args.minimax_h3_adaln_cache_path is not None
|
||||||
|
or component_server_args.minimax_h3_adaln_online
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax-H3 pruned curve checkpoints cannot use a "
|
||||||
|
"separate AdaLN cache or online AdaLN rebuild"
|
||||||
|
)
|
||||||
|
|
||||||
quant_spec = resolve_transformer_quant_load_spec(
|
quant_spec = resolve_transformer_quant_load_spec(
|
||||||
hf_config=config,
|
hf_config=config,
|
||||||
server_args=component_server_args,
|
server_args=component_server_args,
|
||||||
@@ -223,6 +258,7 @@ class TransformerLoader(ComponentLoader):
|
|||||||
cls_name=cls_name,
|
cls_name=cls_name,
|
||||||
component_name=component_name,
|
component_name=component_name,
|
||||||
gguf_file=gguf_file,
|
gguf_file=gguf_file,
|
||||||
|
checkpoint_quant_config=checkpoint_quant_config,
|
||||||
)
|
)
|
||||||
if quant_spec.gguf_file is not None and cls_name == "MiniMaxH3DiTModel":
|
if quant_spec.gguf_file is not None and cls_name == "MiniMaxH3DiTModel":
|
||||||
assert quant_spec.quant_config is not None
|
assert quant_spec.quant_config is not None
|
||||||
@@ -244,6 +280,11 @@ 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:
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax-H3 Comfy FP8 does not support FSDP inference; use TP "
|
||||||
|
"and/or sequence parallelism instead"
|
||||||
|
)
|
||||||
|
|
||||||
if quant_spec.gguf_file is not None:
|
if quant_spec.gguf_file is not None:
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -266,7 +307,9 @@ class TransformerLoader(ComponentLoader):
|
|||||||
"hf_config": config,
|
"hf_config": config,
|
||||||
"quant_config": quant_spec.runtime_quant_config,
|
"quant_config": quant_spec.runtime_quant_config,
|
||||||
}
|
}
|
||||||
checkpoint_key_filter: Callable[[str], bool] | None = None
|
checkpoint_key_filter: Callable[[str], bool] | None = (
|
||||||
|
comfy_quant_key_filter if quant_spec.is_comfy_fp8 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:
|
||||||
if cls_name != "MiniMaxH3DiTModel":
|
if cls_name != "MiniMaxH3DiTModel":
|
||||||
@@ -304,8 +347,9 @@ class TransformerLoader(ComponentLoader):
|
|||||||
init_params["quant_config"] is None
|
init_params["quant_config"] is None
|
||||||
and component_server_args.transformer_weights_path is not None
|
and component_server_args.transformer_weights_path is not None
|
||||||
):
|
):
|
||||||
logger.warning(
|
logger.info(
|
||||||
"transformer_weights_path provided, but quantization config not resolved, which is unexpected and likely to cause errors"
|
"Using an unquantized transformer weight override from %s",
|
||||||
|
component_server_args.transformer_weights_path,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.debug("quantization config: %s", init_params["quant_config"])
|
logger.debug("quantization config: %s", init_params["quant_config"])
|
||||||
|
|||||||
@@ -0,0 +1,139 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Checkpoint inspection for MiniMax-H3 transformer overrides."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from safetensors import safe_open
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import ComfyFp8Config
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def comfy_quant_key_filter(name: str) -> bool:
|
||||||
|
return not name.endswith(".comfy_quant")
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_minimax_h3_safetensors(
|
||||||
|
safetensors_list: list[str],
|
||||||
|
) -> tuple[tuple[int, int] | None, dict[str, dict[str, Any]]]:
|
||||||
|
"""Read H3 architecture metadata and Comfy per-layer format markers."""
|
||||||
|
adaln_curve_shape = None
|
||||||
|
layer_markers: dict[str, dict[str, Any]] = {}
|
||||||
|
checkpoint_keys: set[str] = set()
|
||||||
|
fp8_weight_prefixes: set[str] = set()
|
||||||
|
|
||||||
|
for path in safetensors_list:
|
||||||
|
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
||||||
|
keys = checkpoint.keys()
|
||||||
|
checkpoint_keys.update(keys)
|
||||||
|
if "adaln_t_table" in keys:
|
||||||
|
shape = tuple(checkpoint.get_slice("adaln_t_table").get_shape())
|
||||||
|
if len(shape) != 2 or shape[0] < 2:
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax-H3 adaln_t_table must have shape [N, D] with "
|
||||||
|
f"N >= 2, got {shape} in {path}"
|
||||||
|
)
|
||||||
|
if adaln_curve_shape is not None and adaln_curve_shape != shape:
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax-H3 checkpoint shards disagree on adaln_t_table "
|
||||||
|
f"shape: {adaln_curve_shape} vs {shape}"
|
||||||
|
)
|
||||||
|
adaln_curve_shape = shape
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if (
|
||||||
|
key.endswith(".weight")
|
||||||
|
and checkpoint.get_slice(key).get_dtype() == "F8_E4M3"
|
||||||
|
):
|
||||||
|
fp8_weight_prefixes.add(key.removesuffix(".weight"))
|
||||||
|
if not key.endswith(".comfy_quant"):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
marker = json.loads(checkpoint.get_tensor(key).numpy().tobytes())
|
||||||
|
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid Comfy quantization marker {key!r} in {path}"
|
||||||
|
) from exc
|
||||||
|
if not isinstance(marker, dict):
|
||||||
|
raise ValueError(
|
||||||
|
f"Comfy quantization marker {key!r} must contain a JSON object"
|
||||||
|
)
|
||||||
|
prefix = key.removesuffix(".comfy_quant")
|
||||||
|
previous = layer_markers.get(prefix)
|
||||||
|
if previous is not None and previous != marker:
|
||||||
|
raise ValueError(
|
||||||
|
f"Conflicting Comfy quantization markers for {prefix!r}"
|
||||||
|
)
|
||||||
|
layer_markers[prefix] = marker
|
||||||
|
|
||||||
|
if layer_markers:
|
||||||
|
missing_markers = fp8_weight_prefixes - layer_markers.keys()
|
||||||
|
if missing_markers:
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax-H3 FP8 weights are missing comfy_quant metadata: "
|
||||||
|
f"{sorted(missing_markers)[:5]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for prefix, marker in layer_markers.items():
|
||||||
|
if marker.get("format") != "float8_e4m3fn":
|
||||||
|
continue
|
||||||
|
required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
|
||||||
|
if not marker.get("full_precision_matrix_mult", False):
|
||||||
|
required.add(f"{prefix}.input_scale")
|
||||||
|
missing = required - checkpoint_keys
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"MiniMax-H3 Comfy FP8 layer {prefix!r} is missing checkpoint "
|
||||||
|
f"tensors: {sorted(missing)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return adaln_curve_shape, layer_markers
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_minimax_h3_checkpoint_quantization(
|
||||||
|
layer_markers: dict[str, dict[str, Any]],
|
||||||
|
) -> QuantizationConfig | None:
|
||||||
|
if not layer_markers:
|
||||||
|
return None
|
||||||
|
|
||||||
|
formats = sorted({str(marker.get("format")) for marker in layer_markers.values()})
|
||||||
|
if "int8_tensorwise" in formats:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"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"]:
|
||||||
|
return ComfyFp8Config(layer_markers)
|
||||||
|
raise NotImplementedError(
|
||||||
|
"Unsupported MiniMax-H3 Comfy quantization format(s): " + ", ".join(formats)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def validate_minimax_h3_checkpoint_variant(
|
||||||
|
checkpoint_paths: list[str], selected_variant: str
|
||||||
|
) -> None:
|
||||||
|
names = " ".join(path.lower() for path in checkpoint_paths)
|
||||||
|
checkpoint_variant = next(
|
||||||
|
(variant for variant in ("fl2va", "ref2va") if variant in names), None
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
checkpoint_variant is not None
|
||||||
|
and checkpoint_variant != selected_variant.lower()
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"MiniMax-H3 checkpoint variant {checkpoint_variant!r} does not match "
|
||||||
|
f"--model-variant {selected_variant!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"comfy_quant_key_filter",
|
||||||
|
"inspect_minimax_h3_safetensors",
|
||||||
|
"resolve_minimax_h3_checkpoint_quantization",
|
||||||
|
"validate_minimax_h3_checkpoint_variant",
|
||||||
|
]
|
||||||
@@ -37,6 +37,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
|
hf_hub_download,
|
||||||
maybe_download_model,
|
maybe_download_model,
|
||||||
snapshot_download,
|
snapshot_download,
|
||||||
)
|
)
|
||||||
@@ -61,6 +62,11 @@ _PRECISION_VARIANT_SUFFIX_RE = re.compile(
|
|||||||
r"^(?P<stem>.+?)(?P<precision>\.(?:fp16|bf16|fp32))(?P<shard>-\d+-of-\d+)?(?P<ext>\.safetensors)$"
|
r"^(?P<stem>.+?)(?P<precision>\.(?:fp16|bf16|fp32))(?P<shard>-\d+-of-\d+)?(?P<ext>\.safetensors)$"
|
||||||
)
|
)
|
||||||
_MIXED_SAFETENSORS_RE = re.compile(r".*-mixed(?:-\d+-of-\d+)?\.safetensors$")
|
_MIXED_SAFETENSORS_RE = re.compile(r".*-mixed(?:-\d+-of-\d+)?\.safetensors$")
|
||||||
|
_HF_SAFETENSORS_URL_RE = re.compile(
|
||||||
|
r"https?://huggingface\.co/(?P<repo>[^/]+/[^/]+)/"
|
||||||
|
r"(?:blob|resolve)/(?P<revision>[^/]+)/(?P<filename>.+\.safetensors)$",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _get_quant_config_name(config: Optional[QuantizationConfig]) -> Optional[str]:
|
def _get_quant_config_name(config: Optional[QuantizationConfig]) -> Optional[str]:
|
||||||
@@ -154,6 +160,10 @@ class TransformerQuantLoadSpec:
|
|||||||
def is_modelopt_fp4(self) -> bool:
|
def is_modelopt_fp4(self) -> bool:
|
||||||
return _get_quant_config_name(self.quant_config) == "modelopt_fp4"
|
return _get_quant_config_name(self.quant_config) == "modelopt_fp4"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_comfy_fp8(self) -> bool:
|
||||||
|
return _get_quant_config_name(self.quant_config) == "comfy_fp8"
|
||||||
|
|
||||||
|
|
||||||
class _TransformerQuantAdapter:
|
class _TransformerQuantAdapter:
|
||||||
def prepare(self) -> None:
|
def prepare(self) -> None:
|
||||||
@@ -567,7 +577,31 @@ def resolve_transformer_safetensors_to_load(
|
|||||||
|
|
||||||
if quantized_path:
|
if quantized_path:
|
||||||
original_quantized_path = quantized_path
|
original_quantized_path = quantized_path
|
||||||
quantized_path = maybe_download_model(original_quantized_path)
|
direct_url = _HF_SAFETENSORS_URL_RE.fullmatch(original_quantized_path)
|
||||||
|
if direct_url is not None:
|
||||||
|
quantized_path = hf_hub_download(
|
||||||
|
repo_id=direct_url.group("repo"),
|
||||||
|
filename=direct_url.group("filename"),
|
||||||
|
revision=direct_url.group("revision"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parts = original_quantized_path.strip("/").split("/")
|
||||||
|
is_hub_file = (
|
||||||
|
not os.path.exists(original_quantized_path)
|
||||||
|
and not os.path.isabs(original_quantized_path)
|
||||||
|
and not original_quantized_path.startswith((".", "~"))
|
||||||
|
and len(parts) > 2
|
||||||
|
and original_quantized_path.endswith(".safetensors")
|
||||||
|
)
|
||||||
|
quantized_path = (
|
||||||
|
hf_hub_download(
|
||||||
|
repo_id="/".join(parts[:2]),
|
||||||
|
filename="/".join(parts[2:]),
|
||||||
|
revision=server_args.revision,
|
||||||
|
)
|
||||||
|
if is_hub_file
|
||||||
|
else maybe_download_model(original_quantized_path)
|
||||||
|
)
|
||||||
logger.info("using quantized transformer weights from: %s", quantized_path)
|
logger.info("using quantized transformer weights from: %s", quantized_path)
|
||||||
if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"):
|
if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"):
|
||||||
safetensors_list = [quantized_path]
|
safetensors_list = [quantized_path]
|
||||||
@@ -694,8 +728,11 @@ def resolve_transformer_quant_load_spec(
|
|||||||
cls_name: str,
|
cls_name: str,
|
||||||
component_name: str | None = None,
|
component_name: str | None = None,
|
||||||
gguf_file: str | None = None,
|
gguf_file: str | None = None,
|
||||||
|
checkpoint_quant_config: QuantizationConfig | None = None,
|
||||||
) -> TransformerQuantLoadSpec:
|
) -> TransformerQuantLoadSpec:
|
||||||
if gguf_file is not None:
|
if gguf_file is not None:
|
||||||
|
if checkpoint_quant_config is not None:
|
||||||
|
raise ValueError("GGUF and safetensors quantization metadata conflict")
|
||||||
return _resolve_gguf_quant_load_spec(
|
return _resolve_gguf_quant_load_spec(
|
||||||
gguf_file=gguf_file,
|
gguf_file=gguf_file,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -703,7 +740,19 @@ def resolve_transformer_quant_load_spec(
|
|||||||
component_name=component_name,
|
component_name=component_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
if getattr(model_cls, "handles_checkpoint_quantization", False):
|
if checkpoint_quant_config is not None:
|
||||||
|
if server_args.quantization is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint quantization is encoded in per-layer metadata; do not "
|
||||||
|
"also set --quantization"
|
||||||
|
)
|
||||||
|
if server_args.nunchaku_config is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"Per-layer checkpoint quantization and Nunchaku are mutually "
|
||||||
|
"exclusive"
|
||||||
|
)
|
||||||
|
quant_config = checkpoint_quant_config
|
||||||
|
elif getattr(model_cls, "handles_checkpoint_quantization", False):
|
||||||
quant_config = None
|
quant_config = None
|
||||||
else:
|
else:
|
||||||
quant_config = _resolve_quant_config(
|
quant_config = _resolve_quant_config(
|
||||||
@@ -793,7 +842,7 @@ def _needs_device_weight_postprocess(
|
|||||||
) -> bool:
|
) -> bool:
|
||||||
"""Return whether post-load weight processing needs CUDA/NPU tensors."""
|
"""Return whether post-load weight processing needs CUDA/NPU tensors."""
|
||||||
quant_name = _get_quant_config_name(quant_config)
|
quant_name = _get_quant_config_name(quant_config)
|
||||||
if quant_name == "modelopt_fp8":
|
if quant_name in ("modelopt_fp8", "comfy_fp8"):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
serialized_flag_by_quant_name = {
|
serialized_flag_by_quant_name = {
|
||||||
|
|||||||
@@ -193,8 +193,8 @@ def _copy_grouped_qkv_tp_shard(
|
|||||||
or getattr(param, "output_dim", None) != 0
|
or getattr(param, "output_dim", None) != 0
|
||||||
or getattr(param, "is_sharded_weight", False)
|
or getattr(param, "is_sharded_weight", False)
|
||||||
or getattr(param, "packed_dim", None) is not None
|
or getattr(param, "packed_dim", None) is not None
|
||||||
or param.dtype != _BF16_DTYPE
|
or param.dtype != loaded_weight.dtype
|
||||||
or loaded_weight.dtype != _BF16_DTYPE
|
or param.dtype not in (_BF16_DTYPE, torch.float8_e4m3fn)
|
||||||
or not param.is_contiguous()
|
or not param.is_contiguous()
|
||||||
or not loaded_weight.is_contiguous()
|
or not loaded_weight.is_contiguous()
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -97,36 +97,35 @@ def test_native_weight_names_and_grouped_qkv_reorder():
|
|||||||
).reshape(12, 1)
|
).reshape(12, 1)
|
||||||
torch.testing.assert_close(actual, expected)
|
torch.testing.assert_close(actual, expected)
|
||||||
|
|
||||||
grouped = torch.arange(48, dtype=torch.int16).reshape(24, 2).view(torch.bfloat16)
|
for dtype in (torch.bfloat16, torch.float8_e4m3fn):
|
||||||
reordered = _reorder_grouped_qkv_to_qkv(
|
grouped = torch.arange(48, dtype=torch.float32).reshape(24, 2).to(dtype)
|
||||||
grouped,
|
reordered = _reorder_grouped_qkv_to_qkv(
|
||||||
num_query_groups=4,
|
grouped,
|
||||||
heads_per_group=1,
|
num_query_groups=4,
|
||||||
head_dim=2,
|
heads_per_group=1,
|
||||||
)
|
head_dim=2,
|
||||||
for tp_size in (1, 2, 4):
|
)
|
||||||
local_rows = 8 // tp_size
|
for tp_size in (1, 2, 4):
|
||||||
for tp_rank in range(tp_size):
|
local_rows = 8 // tp_size
|
||||||
start = tp_rank * local_rows
|
for tp_rank in range(tp_size):
|
||||||
param = torch.nn.Parameter(
|
start = tp_rank * local_rows
|
||||||
torch.empty(3 * local_rows, 2, dtype=torch.bfloat16),
|
param = torch.nn.Parameter(
|
||||||
requires_grad=False,
|
torch.empty(3 * local_rows, 2, dtype=dtype),
|
||||||
)
|
requires_grad=False,
|
||||||
param.output_dim = 0
|
)
|
||||||
assert _copy_grouped_qkv_tp_shard(
|
param.output_dim = 0
|
||||||
param,
|
assert _copy_grouped_qkv_tp_shard(
|
||||||
grouped,
|
param,
|
||||||
num_query_groups=4,
|
grouped,
|
||||||
head_dim=2,
|
num_query_groups=4,
|
||||||
tp_rank=tp_rank,
|
head_dim=2,
|
||||||
tp_size=tp_size,
|
tp_rank=tp_rank,
|
||||||
)
|
tp_size=tp_size,
|
||||||
expected_shard = reordered.view(3, 8, 2)[
|
)
|
||||||
:, start : start + local_rows
|
expected_shard = reordered.view(3, 8, 2)[
|
||||||
].reshape(-1, 2)
|
:, start : start + local_rows
|
||||||
assert torch.equal(
|
].reshape(-1, 2)
|
||||||
param.view(torch.int16), expected_shard.view(torch.int16)
|
assert torch.equal(param, expected_shard)
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_pruned_adaln_curve_interpolates_without_timestep_mlp():
|
def test_pruned_adaln_curve_interpolates_without_timestep_mlp():
|
||||||
|
|||||||
@@ -45,11 +45,21 @@ sys.modules.setdefault(
|
|||||||
)
|
)
|
||||||
sys.modules.setdefault("partial_json_parser.core.options", partial_json_parser_options)
|
sys.modules.setdefault("partial_json_parser.core.options", partial_json_parser_options)
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
LinearBase,
|
||||||
|
UnquantizedLinearMethod,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.comfy_fp8 import (
|
||||||
|
ComfyFp8Config,
|
||||||
|
ComfyFullPrecisionFp8LinearMethod,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
||||||
NunchakuConfig,
|
NunchakuConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import (
|
||||||
|
Fp8Config,
|
||||||
|
Fp8LinearMethod,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
||||||
ModelOptFp4Config,
|
ModelOptFp4Config,
|
||||||
ModelOptFp8Config,
|
ModelOptFp8Config,
|
||||||
@@ -61,6 +71,10 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader i
|
|||||||
_resolve_checkpoint_load_device,
|
_resolve_checkpoint_load_device,
|
||||||
_warn_if_expected_param_dtype_missing,
|
_warn_if_expected_param_dtype_missing,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import (
|
||||||
|
inspect_minimax_h3_safetensors,
|
||||||
|
resolve_minimax_h3_checkpoint_quantization,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
||||||
TransformerQuantLoadSpec,
|
TransformerQuantLoadSpec,
|
||||||
_filter_duplicate_precision_variant_safetensors,
|
_filter_duplicate_precision_variant_safetensors,
|
||||||
@@ -119,6 +133,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
nunchaku_config=None,
|
nunchaku_config=None,
|
||||||
quantization=None,
|
quantization=None,
|
||||||
quantization_ignored_layers=None,
|
quantization_ignored_layers=None,
|
||||||
|
revision="test-revision",
|
||||||
tp_size=1,
|
tp_size=1,
|
||||||
dit_cpu_offload=False,
|
dit_cpu_offload=False,
|
||||||
direct_gpu_weight_loading=False,
|
direct_gpu_weight_loading=False,
|
||||||
@@ -170,6 +185,164 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(resolved, [f.name])
|
self.assertEqual(resolved, [f.name])
|
||||||
|
|
||||||
|
@patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.hf_hub_download",
|
||||||
|
return_value="/cache/model.safetensors",
|
||||||
|
)
|
||||||
|
def test_resolve_transformer_safetensors_to_load_uses_hf_file_reference(
|
||||||
|
self, mock_download
|
||||||
|
):
|
||||||
|
filename = "diffusion_models/minimax_h3_fl2va_pruned_bf16.safetensors"
|
||||||
|
references = (
|
||||||
|
(
|
||||||
|
f"https://huggingface.co/Comfy-Org/MiniMax-H3/resolve/main/{filename}",
|
||||||
|
"main",
|
||||||
|
),
|
||||||
|
(f"Comfy-Org/MiniMax-H3/{filename}", "test-revision"),
|
||||||
|
)
|
||||||
|
|
||||||
|
for reference, revision in references:
|
||||||
|
with self.subTest(reference=reference):
|
||||||
|
server_args = self._make_server_args(transformer_weights_path=reference)
|
||||||
|
with patch(
|
||||||
|
"os.path.isfile",
|
||||||
|
side_effect=lambda path: path == "/cache/model.safetensors",
|
||||||
|
):
|
||||||
|
self.assertEqual(
|
||||||
|
resolve_transformer_safetensors_to_load(
|
||||||
|
server_args, "/unused/component/path"
|
||||||
|
),
|
||||||
|
["/cache/model.safetensors"],
|
||||||
|
)
|
||||||
|
mock_download.assert_called_once_with(
|
||||||
|
repo_id="Comfy-Org/MiniMax-H3",
|
||||||
|
filename=filename,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
mock_download.reset_mock()
|
||||||
|
|
||||||
|
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self):
|
||||||
|
marker = json.dumps({"format": "int8_tensorwise", "convrot": True}).encode()
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
||||||
|
save_file(
|
||||||
|
{
|
||||||
|
"adaln_t_table": torch.zeros((1025, 8)),
|
||||||
|
"blocks.0.mlp.fc1.comfy_quant": torch.tensor(
|
||||||
|
list(marker), dtype=torch.uint8
|
||||||
|
),
|
||||||
|
},
|
||||||
|
f.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
curve_shape, comfy_quant = inspect_minimax_h3_safetensors([f.name])
|
||||||
|
|
||||||
|
self.assertEqual(curve_shape, (1025, 8))
|
||||||
|
self.assertEqual(comfy_quant["blocks.0.mlp.fc1"]["format"], "int8_tensorwise")
|
||||||
|
|
||||||
|
def test_inspect_minimax_h3_fp8_validates_required_scales(self):
|
||||||
|
marker = torch.tensor(list(b'{"format":"float8_e4m3fn"}'), dtype=torch.uint8)
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
||||||
|
save_file(
|
||||||
|
{
|
||||||
|
"blocks.0.mlp.fc1.weight": torch.ones(
|
||||||
|
(2, 2), dtype=torch.float8_e4m3fn
|
||||||
|
),
|
||||||
|
"blocks.0.mlp.fc1.weight_scale": torch.tensor(0.5),
|
||||||
|
"blocks.0.mlp.fc1.input_scale": torch.tensor(0.25),
|
||||||
|
"blocks.0.mlp.fc1.comfy_quant": marker,
|
||||||
|
},
|
||||||
|
f.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
_, layer_markers = inspect_minimax_h3_safetensors([f.name])
|
||||||
|
|
||||||
|
self.assertEqual(layer_markers["blocks.0.mlp.fc1"], {"format": "float8_e4m3fn"})
|
||||||
|
|
||||||
|
def test_minimax_h3_comfy_int8_fails_before_weight_loading(self):
|
||||||
|
with self.assertRaisesRegex(NotImplementedError, "regular-Hadamard"):
|
||||||
|
resolve_minimax_h3_checkpoint_quantization(
|
||||||
|
{
|
||||||
|
"blocks.0.mlp.fc1": {
|
||||||
|
"format": "int8_tensorwise",
|
||||||
|
"convrot": True,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_minimax_h3_comfy_fp8_resolves_per_layer_dispatch(self):
|
||||||
|
config = resolve_minimax_h3_checkpoint_quantization(
|
||||||
|
{
|
||||||
|
"blocks.0.attn.qkv_proj": {"format": "float8_e4m3fn"},
|
||||||
|
"blocks.0.mlp.fc2": {
|
||||||
|
"format": "float8_e4m3fn",
|
||||||
|
"full_precision_matrix_mult": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsInstance(config, ComfyFp8Config)
|
||||||
|
layer = LinearBase(input_size=1, output_size=1)
|
||||||
|
self.assertIsInstance(
|
||||||
|
config.get_quant_method(layer, "blocks.0.mlp.fc2"),
|
||||||
|
ComfyFullPrecisionFp8LinearMethod,
|
||||||
|
)
|
||||||
|
self.assertIsInstance(
|
||||||
|
config.get_quant_method(layer, "blocks.0.attn.qkv_proj"),
|
||||||
|
Fp8LinearMethod,
|
||||||
|
)
|
||||||
|
self.assertIsInstance(
|
||||||
|
config.get_quant_method(layer, "unmarked"),
|
||||||
|
UnquantizedLinearMethod,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_comfy_full_precision_fp8_dequantizes_before_linear(self):
|
||||||
|
layer = torch.nn.Module()
|
||||||
|
layer.weight = torch.nn.Parameter(
|
||||||
|
torch.tensor([[2.0, -4.0]], dtype=torch.float8_e4m3fn),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
layer.weight_scale = torch.nn.Parameter(
|
||||||
|
torch.tensor([0.5]), requires_grad=False
|
||||||
|
)
|
||||||
|
|
||||||
|
output = ComfyFullPrecisionFp8LinearMethod().apply(
|
||||||
|
layer, torch.tensor([[3.0, 1.0]])
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.testing.assert_close(output, torch.tensor([[1.0]]))
|
||||||
|
|
||||||
|
def test_checkpoint_quantization_metadata_drives_load_spec(self):
|
||||||
|
config = ComfyFp8Config({})
|
||||||
|
server_args = self._make_server_args()
|
||||||
|
|
||||||
|
spec = resolve_transformer_quant_load_spec(
|
||||||
|
hf_config={},
|
||||||
|
server_args=server_args,
|
||||||
|
safetensors_list=["model.safetensors"],
|
||||||
|
component_model_path="/unused/component/path",
|
||||||
|
model_cls=_FakeFluxTransformer,
|
||||||
|
cls_name=_FakeFluxTransformer.__name__,
|
||||||
|
checkpoint_quant_config=config,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(spec.quant_config, config)
|
||||||
|
self.assertTrue(spec.is_comfy_fp8)
|
||||||
|
self.assertTrue(spec.needs_device_weight_postprocess)
|
||||||
|
|
||||||
|
def test_checkpoint_quantization_rejects_explicit_quantization(self):
|
||||||
|
server_args = self._make_server_args(quantization="fp8")
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "per-layer metadata"):
|
||||||
|
resolve_transformer_quant_load_spec(
|
||||||
|
hf_config={},
|
||||||
|
server_args=server_args,
|
||||||
|
safetensors_list=["model.safetensors"],
|
||||||
|
component_model_path="/unused/component/path",
|
||||||
|
model_cls=_FakeFluxTransformer,
|
||||||
|
cls_name=_FakeFluxTransformer.__name__,
|
||||||
|
checkpoint_quant_config=ComfyFp8Config({}),
|
||||||
|
)
|
||||||
|
|
||||||
@patch(
|
@patch(
|
||||||
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model",
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model",
|
||||||
side_effect=lambda path, **kw: path,
|
side_effect=lambda path, **kw: path,
|
||||||
@@ -355,6 +528,11 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_comfy_fp8_needs_device_weight_postprocess(self):
|
||||||
|
self.assertTrue(
|
||||||
|
_needs_device_weight_postprocess(_make_quant_config("comfy_fp8"))
|
||||||
|
)
|
||||||
|
|
||||||
def test_online_fp8_receives_cli_ignored_layer_patterns(self):
|
def test_online_fp8_receives_cli_ignored_layer_patterns(self):
|
||||||
ignored_layers = ["blocks.0.attn.out_proj", "condition_proj"]
|
ignored_layers = ["blocks.0.attn.out_proj", "condition_proj"]
|
||||||
server_args = self._make_server_args(
|
server_args = self._make_server_args(
|
||||||
|
|||||||
Reference in New Issue
Block a user