[diffusion] model: support MiniMax-H3 singularity hybrid checkpoints (#38455)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
@@ -38,7 +38,7 @@ def minimax_h3_tasks_for_server(server_args: "ServerArgs") -> tuple[str, ...]:
|
||||
return tuple(
|
||||
task
|
||||
for task, task_partition in MINIMAX_H3_TASK_PARTITIONS.items()
|
||||
if task_partition == partition
|
||||
if partition == "hybrid" or task_partition == partition
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
|
||||
# H3 fuses Q/K/V, so split projections are stacked for the fused LoRA layer
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^model\.diffusion_model\.(.*)$": r"\1",
|
||||
r"^(.*)\.weight_scale$": r"\1.weight_scale_inv",
|
||||
r"^(.*\.lora_[AB])\.[^.]+$": r"\1",
|
||||
r"^base_model\.model\.(.*\.lora_[AB])$": r"\1",
|
||||
|
||||
@@ -370,7 +370,9 @@ def maybe_load_fsdp_model(
|
||||
)
|
||||
register_fsdp_entrypoints(model)
|
||||
|
||||
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
|
||||
param_names_mapping_fn = get_param_names_mapping(
|
||||
model.param_names_mapping, valid_target_names=set(model.state_dict())
|
||||
)
|
||||
|
||||
# 2. load model from disk
|
||||
preprocess_loaded_state_dict = getattr(model, "preprocess_loaded_state_dict", None)
|
||||
|
||||
@@ -5,12 +5,14 @@ from typing import Any
|
||||
|
||||
from safetensors import safe_open
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.minimax_h3 import MiniMaxH3DiTArchConfig
|
||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||
QuantizationConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
||||
ModelOptFp4Config,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
|
||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||
build_nvfp4_config_from_safetensors_list,
|
||||
inspect_comfy_quant_markers,
|
||||
@@ -27,13 +29,20 @@ def inspect_minimax_h3_safetensors(
|
||||
) -> 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 = inspect_comfy_quant_markers(safetensors_list)
|
||||
mapping = get_param_names_mapping(MiniMaxH3DiTArchConfig().param_names_mapping)
|
||||
layer_markers = inspect_comfy_quant_markers(
|
||||
safetensors_list,
|
||||
param_name_mapper=lambda prefix: mapping(f"{prefix}.weight")[0].removesuffix(
|
||||
".weight"
|
||||
),
|
||||
)
|
||||
|
||||
for path in safetensors_list:
|
||||
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
||||
keys = checkpoint.keys()
|
||||
if "adaln_t_table" in keys:
|
||||
shape = tuple(checkpoint.get_slice("adaln_t_table").get_shape())
|
||||
for key in checkpoint.keys():
|
||||
if mapping(key)[0] != "adaln_t_table":
|
||||
continue
|
||||
shape = tuple(checkpoint.get_slice(key).get_shape())
|
||||
if len(shape) != 2 or shape[0] < 2:
|
||||
raise ValueError(
|
||||
"MiniMax-H3 adaln_t_table must have shape [N, D] with "
|
||||
@@ -84,6 +93,9 @@ def resolve_minimax_h3_checkpoint_quantization(
|
||||
def validate_minimax_h3_checkpoint_variant(
|
||||
checkpoint_paths: list[str], selected_variant: str
|
||||
) -> None:
|
||||
# hybrid is an explicit cross-partition override, not a filename heuristic
|
||||
if selected_variant.strip().lower() == "hybrid":
|
||||
return
|
||||
names = " ".join(path.lower() for path in checkpoint_paths)
|
||||
checkpoint_variants = {
|
||||
variant for variant in ("fl2va", "ref2va") if variant in names
|
||||
|
||||
@@ -121,12 +121,15 @@ def load_model_state_dict(
|
||||
|
||||
def get_param_names_mapping(
|
||||
mapping_dict: dict[str, str | tuple[str, int, int]],
|
||||
valid_target_names: set[str] | None = None,
|
||||
) -> Callable[[str], tuple[str, Any, Any]]:
|
||||
"""
|
||||
Creates a mapping function that transforms parameter names using regex patterns.
|
||||
|
||||
Args:
|
||||
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
|
||||
valid_target_names: Keep a valid intermediate mapping when a later
|
||||
alias does not exist in the constructed model.
|
||||
|
||||
Returns:
|
||||
Callable[[str], str]: A function that maps parameter names from source to target format
|
||||
@@ -140,6 +143,7 @@ def get_param_names_mapping(
|
||||
max_steps = max(8, len(mapping_dict) * 2)
|
||||
applied_patterns: set[str] = set()
|
||||
visited_names: set[str] = {name}
|
||||
valid_mapping = None
|
||||
|
||||
for _ in range(max_steps):
|
||||
transformed = False
|
||||
@@ -166,6 +170,8 @@ def get_param_names_mapping(
|
||||
|
||||
name = new_name
|
||||
applied_patterns.add(pattern)
|
||||
if valid_target_names is not None and name in valid_target_names:
|
||||
valid_mapping = (name, merge_index, total_split_params)
|
||||
if name in visited_names:
|
||||
transformed = False
|
||||
break
|
||||
@@ -176,6 +182,16 @@ def get_param_names_mapping(
|
||||
if not transformed:
|
||||
break
|
||||
|
||||
# Prefer the complete mapping. If a later alias does not exist in this
|
||||
# model (e.g. INT8 weight_scale -> FP8 weight_scale_inv), retain the
|
||||
# last valid intermediate name, including any required QKV merge.
|
||||
if (
|
||||
name
|
||||
and valid_mapping is not None
|
||||
and valid_target_names is not None
|
||||
and name not in valid_target_names
|
||||
):
|
||||
return valid_mapping
|
||||
return name, merge_index, total_split_params
|
||||
|
||||
return mapping_fn
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
from dataclasses import replace
|
||||
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import (
|
||||
MiniMaxH3PipelineConfig,
|
||||
@@ -26,6 +27,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.m
|
||||
MiniMaxH3PartitionAdmissionStage,
|
||||
MiniMaxH3ReleaseMetadata,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.task_profiles import (
|
||||
MINIMAX_H3_TASK_PARTITIONS,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
@@ -75,6 +79,7 @@ class MiniMaxH3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
subfolders = {
|
||||
"fl2va": "FL2VA",
|
||||
"ref2va": "Ref2VA",
|
||||
"hybrid": "Ref2VA",
|
||||
}
|
||||
try:
|
||||
return subfolders[normalized]
|
||||
@@ -88,6 +93,14 @@ class MiniMaxH3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
model_variant = self.server_args.model_variant
|
||||
if model_variant is not None:
|
||||
semantic_subfolder = self.model_subfolder_for_variant(model_variant)
|
||||
if model_variant.strip().lower() == "hybrid" and not (
|
||||
self.server_args.component_weights_paths.get("transformer")
|
||||
or self.server_args.transformer_weights_path
|
||||
):
|
||||
raise ValueError(
|
||||
"MiniMax H3 --model-variant hybrid requires explicit merged "
|
||||
"weights via --component-weights-paths.transformer"
|
||||
)
|
||||
explicit_subfolder = self.server_args.model_subfolder
|
||||
if (
|
||||
explicit_subfolder is not None
|
||||
@@ -104,12 +117,20 @@ class MiniMaxH3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.release_metadata = MiniMaxH3ReleaseMetadata.from_model_index(model_index)
|
||||
if (
|
||||
model_variant is not None
|
||||
and self.release_metadata.partition != model_variant.strip().lower()
|
||||
and self.release_metadata.partition != semantic_subfolder.lower()
|
||||
):
|
||||
raise ValueError(
|
||||
"MiniMax H3 loaded checkpoint partition does not match "
|
||||
f"--model-variant {model_variant!r}"
|
||||
)
|
||||
if model_variant is not None and model_variant.strip().lower() == "hybrid":
|
||||
# merged checkpoints share the native graph across all three tasks
|
||||
# keep the base partition contract strict unless explicitly selected
|
||||
self.release_metadata = replace(
|
||||
self.release_metadata,
|
||||
partition="hybrid",
|
||||
tasks=tuple(MINIMAX_H3_TASK_PARTITIONS),
|
||||
)
|
||||
return model_index
|
||||
|
||||
def validate_disagg_role(self, role: RoleType) -> None:
|
||||
|
||||
+4
-1
@@ -131,7 +131,10 @@ class MiniMaxH3ReleaseMetadata:
|
||||
f"task {task!r} is not served by MiniMax H3 partition {self.partition!r}; "
|
||||
f"supported tasks: {list(self.tasks)!r}"
|
||||
)
|
||||
if partition_for_task(canonical) != self.partition:
|
||||
if (
|
||||
self.partition != "hybrid"
|
||||
and partition_for_task(canonical) != self.partition
|
||||
):
|
||||
raise ValueError(
|
||||
f"task {task!r} resolves outside partition {self.partition!r}"
|
||||
)
|
||||
|
||||
@@ -20,10 +20,19 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionRequirements,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import (
|
||||
validate_minimax_h3_checkpoint_variant,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||
LAYERWISE_OFFLOAD,
|
||||
RESIDENT,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
|
||||
MiniMaxH3Pipeline,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.release_metadata import (
|
||||
MiniMaxH3PartitionAdmissionStage,
|
||||
@@ -206,6 +215,50 @@ def test_loaded_weight_partition_admits_only_its_declared_tasks(partition, tasks
|
||||
metadata.canonical_task(rejected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("weights", [None, "owner/repo/merged_ref2va_int8.safetensors"])
|
||||
def test_hybrid_override_loads_ref_config_and_admits_all_native_tasks(weights):
|
||||
model_index = {
|
||||
"_minimax_h3": {
|
||||
"schema_version": 1,
|
||||
"partition": "ref2va",
|
||||
"tasks": ["ref2va"],
|
||||
"sigma_shift_scales": {"video": 12.0, "audio": 3.0},
|
||||
}
|
||||
}
|
||||
pipeline = MiniMaxH3Pipeline.__new__(MiniMaxH3Pipeline)
|
||||
pipeline.server_args = SimpleNamespace(
|
||||
model_variant="hybrid",
|
||||
model_subfolder=None,
|
||||
component_weights_paths={"transformer": weights} if weights else {},
|
||||
transformer_weights_path=None,
|
||||
)
|
||||
with patch.object(
|
||||
ComposedPipelineBase, "_load_config", return_value=model_index
|
||||
) as load:
|
||||
if weights is None:
|
||||
with pytest.raises(ValueError, match="requires explicit merged weights"):
|
||||
pipeline._load_config()
|
||||
load.assert_not_called()
|
||||
return
|
||||
pipeline._load_config()
|
||||
|
||||
assert pipeline.server_args.model_subfolder == "Ref2VA"
|
||||
assert model_index["_minimax_h3"]["tasks"] == ["ref2va"]
|
||||
validate_minimax_h3_checkpoint_variant([weights], "hybrid")
|
||||
stage = MiniMaxH3PartitionAdmissionStage(pipeline.release_metadata)
|
||||
for task in ("t2va", "fl2va", "ref2va"):
|
||||
batch = SimpleNamespace(
|
||||
sampling_params=SimpleNamespace(task=task, quality="lossless"),
|
||||
num_inference_steps=50,
|
||||
)
|
||||
assert (
|
||||
stage.forward(batch, SimpleNamespace(minimax_h3_adaln_online=False))
|
||||
is batch
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
pipeline.release_metadata.canonical_task("unknown")
|
||||
|
||||
|
||||
def test_synthetic_warmup_target_honors_warmup_flags():
|
||||
def target(num_frames=None, resolution=None):
|
||||
width, height = map(int, (resolution or "896x512").split("x"))
|
||||
|
||||
@@ -13,6 +13,8 @@ from unittest.mock import patch
|
||||
import torch
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.minimax_h3 import MiniMaxH3DiTArchConfig
|
||||
|
||||
partial_json_parser = types.ModuleType("partial_json_parser")
|
||||
partial_json_parser_core = types.ModuleType("partial_json_parser.core")
|
||||
partial_json_parser_exceptions = types.ModuleType("partial_json_parser.core.exceptions")
|
||||
@@ -704,7 +706,11 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
||||
["/cache/minimax_h3_fl2va.safetensors"], "ref2va"
|
||||
)
|
||||
|
||||
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self):
|
||||
@patch(
|
||||
"sglang.multimodal_gen.runtime.layers.quantization.kitchen_int8."
|
||||
"_load_comfy_kitchen"
|
||||
)
|
||||
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self, _load):
|
||||
marker = json.dumps(
|
||||
{
|
||||
"format": "int8_tensorwise",
|
||||
@@ -712,23 +718,80 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
||||
"convrot_groupsize": 256,
|
||||
}
|
||||
).encode()
|
||||
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
||||
save_file(
|
||||
{
|
||||
"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(
|
||||
mapping = get_param_names_mapping(MiniMaxH3DiTArchConfig().param_names_mapping)
|
||||
for prefix in ("", "model.diffusion_model."):
|
||||
with (
|
||||
self.subTest(prefix=prefix),
|
||||
tempfile.NamedTemporaryFile(suffix=".safetensors") as f,
|
||||
):
|
||||
weights = {
|
||||
prefix + "adaln_t_table": torch.zeros((1025, 8)),
|
||||
prefix + "blocks.0.mlp.fc1.weight": torch.ones(
|
||||
(2, 256), dtype=torch.int8
|
||||
),
|
||||
prefix + "blocks.0.mlp.fc1.weight_scale": torch.ones((2, 1)),
|
||||
prefix + "blocks.0.mlp.fc1.comfy_quant": torch.tensor(
|
||||
list(marker), dtype=torch.uint8
|
||||
),
|
||||
},
|
||||
f.name,
|
||||
)
|
||||
}
|
||||
save_file(weights, f.name)
|
||||
curve_shape, comfy_quant = inspect_minimax_h3_safetensors([f.name])
|
||||
config = resolve_minimax_h3_checkpoint_quantization(comfy_quant)
|
||||
layer = ReplicatedLinear(
|
||||
256,
|
||||
2,
|
||||
bias=False,
|
||||
params_dtype=torch.bfloat16,
|
||||
quant_config=config,
|
||||
prefix="blocks.0.mlp.fc1",
|
||||
)
|
||||
target_names = {
|
||||
"blocks.0.mlp.fc1." + name for name in layer.state_dict()
|
||||
} | {"adaln_t_table"}
|
||||
native_mapping = get_param_names_mapping(
|
||||
MiniMaxH3DiTArchConfig().param_names_mapping,
|
||||
valid_target_names=target_names,
|
||||
)
|
||||
mapped, _ = hf_to_custom_state_dict(
|
||||
weights, native_mapping, valid_target_names=target_names
|
||||
)
|
||||
|
||||
curve_shape, comfy_quant = inspect_minimax_h3_safetensors([f.name])
|
||||
self.assertEqual(curve_shape, (1025, 8))
|
||||
self.assertIn("adaln_t_table", mapped)
|
||||
self.assertEqual(set(config.layer_markers), {"blocks.0.mlp.fc1"})
|
||||
self.assertEqual(mapped["blocks.0.mlp.fc1.weight"].dtype, torch.int8)
|
||||
self.assertIn("blocks.0.mlp.fc1.weight_scale", mapped)
|
||||
layer.load_state_dict(
|
||||
{
|
||||
name.removeprefix("blocks.0.mlp.fc1."): tensor
|
||||
for name, tensor in mapped.items()
|
||||
if name in target_names and name != "adaln_t_table"
|
||||
},
|
||||
strict=True,
|
||||
)
|
||||
fp8_mapped, _ = hf_to_custom_state_dict(weights, mapping)
|
||||
self.assertIn("blocks.0.mlp.fc1.weight_scale_inv", fp8_mapped)
|
||||
both_scales = target_names | {"blocks.0.mlp.fc1.weight_scale_inv"}
|
||||
final_mapping = get_param_names_mapping(
|
||||
MiniMaxH3DiTArchConfig().param_names_mapping, both_scales
|
||||
)
|
||||
self.assertEqual(
|
||||
final_mapping(prefix + "blocks.0.mlp.fc1.weight_scale")[0],
|
||||
"blocks.0.mlp.fc1.weight_scale_inv",
|
||||
)
|
||||
|
||||
self.assertEqual(curve_shape, (1025, 8))
|
||||
self.assertEqual(comfy_quant["blocks.0.mlp.fc1"]["format"], "int8_tensorwise")
|
||||
def test_mapping_fallback_preserves_merge_and_explicit_drop(self):
|
||||
mapping = get_param_names_mapping(
|
||||
{
|
||||
r"^wrapper\.(.*)$": r"\1",
|
||||
r"^q\.(.*)$": (r"qkv.\1", 0, 3),
|
||||
r"^qkv\.weight_scale$": "qkv.weight_scale_inv",
|
||||
r"^ignored$": "",
|
||||
},
|
||||
{"qkv.weight_scale", "ignored"},
|
||||
)
|
||||
self.assertEqual(mapping("wrapper.q.weight_scale"), ("qkv.weight_scale", 0, 3))
|
||||
self.assertEqual(mapping("wrapper.ignored"), ("", None, None))
|
||||
|
||||
def test_inspect_minimax_h3_fp8_detects_static_activation_scale(self):
|
||||
marker = torch.tensor(list(b'{"format":"float8_e4m3fn"}'), dtype=torch.uint8)
|
||||
|
||||
@@ -44,6 +44,9 @@ def test_h3_webui_tasks_follow_loaded_partition():
|
||||
assert minimax_h3_tasks_for_server(
|
||||
SimpleNamespace(model_variant="ref2va", model_subfolder=None)
|
||||
) == ("ref2va",)
|
||||
assert minimax_h3_tasks_for_server(
|
||||
SimpleNamespace(model_variant="hybrid", model_subfolder=None)
|
||||
) == ("t2va", "fl2va", "ref2va")
|
||||
|
||||
|
||||
def test_h3_t2va_uses_native_contract_without_generic_cfg_fields():
|
||||
|
||||
Reference in New Issue
Block a user