From 51b27f747ade4a267f37a29457bc8c690100ca9e Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 24 Aug 2026 16:39:37 +0800 Subject: [PATCH] [diffusion] feat: support loading pruned minimax h3 components natively (#36070) --- .../configs/models/dits/minimax_h3.py | 15 ++- .../runtime/layers/lora/linear.py | 107 ++++++++++++++++-- .../component_loaders/transformer_loader.py | 2 +- .../runtime/models/dits/base.py | 6 + .../runtime/models/dits/minimax_h3.py | 89 ++++++++++++++- .../runtime/pipelines_core/lora/pipeline.py | 9 ++ .../test/unit/test_lora_commit_as_base.py | 19 ++++ .../test/unit/test_minimax_h3_dit_contract.py | 38 +++++++ 8 files changed, 269 insertions(+), 16 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py index 3bdbc3e68..49b434003 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py @@ -22,11 +22,14 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig): r"^context_embedder\.(.*)$": r"condition_proj.\1", r"^time_embedder\.linear_1\.(.*)$": r"time_embedder.proj_in.\1", r"^time_embedder\.linear_2\.(.*)$": r"time_embedder.proj_out.\1", + r"^time_embedder\.table$": r"adaln_t_table", r"^norm_out\.norm\.(.*)$": r"final_layer.norm.\1", + r"^norm_out\.folded_bias$": r"final_layer.adaln_proj.linear.bias", r"^norm_out\.linear\.(.*)$": r"final_layer.adaln_proj.linear.\1", r"^proj_out\.(.*)$": r"final_layer.video_out.\1", r"^audio_proj_out\.(.*)$": r"final_layer.audio_out.\1", r"^transformer_blocks\.(\d+)\.adaln_proj\.linear\.(.*)$": r"blocks.\1.adaln_proj.linear.\2", + r"^transformer_blocks\.(\d+)\.adaln_proj\.folded_bias$": r"blocks.\1.adaln_proj.linear.bias", r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": ( r"blocks.\1.attn.qkv_proj.\2", 0, @@ -95,6 +98,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig): qk_norm_eps: float = 1e-5 final_norm_eps: float = 1e-5 checkpoint_uses_diffusers_layout: bool = False + adaln_affine_input_dim: int | None = None def __post_init__(self) -> None: super().__post_init__() @@ -119,9 +123,14 @@ class MiniMaxH3DiTConfig(DiTConfig): "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()} - ) + model_dict = { + aliases.get(key, key): value for key, value in source_model_dict.items() + } + if source_model_dict.get("_class_name") == "MiniMaxH3PrunedTransformer3DModel": + model_dict["adaln_affine_input_dim"] = source_model_dict["time_embed_dim"] + model_dict["time_embed_dim"] = source_model_dict["adaln_rank"] + model_dict["adaln_curve_grid"] = source_model_dict["time_table_size"] + super().update_model_arch(model_dict) __all__ = [ diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index 542144779..2bba15051 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -44,6 +44,7 @@ LoRAWeightEntry = tuple[ float, int | None, int | None, + torch.nn.Parameter | None, ] @@ -104,6 +105,8 @@ class BaseLayerWithLoRA(nn.Module): self.lora_A = None self.lora_B = None + self.lora_output_offset = None + self.has_lora_output_offset = False @property def weight(self): @@ -138,10 +141,10 @@ class BaseLayerWithLoRA(nn.Module): ) # type: ignore delta = delta * self.strength out, output_bias = self.base_layer(x) - return out + delta.to(dtype=out.dtype), output_bias + out = out + delta.to(dtype=out.dtype) else: out, output_bias = self.base_layer(x) - return out, output_bias + return self._add_lora_output_offset(out), output_bias def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor: return A @@ -149,6 +152,44 @@ class BaseLayerWithLoRA(nn.Module): def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor: return B + def _scaled_lora_output_offset( + self, + offset: torch.Tensor | None, + strength: float, + rank: int | None, + alpha: int | None, + ) -> torch.Tensor | None: + if offset is None: + return None + offset = self.slice_lora_b_weights(offset.unsqueeze(-1)).squeeze(-1) + scale = strength + if rank is not None and alpha is not None and rank != alpha: + scale *= alpha / rank + return offset if scale == 1.0 else offset * scale + + def _active_lora_output_offset(self) -> torch.Tensor | None: + if self.disable_lora or not self.has_lora_output_offset: + return None + if not self.merged: + return self._scaled_lora_output_offset( + self.lora_output_offset, + self.strength, + self.lora_rank, + self.lora_alpha, + ) + combined = None + for _, _, _, strength, rank, alpha, offset in self.lora_weights_list: + scaled = self._scaled_lora_output_offset(offset, strength, rank, alpha) + if scaled is not None: + combined = scaled if combined is None else combined + scaled + return combined + + def _add_lora_output_offset(self, output: torch.Tensor) -> torch.Tensor: + offset = self._active_lora_output_offset() + if offset is None: + return output + return output + offset.to(device=output.device, dtype=output.dtype) + @staticmethod def _as_mutable_tensor(tensor: torch.Tensor) -> torch.Tensor: # lora can be reconfigured after executor forwards create inference tensors @@ -165,6 +206,7 @@ class BaseLayerWithLoRA(nn.Module): strength: float = 1.0, clear_existing: bool = False, merge_weights: bool = True, + output_offset: torch.Tensor | None = None, ) -> None: """ Set LoRA weights. Supports multiple LoRA adapters. @@ -176,17 +218,25 @@ class BaseLayerWithLoRA(nn.Module): strength: LoRA strength clear_existing: If True, clear existing LoRA weights before adding new one. If False, append to existing list (for multi-LoRA support). + output_offset: Optional constant output term paired with this adapter """ lora_A_param = torch.nn.Parameter( A ) # share storage with weights in the pipeline lora_B_param = torch.nn.Parameter(B) + output_offset_param = ( + torch.nn.Parameter(output_offset, requires_grad=False) + if output_offset is not None + else None + ) if clear_existing: self.lora_weights_list.clear() # Also clear backward compatibility attributes self.lora_A = None self.lora_B = None + self.lora_output_offset = None + self.has_lora_output_offset = False self.lora_path = None self.strength = 1.0 @@ -199,6 +249,7 @@ class BaseLayerWithLoRA(nn.Module): strength, self.lora_rank, self.lora_alpha, + output_offset_param, ) ) @@ -206,6 +257,8 @@ class BaseLayerWithLoRA(nn.Module): # This ensures backward compatibility while supporting multiple LoRA self.lora_A = lora_A_param self.lora_B = lora_B_param + self.lora_output_offset = output_offset_param + self.has_lora_output_offset |= output_offset_param is not None self.lora_path = lora_path self.strength = strength @@ -233,10 +286,18 @@ class BaseLayerWithLoRA(nn.Module): Args: data: The base weight tensor to merge LoRA into (modified in-place) - lora_list: List of (lora_A, lora_B, lora_path, lora_strength, rank, alpha) tuples + lora_list: Adapter factors, path, scale metadata, and output offset """ # Merge all LoRA adapters in order - for lora_A, lora_B, _, lora_strength, lora_rank, lora_alpha in lora_list: + for ( + lora_A, + lora_B, + _, + lora_strength, + lora_rank, + lora_alpha, + _, + ) in lora_list: lora_A_sliced = self.slice_lora_a_weights(lora_A.to(data)) lora_B_sliced = self.slice_lora_b_weights(lora_B.to(data)) @@ -286,7 +347,7 @@ class BaseLayerWithLoRA(nn.Module): ) -> bool: if os.getenv("SGLANG_DIFFUSION_LORA_MERGE_FP32", "1") != "1": return False - for _, _, lora_path, _, _, _ in lora_list: + for _, _, lora_path, _, _, _, _ in lora_list: if lora_path and "distilled-lora" in lora_path.lower(): return False return True @@ -332,7 +393,15 @@ class BaseLayerWithLoRA(nn.Module): self.strength = strength if self.lora_weights_list: self.lora_weights_list = [ - (lora_A, lora_B, lora_path, strength, lora_rank, lora_alpha) + ( + lora_A, + lora_B, + lora_path, + strength, + lora_rank, + lora_alpha, + output_offset, + ) for ( lora_A, lora_B, @@ -340,6 +409,7 @@ class BaseLayerWithLoRA(nn.Module): _, lora_rank, lora_alpha, + output_offset, ) in self.lora_weights_list ] @@ -361,11 +431,18 @@ class BaseLayerWithLoRA(nn.Module): self.strength, self.lora_rank, self.lora_alpha, + self.lora_output_offset, ) ] if not lora_list: raise ValueError("LoRA weights not set. Please set them first.") + if isinstance(self.base_layer.weight, DTensor) and any( + output_offset is not None for *_, output_offset in lora_list + ): + raise ValueError( + "LoRA output offsets require dynamic mode with FSDP-sharded weights." + ) merge_in_fp32 = self._should_merge_in_fp32(lora_list) @@ -488,6 +565,11 @@ class BaseLayerWithLoRA(nn.Module): """ if not self.merged: return + if self._active_lora_output_offset() is not None: + raise ValueError( + "A LoRA with a constant output offset cannot be committed as a " + "weight-only base." + ) weight = self.base_layer.weight if isinstance(weight, DTensor): weight = weight.to_local() @@ -498,6 +580,8 @@ class BaseLayerWithLoRA(nn.Module): self.lora_weights_list = [] self.lora_A = None self.lora_B = None + self.lora_output_offset = None + self.has_lora_output_offset = False self.lora_path = None self.strength = 1.0 @@ -534,7 +618,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def forward(self, input_: torch.Tensor) -> torch.Tensor: - if self.merged or self.disable_lora: + if self.disable_lora or (self.merged and not self.has_lora_output_offset): return self.base_layer(input_) lora_A = self.lora_A @@ -567,6 +651,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): output_parallel = output_parallel + delta_parallel.to( dtype=output_parallel.dtype ) + output_parallel = self._add_lora_output_offset(output_parallel) if self.base_layer.gather_output: output = tensor_model_parallel_all_gather(output_parallel) else: @@ -667,7 +752,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def forward(self, input_: torch.Tensor): - if self.merged or self.disable_lora: + if self.disable_lora or (self.merged and not self.has_lora_output_offset): return self.base_layer(input_) lora_A = self.lora_A @@ -723,7 +808,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): else: output = output_ output_bias = self.base_layer.bias - return output, output_bias + return self._add_lora_output_offset(output), output_bias def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor: tp_rank = get_tp_rank() @@ -779,11 +864,11 @@ class LinearWithLoRA(BaseLayerWithLoRA): delta = delta * self.strength # nn.Linear.forward() returns a single tensor, not a tuple out = self.base_layer(x) - return out + delta.to(dtype=out.dtype) + out = out + delta.to(dtype=out.dtype) else: # nn.Linear.forward() returns a single tensor out = self.base_layer(x) - return out + return self._add_lora_output_offset(out) def _use_owned_base_snapshot(snapshot_base: bool, device_type: str) -> bool: diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index ed74ac686..251dc6bad 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -232,7 +232,7 @@ class TransformerLoader(ComponentLoader): is_minimax_h3 = model_cls.__name__ == "MiniMaxH3DiTModel" if is_minimax_h3: dit_config.arch_config.checkpoint_uses_diffusers_layout = ( - cls_name == "MiniMaxH3Transformer3DModel" + cls_name != model_cls.__name__ ) checkpoint_quant_config = None diff --git a/python/sglang/multimodal_gen/runtime/models/dits/base.py b/python/sglang/multimodal_gen/runtime/models/dits/base.py index 849a05ce2..e5d545ca9 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/base.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/base.py @@ -102,6 +102,12 @@ class BaseDiT(nn.Module, ABC): """Run model-specific post-load weight fixups after all parameters are materialized.""" return None + def prepare_lora_adapter( + self, adapter: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Apply model-specific LoRA transforms after names are normalized.""" + return adapter + @property def supported_attention_backends(self) -> set[AttentionBackendEnum]: return self._supported_attention_backends diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index 8f190f54d..ecd0d34b2 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -18,6 +18,7 @@ from typing import Any, Callable import torch import torch.nn as nn from safetensors.torch import safe_open +from torch.distributed.tensor import DTensor from sglang.kernels.ops.activation.activation import ( silu_and_mul_with_activation_rounding_, @@ -1653,7 +1654,10 @@ class MiniMaxH3FinalLayer(nn.Module): class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin): - _aliases = ["MiniMaxH3Transformer3DModel"] + _aliases = [ + "MiniMaxH3Transformer3DModel", + "MiniMaxH3PrunedTransformer3DModel", + ] _fsdp_shard_conditions = [is_block] # refine_prompt_embeds drives a forward pass outside __call__. _fsdp_forward_methods = ("refine_prompt_embeds",) @@ -1666,6 +1670,65 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin): reverse_param_names_mapping = _ARCH_DEFAULTS.reverse_param_names_mapping lora_param_names_mapping = _ARCH_DEFAULTS.lora_param_names_mapping + def prepare_lora_adapter( + self, adapter: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Project released-checkpoint AdaLN LoRAs onto pruned coordinates.""" + full_width = self.arch.adaln_affine_input_dim + if full_width is None: + return adapter + + suffix = ".adaln_proj.linear.lora_A" + a_keys = sorted(key for key in adapter if key.endswith(suffix)) + if not a_keys: + return adapter + widths = {int(adapter[key].shape[-1]) for key in a_keys} + if widths == {self.arch.time_embed_dim}: + return adapter + if widths != {full_width}: + raise ValueError( + "MiniMax H3 pruned AdaLN LoRA inputs must be uniformly " + f"{self.arch.time_embed_dim} or {full_width} wide, got " + f"{sorted(widths)}." + ) + + basis = self.adaln_basis + mean = self.adaln_mean + assert basis is not None and mean is not None + if isinstance(basis, DTensor): + basis = basis.full_tensor() + mean = mean.full_tensor() + if torch.count_nonzero(basis).item() == 0: + raise ValueError( + "MiniMax H3 pruned LoRA projection requires adaln_basis and " + "adaln_mean from the component checkpoint." + ) + + projected = dict(adapter) + work_device = adapter[a_keys[0]].device + work_basis = basis.to(device=work_device, dtype=torch.float64) + work_mean = mean.to(device=work_device, dtype=torch.float64) + for a_key in a_keys: + b_key = a_key[: -len("lora_A")] + "lora_B" + if b_key not in adapter: + raise ValueError(f"MiniMax H3 AdaLN LoRA is missing {b_key!r}.") + a = adapter[a_key] + b = adapter[b_key] + a64 = a.to(torch.float64) + b64 = b.to(device=work_device, dtype=torch.float64) + projected[a_key] = (a64 @ work_basis.T).to(torch.float32) + projected[a_key[: -len("lora_A")] + "lora_output_offset"] = ( + b64 @ (a64 @ work_mean) + ).to(torch.float32) + + logger.info( + "Projected %d MiniMax H3 AdaLN LoRA modules from width %d to %d", + len(a_keys), + full_width, + self.arch.time_embed_dim, + ) + return projected + def prepare_adaln_plans(self, step_timesteps: list[torch.Tensor]) -> None: """Fill the AdaLN cache for this request before denoising starts. @@ -1845,6 +1908,28 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin): ), requires_grad=False, ) + if arch.adaln_affine_input_dim is None: + self.register_parameter("adaln_basis", None) + self.register_parameter("adaln_mean", None) + else: + self.register_parameter( + "adaln_basis", + nn.Parameter( + torch.empty( + arch.time_embed_dim, + arch.adaln_affine_input_dim, + dtype=_FP32_DTYPE, + ), + requires_grad=False, + ), + ) + self.register_parameter( + "adaln_mean", + nn.Parameter( + torch.empty(arch.adaln_affine_input_dim, dtype=_FP32_DTYPE), + requires_grad=False, + ), + ) self.rope = MiniMaxH3Rope(arch.rope_inv_freq_len) self.token_refiner = MiniMaxH3TokenRefiner( arch, @@ -1929,6 +2014,8 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin): if not name.startswith("time_embedder.") ] fp32_param_names.append("adaln_t_table") + if self.adaln_basis is not None: + fp32_param_names.extend(("adaln_basis", "adaln_mean")) for name in fp32_param_names: param = self.get_parameter(name) if param.dtype != _FP32_DTYPE: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py index a1afd5bec..f4412c0f1 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py @@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, ) +from sglang.multimodal_gen.runtime.models.dits.base import BaseDiT from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( ComposedPipelineBase, ) @@ -688,6 +689,9 @@ class LoRAPipeline(ComposedPipelineBase): layer.set_lora_weights( self.lora_adapters[nickname][lora_A_name], self.lora_adapters[nickname][lora_B_name], + output_offset=self.lora_adapters[nickname].get( + name + ".lora_output_offset" + ), lora_path=path, strength=lora_strength, merge_weights=merge_weights and not use_cache, @@ -918,6 +922,11 @@ class LoRAPipeline(ComposedPipelineBase): adapter_lora_alpha, self.device, ) + transformer = self.modules["transformer"] + if isinstance(transformer, BaseDiT): + self.lora_adapters[lora_nickname] = transformer.prepare_lora_adapter( + self.lora_adapters[lora_nickname] + ) self.loaded_adapter_paths[lora_nickname] = lora_path self.loaded_adapter_alphas[lora_nickname] = adapter_lora_alpha diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py b/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py index 0eac3bf91..dfc630e7c 100644 --- a/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py +++ b/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py @@ -6,6 +6,7 @@ invariants are checked on CPU with a plain nn.Linear base (the parallel forward path needs a distributed runtime and is covered by the diffusion server tests). """ +import pytest import torch from torch import nn @@ -47,6 +48,24 @@ def test_commit_merged_as_base_promotes_weights_and_resets_state(): assert layer.lora_A is None and layer.lora_B is None +def test_lora_output_offset_tracks_dynamic_and_merged_scale(): + layer = _make_layer(torch.zeros(1, 2), rank=1, alpha=2) + inputs = torch.zeros(3, 2) + layer.set_lora_weights( + torch.zeros(1, 2), + torch.zeros(1, 1), + output_offset=torch.tensor([3.0]), + strength=0.5, + clear_existing=True, + merge_weights=False, + ) + torch.testing.assert_close(layer(inputs), torch.full((3, 1), 3.0)) + layer.merge_lora_weights() + torch.testing.assert_close(layer(inputs), torch.full((3, 1), 3.0)) + with pytest.raises(ValueError, match="constant output offset"): + layer.commit_merged_as_base() + + def test_dynamic_delta_after_commit_does_not_unmerge_base(): torch.manual_seed(1) out_f, in_f, rank = 4, 5, 2 diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py index de3d8f3d6..44c421666 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 """Mixed-precision weight and TP/Ulysses numerical contracts for H3 DiT.""" +from types import SimpleNamespace from unittest.mock import patch import pytest @@ -44,6 +45,24 @@ def _ensure_single_process_parallel_runtime() -> None: maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1) +def test_pruned_adaln_lora_projection_preserves_affine_term(): + model = SimpleNamespace( + arch=SimpleNamespace(adaln_affine_input_dim=3, time_embed_dim=2), + adaln_basis=torch.tensor([[1.0, 0.0, 2.0], [0.0, 1.0, -1.0]]), + adaln_mean=torch.tensor([1.0, 2.0, 3.0]), + ) + prefix = "blocks.0.adaln_proj.linear." + a = torch.tensor([[2.0, 3.0, 4.0]]) + b = torch.tensor([[5.0], [6.0]]) + actual = MiniMaxH3DiTModel.prepare_lora_adapter( + model, {prefix + "lora_A": a, prefix + "lora_B": b} + ) + torch.testing.assert_close(actual[prefix + "lora_A"], a @ model.adaln_basis.T) + torch.testing.assert_close( + actual[prefix + "lora_output_offset"], b @ (a @ model.adaln_mean) + ) + + def test_native_weight_names_and_grouped_qkv_reorder(): arch = MiniMaxH3DiTArchConfig() assert arch.reverse_param_names_mapping == {} @@ -90,6 +109,12 @@ def test_native_weight_names_and_grouped_qkv_reorder(): 1, 3, ) + assert mapping("time_embedder.table") == ("adaln_t_table", None, None) + assert mapping("transformer_blocks.7.adaln_proj.folded_bias") == ( + "blocks.7.adaln_proj.linear.bias", + None, + None, + ) diffusers_weights = [ ("transformer_blocks.0.attn.to_q.qweight", torch.full((2, 3), 1)), @@ -110,6 +135,19 @@ def test_native_weight_names_and_grouped_qkv_reorder(): torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]), ) + pruned_config = MiniMaxH3DiTConfig() + pruned_config.update_model_arch( + { + "_class_name": "MiniMaxH3PrunedTransformer3DModel", + "adaln_rank": 8, + "time_embed_dim": 2688, + "time_table_size": 1025, + } + ) + assert pruned_config.arch_config.time_embed_dim == 8 + assert pruned_config.arch_config.adaln_curve_grid == 1025 + assert pruned_config.arch_config.adaln_affine_input_dim == 2688 + weight = torch.arange(12, dtype=torch.float32).reshape(12, 1) actual = _reorder_grouped_qkv_to_qkv( weight,