[diffusion] feat: support loading native diffusers miniMax h3 components (#36067)

This commit is contained in:
Mick
2026-08-23 22:13:59 +08:00
committed by GitHub
parent de6a1dbd7a
commit d1af3c8923
4 changed files with 130 additions and 25 deletions
@@ -16,51 +16,59 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
r"^(.*\.lora_[AB])\.[^.]+$": r"\1", r"^(.*\.lora_[AB])\.[^.]+$": r"\1",
r"^base_model\.model\.(.*\.lora_[AB])$": r"\1", r"^base_model\.model\.(.*\.lora_[AB])$": r"\1",
r"^transformer\.(.*\.lora_[AB])$": r"\1", r"^transformer\.(.*\.lora_[AB])$": r"\1",
r"^proj_in\.(lora_[AB])$": r"video_patch_proj.\1", r"^proj_in\.(.*)$": r"video_patch_proj.\1",
r"^audio_proj_in\.(lora_[AB])$": r"audio_patch_proj.\1", r"^audio_proj_in\.(.*)$": r"audio_patch_proj.\1",
r"^context_embedder\.(lora_[AB])$": r"condition_proj.\1", r"^context_embedder\.(.*)$": r"condition_proj.\1",
r"^time_embedder\.linear_1\.(lora_[AB])$": r"time_embedder.proj_in.\1", r"^time_embedder\.linear_1\.(.*)$": r"time_embedder.proj_in.\1",
r"^time_embedder\.linear_2\.(lora_[AB])$": r"time_embedder.proj_out.\1", r"^time_embedder\.linear_2\.(.*)$": r"time_embedder.proj_out.\1",
r"^norm_out\.linear\.(lora_[AB])$": r"final_layer.adaln_proj.linear.\1", r"^norm_out\.norm\.(.*)$": r"final_layer.norm.\1",
r"^proj_out\.(lora_[AB])$": r"final_layer.video_out.\1", r"^norm_out\.linear\.(.*)$": r"final_layer.adaln_proj.linear.\1",
r"^audio_proj_out\.(lora_[AB])$": r"final_layer.audio_out.\1", r"^proj_out\.(.*)$": r"final_layer.video_out.\1",
r"^transformer_blocks\.(\d+)\.adaln_proj\.linear\.(lora_[AB])$": r"blocks.\1.adaln_proj.linear.\2", r"^audio_proj_out\.(.*)$": r"final_layer.audio_out.\1",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(lora_[AB])$": ( r"^transformer_blocks\.(\d+)\.adaln_proj\.linear\.(.*)$": r"blocks.\1.adaln_proj.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (
r"blocks.\1.attn.qkv_proj.\2", r"blocks.\1.attn.qkv_proj.\2",
0, 0,
3, 3,
), ),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(lora_[AB])$": ( r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (
r"blocks.\1.attn.qkv_proj.\2", r"blocks.\1.attn.qkv_proj.\2",
1, 1,
3, 3,
), ),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(lora_[AB])$": ( r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (
r"blocks.\1.attn.qkv_proj.\2", r"blocks.\1.attn.qkv_proj.\2",
2, 2,
3, 3,
), ),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(lora_[AB])$": r"blocks.\1.attn.out_proj.\2", r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"blocks.\1.attn.out_proj.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(lora_[AB])$": r"blocks.\1.mlp.fc1.\2", r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"blocks.\1.attn.q_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(lora_[AB])$": r"blocks.\1.mlp.fc2.\2", r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"blocks.\1.attn.k_norm.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(lora_[AB])$": ( r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"blocks.\1.mlp.fc1.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"blocks.\1.mlp.fc2.\2",
r"^transformer_blocks\.(\d+)\.norm([12])\.(.*)$": r"blocks.\1.norm\2.\3",
r"^token_refiner\.final_norm\.(.*)$": r"token_refiner.final_norm.\1",
r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$": (
r"token_refiner.blocks.\1.attn.qkv_proj.\2", r"token_refiner.blocks.\1.attn.qkv_proj.\2",
0, 0,
3, 3,
), ),
r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(lora_[AB])$": ( r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$": (
r"token_refiner.blocks.\1.attn.qkv_proj.\2", r"token_refiner.blocks.\1.attn.qkv_proj.\2",
1, 1,
3, 3,
), ),
r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(lora_[AB])$": ( r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$": (
r"token_refiner.blocks.\1.attn.qkv_proj.\2", r"token_refiner.blocks.\1.attn.qkv_proj.\2",
2, 2,
3, 3,
), ),
r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(lora_[AB])$": r"token_refiner.blocks.\1.attn.out_proj.\2", r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"token_refiner.blocks.\1.attn.out_proj.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0\.proj\.(lora_[AB])$": r"token_refiner.blocks.\1.mlp.fc1.\2", r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"token_refiner.blocks.\1.attn.q_norm.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2\.(lora_[AB])$": r"token_refiner.blocks.\1.mlp.fc2.\2", r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"token_refiner.blocks.\1.attn.k_norm.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"token_refiner.blocks.\1.mlp.fc1.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"token_refiner.blocks.\1.mlp.fc2.\2",
r"^token_refiner\.refiner_blocks\.(\d+)\.norm([12])\.(.*)$": r"token_refiner.blocks.\1.norm\2.\3",
} }
) )
@@ -85,6 +93,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
norm_eps: float = 1e-5 norm_eps: float = 1e-5
qk_norm_eps: float = 1e-5 qk_norm_eps: float = 1e-5
final_norm_eps: float = 1e-5 final_norm_eps: float = 1e-5
checkpoint_uses_diffusers_layout: bool = False
def __post_init__(self) -> None: def __post_init__(self) -> None:
super().__post_init__() super().__post_init__()
@@ -99,6 +108,20 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
class MiniMaxH3DiTConfig(DiTConfig): class MiniMaxH3DiTConfig(DiTConfig):
arch_config: MiniMaxH3DiTArchConfig = field(default_factory=MiniMaxH3DiTArchConfig) arch_config: MiniMaxH3DiTArchConfig = field(default_factory=MiniMaxH3DiTArchConfig)
def update_model_arch(self, source_model_dict: dict) -> None:
aliases = {
"num_refiner_layers": "token_refiner_num_layers",
"ffn_dim": "ffn_hidden_size",
"in_channels": "latents_dim",
"audio_in_channels": "audio_latents_dim",
"freq_dim": "timestep_input_dim",
"time_embed_hidden_dim": "time_embed_hidden_size",
"rope_freq_dim": "rope_inv_freq_len",
}
super().update_model_arch(
{aliases.get(key, key): value for key, value in source_model_dict.items()}
)
__all__ = [ __all__ = [
"MINIMAX_H3_ADALN_MODALITY_NUM", "MINIMAX_H3_ADALN_MODALITY_NUM",
@@ -219,9 +219,14 @@ 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)
is_minimax_h3 = model_cls.__name__ == "MiniMaxH3DiTModel"
if is_minimax_h3:
dit_config.arch_config.checkpoint_uses_diffusers_layout = (
cls_name == "MiniMaxH3Transformer3DModel"
)
checkpoint_quant_config = None checkpoint_quant_config = None
if cls_name == "MiniMaxH3DiTModel": if is_minimax_h3:
selected_variant = str(component_server_args.model_variant or "fl2va") selected_variant = str(component_server_args.model_variant or "fl2va")
if gguf_file is not None: if gguf_file is not None:
validate_minimax_h3_checkpoint_variant([gguf_file], selected_variant) validate_minimax_h3_checkpoint_variant([gguf_file], selected_variant)
@@ -260,7 +265,7 @@ class TransformerLoader(ComponentLoader):
gguf_file=gguf_file, gguf_file=gguf_file,
checkpoint_quant_config=checkpoint_quant_config, 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 is_minimax_h3:
assert quant_spec.quant_config is not None assert quant_spec.quant_config is not None
curve = quant_spec.quant_config.tensor_meta.get("adaln_t_table") curve = quant_spec.quant_config.tensor_meta.get("adaln_t_table")
if curve is not None: if curve is not None:
@@ -312,7 +317,7 @@ class TransformerLoader(ComponentLoader):
) )
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 not is_minimax_h3:
raise ValueError( raise ValueError(
"--minimax-h3-adaln-cache-path is only supported by MiniMax H3" "--minimax-h3-adaln-cache-path is only supported by MiniMax H3"
) )
@@ -326,7 +331,7 @@ class TransformerLoader(ComponentLoader):
) )
checkpoint_key_filter = _minimax_h3_adaln_cache_key_filter checkpoint_key_filter = _minimax_h3_adaln_cache_key_filter
if component_server_args.minimax_h3_adaln_online: if component_server_args.minimax_h3_adaln_online:
if cls_name != "MiniMaxH3DiTModel": if not is_minimax_h3:
raise ValueError( raise ValueError(
"--minimax-h3-adaln-online is only supported by MiniMax H3" "--minimax-h3-adaln-online is only supported by MiniMax H3"
) )
@@ -10,6 +10,8 @@ from __future__ import annotations
import math import math
import os import os
import struct import struct
from collections import defaultdict
from collections.abc import Iterable, Iterator
from contextlib import ExitStack from contextlib import ExitStack
from typing import Any, Callable from typing import Any, Callable
@@ -58,6 +60,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig, QuantizationConfig,
) )
from sglang.multimodal_gen.runtime.layers.usp import _ring_attention_varlen from sglang.multimodal_gen.runtime.layers.usp import _ring_attention_varlen
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin, LayerwiseOffloadableModuleMixin,
is_layerwise_offloaded_module, is_layerwise_offloaded_module,
@@ -75,6 +78,48 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import
logger = init_logger(__name__) logger = init_logger(__name__)
_ARCH_DEFAULTS = MiniMaxH3DiTArchConfig() _ARCH_DEFAULTS = MiniMaxH3DiTArchConfig()
def _diffusers_h3_checkpoint(
iterator: Iterable[tuple[str, torch.Tensor]],
) -> Iterator[tuple[str, torch.Tensor]]:
"""Map Diffusers H3 names/layout to the fused native checkpoint layout."""
mapping = get_param_names_mapping(_ARCH_DEFAULTS.param_names_mapping)
pending: dict[str, dict[int, torch.Tensor]] = defaultdict(dict)
for source_name, tensor in iterator:
target_name, merge_index, merge_count = mapping(source_name)
# Diffusers SwiGLU stores [value, gate]; the native fused MLP consumes
# [gate, value]. Packed GPTQ tensors carry output channels on dim 1.
if ".ff.net.0.proj." in source_name:
output_dim = (
1 if source_name.endswith((".qweight", ".qzeros", ".scales")) else 0
)
value, gate = tensor.chunk(2, dim=output_dim)
tensor = torch.cat((gate, value), dim=output_dim)
if merge_index is None:
yield target_name, tensor
continue
assert merge_count is not None
pending[target_name][merge_index] = tensor
if len(pending[target_name]) != merge_count:
continue
merge_dim = 1 if target_name.endswith((".qweight", ".qzeros", ".scales")) else 0
yield target_name, torch.cat(
[pending[target_name][index] for index in range(merge_count)],
dim=merge_dim,
)
del pending[target_name]
if pending:
incomplete = ", ".join(sorted(pending))
raise ValueError(f"Incomplete Diffusers H3 fused parameters: {incomplete}")
_BF16_DTYPE = torch.bfloat16 _BF16_DTYPE = torch.bfloat16
_FP32_DTYPE = torch.float32 _FP32_DTYPE = torch.float32
_MPS_MLP_TOKEN_CHUNK_SIZE = 128 _MPS_MLP_TOKEN_CHUNK_SIZE = 128
@@ -618,6 +663,9 @@ class MiniMaxH3Attention(nn.Module):
quant_config.get_name() == "gguf" quant_config.get_name() == "gguf"
or quant_config.checkpoint_uses_native_qkv_layout or quant_config.checkpoint_uses_native_qkv_layout
) )
checkpoint_qkv_is_native = (
checkpoint_qkv_is_native or arch.checkpoint_uses_diffusers_layout
)
if not checkpoint_qkv_is_native: 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)
@@ -1605,6 +1653,7 @@ class MiniMaxH3FinalLayer(nn.Module):
class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin): class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
_aliases = ["MiniMaxH3Transformer3DModel"]
_fsdp_shard_conditions = [is_block] _fsdp_shard_conditions = [is_block]
# refine_prompt_embeds drives a forward pass outside __call__. # refine_prompt_embeds drives a forward pass outside __call__.
_fsdp_forward_methods = ("refine_prompt_embeds",) _fsdp_forward_methods = ("refine_prompt_embeds",)
@@ -1735,6 +1784,8 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
adaln_cache_path is not None or adaln_weight_files is not None adaln_cache_path is not None or adaln_weight_files is not None
) )
self.arch = arch self.arch = arch
if arch.checkpoint_uses_diffusers_layout:
self.preprocess_loaded_state_dict = _diffusers_h3_checkpoint
self.hidden_size = arch.hidden_size self.hidden_size = arch.hidden_size
self.num_attention_heads = arch.num_attention_heads self.num_attention_heads = arch.num_attention_heads
self.num_channels_latents = arch.latents_dim self.num_channels_latents = arch.latents_dim
@@ -28,6 +28,7 @@ from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import (
MiniMaxH3DiTBlock, MiniMaxH3DiTBlock,
MiniMaxH3DiTModel, MiniMaxH3DiTModel,
_copy_grouped_qkv_tp_shard, _copy_grouped_qkv_tp_shard,
_diffusers_h3_checkpoint,
_modulate_gate, _modulate_gate,
_reorder_grouped_qkv_to_qkv, _reorder_grouped_qkv_to_qkv,
) )
@@ -84,6 +85,31 @@ def test_native_weight_names_and_grouped_qkv_reorder():
None, None,
) )
assert mapping("transformer_blocks.7.attn.to_k.qweight") == (
"blocks.7.attn.qkv_proj.qweight",
1,
3,
)
diffusers_weights = [
("transformer_blocks.0.attn.to_q.qweight", torch.full((2, 3), 1)),
("transformer_blocks.0.attn.to_k.qweight", torch.full((2, 3), 2)),
("transformer_blocks.0.attn.to_v.qweight", torch.full((2, 3), 3)),
(
"transformer_blocks.0.ff.net.0.proj.weight",
torch.arange(8).reshape(4, 2),
),
]
converted = dict(_diffusers_h3_checkpoint(diffusers_weights))
assert torch.equal(
converted["blocks.0.attn.qkv_proj.qweight"],
torch.cat([tensor for _, tensor in diffusers_weights[:3]], dim=1),
)
assert torch.equal(
converted["blocks.0.mlp.fc1.weight"],
torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]),
)
weight = torch.arange(12, dtype=torch.float32).reshape(12, 1) weight = torch.arange(12, dtype=torch.float32).reshape(12, 1)
actual = _reorder_grouped_qkv_to_qkv( actual = _reorder_grouped_qkv_to_qkv(
weight, weight,