Skip to content

minimax_h3

FastVideo-native MiniMax H3 joint audio-video diffusion transformer.

Classes

fastvideo.models.dits.minimax_h3.MiniMaxH3AdaLayerNormModulation

MiniMaxH3AdaLayerNormModulation(time_embed_dim: int, hidden_size: int, quant_config: QuantizationConfig | None = None, prefix: str = '', apply_silu: bool = True)

Bases: Module

Produce six modulation tables for every (timestep, modality) pair.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    time_embed_dim: int,
    hidden_size: int,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
    apply_silu: bool = True,
) -> None:
    super().__init__()
    self.hidden_size = hidden_size
    self.apply_silu = apply_silu
    self.linear = ReplicatedLinear(
        time_embed_dim,
        6 * hidden_size * MINIMAX_H3_MODALITY_NUM,
        bias=True,
        quant_config=quant_config,
        prefix=f"{prefix}.linear",
    )

fastvideo.models.dits.minimax_h3.MiniMaxH3AdaLayerNormOut

MiniMaxH3AdaLayerNormOut(hidden_size: int, time_embed_dim: int, eps: float, quant_config: QuantizationConfig | None = None, prefix: str = '', apply_silu: bool = True)

Bases: Module

Final RMSNorm with per-timestep row modulation.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    time_embed_dim: int,
    eps: float,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
    apply_silu: bool = True,
) -> None:
    super().__init__()
    self.norm = nn.RMSNorm(hidden_size, eps=eps)
    self.apply_silu = apply_silu
    self.linear = ReplicatedLinear(
        time_embed_dim,
        2 * hidden_size,
        bias=True,
        quant_config=quant_config,
        prefix=f"{prefix}.linear",
    )

fastvideo.models.dits.minimax_h3.MiniMaxH3Attention

MiniMaxH3Attention(hidden_size: int, num_attention_heads: int, attention_head_dim: int, qk_norm_eps: float, supported_attention_backends: tuple[AttentionBackendEnum, ...], quant_config: QuantizationConfig | None, prefix: str, fuse_qknorm_rope: bool = False, fa4_packed_varlen: bool = False)

Bases: Module

Full self-attention over one sequence-parallel packed document.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    num_attention_heads: int,
    attention_head_dim: int,
    qk_norm_eps: float,
    supported_attention_backends: tuple[AttentionBackendEnum, ...],
    quant_config: QuantizationConfig | None,
    prefix: str,
    fuse_qknorm_rope: bool = False,
    fa4_packed_varlen: bool = False,
) -> None:
    super().__init__()
    self.num_attention_heads = num_attention_heads
    self.attention_head_dim = attention_head_dim
    inner_dim = num_attention_heads * attention_head_dim
    self.to_q = ReplicatedLinear(
        hidden_size,
        inner_dim,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.to_q",
    )
    self.to_k = ReplicatedLinear(
        hidden_size,
        inner_dim,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.to_k",
    )
    self.to_v = ReplicatedLinear(
        hidden_size,
        inner_dim,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.to_v",
    )
    self.norm_q = nn.RMSNorm(attention_head_dim, eps=qk_norm_eps)
    self.norm_k = nn.RMSNorm(attention_head_dim, eps=qk_norm_eps)
    self.to_out = ReplicatedLinear(
        inner_dim,
        hidden_size,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.to_out",
    )
    self.fuse_qknorm_rope = fuse_qknorm_rope
    # VSA carries a learned gate on its pooled-compression branch. The H3
    # checkpoint has no such weight, so the loader zero-initializes it
    # (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
    # until finetuned. Built only when VSA-H3 actually resolves, keeping
    # the FLASH/SDPA paths and their state_dict untouched.
    resolved_backend = get_attn_backend(attention_head_dim,
                                        get_compute_dtype(),
                                        supported_attention_backends=supported_attention_backends)
    use_vsa = resolved_backend.get_name() == "VIDEO_SPARSE_ATTN_H3"
    attention_cls = DistributedAttention_VSA if use_vsa else DistributedAttention
    self.distributed_attention = attention_cls(
        num_heads=num_attention_heads,
        head_size=attention_head_dim,
        causal=False,
        supported_attention_backends=supported_attention_backends,
        prefix=prefix,
        fa4_packed_varlen=fa4_packed_varlen,
    )
    self.to_gate_compress: ReplicatedLinear | None = None
    # None = unchecked; the first forward tests the loaded weight once and
    # skips the gate branch entirely while it is structurally zero.
    self._gate_compress_active: bool | None = None
    if use_vsa:
        self.to_gate_compress = ReplicatedLinear(
            hidden_size,
            inner_dim,
            bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.to_gate_compress",
        )

fastvideo.models.dits.minimax_h3.MiniMaxH3FeedForward

MiniMaxH3FeedForward(hidden_size: int, ffn_dim: int, quant_config: QuantizationConfig | None = None, prefix: str = '', fuse_swiglu: bool = False)

Bases: Module

Bias-free H3 SwiGLU with value-first packed halves.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    ffn_dim: int,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
    fuse_swiglu: bool = False,
) -> None:
    super().__init__()
    self.fc_in = ReplicatedLinear(
        hidden_size,
        2 * ffn_dim,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.fc_in",
    )
    self.fc_out = ReplicatedLinear(
        ffn_dim,
        hidden_size,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.fc_out",
    )
    self.fuse_swiglu = fuse_swiglu

fastvideo.models.dits.minimax_h3.MiniMaxH3RotaryPosEmbed

MiniMaxH3RotaryPosEmbed(rope_freq_dim: int, rope_theta: float)

Bases: Module

Three-axis rotary frequencies over packed (t, h, w) coordinates.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(self, rope_freq_dim: int, rope_theta: float) -> None:
    super().__init__()
    inv_freq = 1.0 / (rope_theta**(torch.arange(0, 2 * rope_freq_dim, 2, dtype=torch.float32) /
                                   (2 * rope_freq_dim)))
    self.register_buffer("inv_freq", inv_freq, persistent=False)

Methods:

fastvideo.models.dits.minimax_h3.MiniMaxH3RotaryPosEmbed.forward
forward(position_ids: Tensor) -> tuple[Tensor, Tensor]

Build rotary tensors on the device that owns the packed positions.

Source code in fastvideo/models/dits/minimax_h3.py
def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Build rotary tensors on the device that owns the packed positions."""
    position_ids = position_ids.to(torch.float32)
    # Analytic rotary positional embedding (RoPE) state is non-persistent,
    # so runtime coordinates own the device after loading or state offload.
    inv_freq = self.inv_freq.to(position_ids.device)
    freqs = position_ids.unsqueeze(-1) * inv_freq.view(1, 1, -1)
    freqs_t, freqs_h, freqs_w = freqs.unbind(dim=1)
    freqs = torch.cat((freqs_t, freqs_h, freqs_w), dim=-1)
    freqs = torch.cat((freqs, freqs), dim=-1)
    return freqs.cos(), freqs.sin()

fastvideo.models.dits.minimax_h3.MiniMaxH3TokenRefiner

MiniMaxH3TokenRefiner(hidden_size: int, num_attention_heads: int, attention_head_dim: int, ffn_dim: int, num_layers: int, norm_eps: float, qk_norm_eps: float, final_norm_eps: float, supported_attention_backends: tuple[AttentionBackendEnum, ...], quant_config: QuantizationConfig | None, prefix: str)

Bases: Module

Two-block text refiner used before packing the modalities.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    num_attention_heads: int,
    attention_head_dim: int,
    ffn_dim: int,
    num_layers: int,
    norm_eps: float,
    qk_norm_eps: float,
    final_norm_eps: float,
    supported_attention_backends: tuple[AttentionBackendEnum, ...],
    quant_config: QuantizationConfig | None,
    prefix: str,
) -> None:
    super().__init__()
    self.refiner_blocks = nn.ModuleList([
        MiniMaxH3TokenRefinerBlock(
            hidden_size,
            num_attention_heads,
            attention_head_dim,
            ffn_dim,
            norm_eps,
            qk_norm_eps,
            supported_attention_backends,
            quant_config,
            prefix=f"{prefix}.refiner_blocks.{index}",
        ) for index in range(num_layers)
    ])
    self.final_norm = nn.RMSNorm(hidden_size, eps=final_norm_eps)

fastvideo.models.dits.minimax_h3.MiniMaxH3TokenRefinerBlock

MiniMaxH3TokenRefinerBlock(hidden_size: int, num_attention_heads: int, attention_head_dim: int, ffn_dim: int, norm_eps: float, qk_norm_eps: float, supported_attention_backends: tuple[AttentionBackendEnum, ...], quant_config: QuantizationConfig | None, prefix: str)

Bases: Module

Plain pre-norm Transformer block for the projected text stream.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    num_attention_heads: int,
    attention_head_dim: int,
    ffn_dim: int,
    norm_eps: float,
    qk_norm_eps: float,
    supported_attention_backends: tuple[AttentionBackendEnum, ...],
    quant_config: QuantizationConfig | None,
    prefix: str,
) -> None:
    super().__init__()
    self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
    self.attn = MiniMaxH3Attention(
        hidden_size,
        num_attention_heads,
        attention_head_dim,
        qk_norm_eps,
        supported_attention_backends,
        quant_config,
        prefix=f"{prefix}.attn",
    )
    self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
    self.ff = MiniMaxH3FeedForward(
        hidden_size,
        ffn_dim,
        quant_config=quant_config,
        prefix=f"{prefix}.ff",
    )

fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel

MiniMaxH3Transformer3DModel(config: MiniMaxH3Config, hf_config: dict[str, Any])

Bases: BaseDiT

Joint H3 Transformer over one padless text/audio/video document.

The layout builder validates semantic rows before denoising. Sequence- parallel padding is transport-only and DistributedAttention trims it before attention.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
    super().__init__(config, hf_config)
    arch = config.arch_config
    self.enabled_fusions = _enabled_minimax_h3_fusions()
    if self.enabled_fusions:
        if HAVE_TRITON:
            logger.info(
                "MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled forwards "
                "fall back to eager; torch.compile captures opaque custom-op boundaries).",
                ",".join(sorted(self.enabled_fusions)))
        else:
            logger.warning(
                "FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
                "every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
    sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
    if arch.num_attention_heads % sp_world_size:
        raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
                         f"sequence parallel size ({sp_world_size}).")

    self.hidden_size = arch.hidden_size
    self.num_attention_heads = arch.num_attention_heads
    self.num_channels_latents = arch.in_channels
    self.patch_size = tuple(int(value) for value in arch.patch_size)
    video_patch_dim = arch.in_channels * math.prod(arch.patch_size)

    self.proj_in = ReplicatedLinear(
        video_patch_dim,
        arch.hidden_size,
        bias=True,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.proj_in",
    )
    self.audio_proj_in = ReplicatedLinear(
        arch.audio_in_channels,
        arch.hidden_size,
        bias=True,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.audio_proj_in",
    )
    self.context_embedder = ReplicatedLinear(
        arch.text_dim,
        arch.hidden_size,
        bias=True,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.context_embedder",
    )
    self.time_proj = Timesteps(
        num_channels=arch.freq_dim,
        flip_sin_to_cos=True,
        downscale_freq_shift=0,
    )
    self.time_embedder = MLP(
        arch.freq_dim,
        arch.time_embed_hidden_dim,
        arch.time_embed_dim,
        act_type="silu",
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.time_embedder",
    )
    self.adaln_rank: int | None = arch.adaln_rank
    if self.adaln_rank is not None and config.uniform_parameter_dtype:
        raise ValueError(
            "Rank-reduced AdaLN checkpoints (adaln_rank set) cannot be trained: "
            "uniform_parameter_dtype needs one dtype for every trainable "
            "parameter, but factorized AdaLN weights are pinned to FP16 "
            "(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank "
            "checkpoint instead, then re-fit the basis with "
            "scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
    adaln_dim = self.adaln_rank or arch.time_embed_dim
    self.adaln_basis = ReplicatedLinear(
        arch.time_embed_dim,
        self.adaln_rank,
        bias=False,
        quant_config=None,
        prefix=f"{config.prefix}.adaln_basis",
    ) if self.adaln_rank else None

    self.rope = MiniMaxH3RotaryPosEmbed(arch.rope_freq_dim, arch.rope_theta)
    # per-generation caches for loop-invariant work (see _rotary_for /
    # _refined_text); plain attrs, never in state_dict
    self._rope_cache: tuple | None = None
    self._text_cache: tuple | None = None
    self.token_refiner = MiniMaxH3TokenRefiner(
        arch.hidden_size,
        arch.num_attention_heads,
        arch.attention_head_dim,
        arch.ffn_dim,
        arch.num_refiner_layers,
        arch.norm_eps,
        arch.qk_norm_eps,
        arch.final_norm_eps,
        # The refiner attends over the text stream only; the packed-sequence
        # VSA backend must never be selected for it.
        tuple(backend for backend in self.supported_attention_backends
              if backend != AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3),
        config.quant_config,
        prefix=f"{config.prefix}.token_refiner",
    )
    self.transformer_blocks = nn.ModuleList([
        MiniMaxH3TransformerBlock(
            arch.hidden_size,
            arch.num_attention_heads,
            arch.attention_head_dim,
            arch.ffn_dim,
            adaln_dim,
            arch.norm_eps,
            arch.qk_norm_eps,
            self.supported_attention_backends,
            config.quant_config,
            prefix=f"{config.prefix}.transformer_blocks.{index}",
            adaln_apply_silu=self.adaln_rank is None,
            fuse_modulate="modulate" in self.enabled_fusions,
            fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
            fuse_swiglu="swiglu" in self.enabled_fusions,
            fa4_packed_varlen=envs.FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN,
        ) for index in range(arch.num_layers)
    ])
    self.norm_out = MiniMaxH3AdaLayerNormOut(
        arch.hidden_size,
        adaln_dim,
        arch.final_norm_eps,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.norm_out",
        apply_silu=self.adaln_rank is None,
    )
    self.proj_out = ReplicatedLinear(
        arch.hidden_size,
        video_patch_dim,
        bias=True,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.proj_out",
    )
    self.audio_proj_out = ReplicatedLinear(
        arch.hidden_size,
        arch.audio_in_channels,
        bias=True,
        quant_config=config.quant_config,
        prefix=f"{config.prefix}.audio_proj_out",
    )
    self.__post_init__()

Methods:

fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.forward
forward(hidden_states: Tensor, audio_hidden_states: Tensor, encoder_hidden_states: Tensor, timestep: Tensor, timestep_indices: Tensor, token_tags: Tensor, position_ids: Tensor, video_indices: Tensor, audio_indices: Tensor, text_indices: Tensor) -> tuple[Tensor, Tensor]

Predict video and audio velocities from one caller-defined packed layout.

Source code in fastvideo/models/dits/minimax_h3.py
def forward(
    self,
    hidden_states: torch.Tensor,
    audio_hidden_states: torch.Tensor,
    encoder_hidden_states: torch.Tensor,
    timestep: torch.Tensor,
    timestep_indices: torch.Tensor,
    token_tags: torch.Tensor,
    position_ids: torch.Tensor,
    video_indices: torch.Tensor,
    audio_indices: torch.Tensor,
    text_indices: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Predict video and audio velocities from one caller-defined packed layout."""
    if position_ids.ndim != 2 or position_ids.shape[-1] != 3:
        raise ValueError(f"position_ids must have shape (seq_len, 3), got {tuple(position_ids.shape)}.")
    sequence_length = position_ids.shape[0]
    if token_tags.shape != (sequence_length, ) or timestep_indices.shape != (sequence_length, ):
        raise ValueError("token_tags and timestep_indices must both match the packed sequence length.")
    if hidden_states.shape[1] != video_indices.numel():
        raise ValueError("hidden_states row count must match video_indices.")
    if audio_hidden_states.shape[1] != audio_indices.numel():
        raise ValueError("audio_hidden_states row count must match audio_indices.")
    if encoder_hidden_states.shape[1] != text_indices.numel():
        raise ValueError("encoder_hidden_states row count must match text_indices.")

    video_embeds, _ = self.proj_in(hidden_states.to(self.proj_in.weight.dtype))
    audio_embeds, _ = self.audio_proj_in(audio_hidden_states.to(self.audio_proj_in.weight.dtype))
    text_embeds = self._refined_text(encoder_hidden_states)
    rotary_emb = self._rotary_for(position_ids, text_embeds.dtype)
    sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1

    # text/video/audio indices partition [0, sequence_length), so the
    # uninitialized buffer is fully overwritten; in-place index_copy_ avoids
    # the three full-buffer clones out-of-place index_copy would make.
    packed_hidden_states = text_embeds.new_empty((text_embeds.shape[0], sequence_length, text_embeds.shape[-1]))
    packed_hidden_states.index_copy_(1, text_indices, text_embeds)
    packed_hidden_states.index_copy_(1, video_indices, video_embeds.to(text_embeds.dtype))
    packed_hidden_states.index_copy_(1, audio_indices, audio_embeds.to(text_embeds.dtype))

    temb = self.time_proj(timestep)
    temb = self.time_embedder(temb.to(self.time_embedder.fc_in.weight.dtype))
    if self.adaln_basis is not None:
        temb, _ = self.adaln_basis(F.silu(temb).to(self.adaln_basis.weight.dtype))
    adaln_indices = timestep_indices * MINIMAX_H3_MODALITY_NUM + token_tags
    local_timestep_indices = timestep_indices
    original_seq_len = sequence_length

    if sp_world_size > 1:
        packed_hidden_states, _ = sequence_model_parallel_shard(packed_hidden_states, dim=1)
        rotary_cos, _ = sequence_model_parallel_shard(rotary_emb[0], dim=0)
        rotary_sin, _ = sequence_model_parallel_shard(rotary_emb[1], dim=0)
        adaln_indices, _ = sequence_model_parallel_shard(adaln_indices, dim=0)
        local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
        rotary_emb = (rotary_cos, rotary_sin)

    # The eager driver owns profiling markers while each block's compiled
    # forward owns the graph that the marker surrounds.
    for block_index, block in enumerate(self.transformer_blocks):
        with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
            packed_hidden_states = block(
                packed_hidden_states,
                temb,
                adaln_indices,
                rotary_emb,
                original_seq_len,
            )

    packed_hidden_states = self.norm_out(
        packed_hidden_states,
        temb,
        local_timestep_indices,
    ).to(self.proj_out.weight.dtype)
    video_output, _ = self.proj_out(packed_hidden_states)
    audio_output, _ = self.audio_proj_out(packed_hidden_states)
    if sp_world_size > 1:
        video_output = sequence_model_parallel_all_gather_with_unpad(video_output, original_seq_len, dim=1)
        audio_output = sequence_model_parallel_all_gather_with_unpad(audio_output, original_seq_len, dim=1)
    video_output = video_output.index_select(1, video_indices)
    audio_output = audio_output.index_select(1, audio_indices)

    return video_output, audio_output
fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.materialize_non_persistent_buffers
materialize_non_persistent_buffers(device: device, dtype: dtype | None = None) -> None

Rebuild analytic RoPE state on the checkpoint loader device.

RoPE frequencies are absent from the checkpoint, so meta-device model construction and device moves must derive the buffer from architecture fields before the first forward pass.

Source code in fastvideo/models/dits/minimax_h3.py
def materialize_non_persistent_buffers(
    self,
    device: torch.device,
    dtype: torch.dtype | None = None,
) -> None:
    """Rebuild analytic RoPE state on the checkpoint loader device.

    RoPE frequencies are absent from the checkpoint, so meta-device model
    construction and device moves must derive the buffer from architecture
    fields before the first forward pass.
    """
    del dtype
    if self.rope.inv_freq.is_meta or self.rope.inv_freq.device != device:
        arch = self.config.arch_config
        inv_freq = 1.0 / (arch.rope_theta
                          **(torch.arange(0, 2 * arch.rope_freq_dim, 2, device=device, dtype=torch.float32) /
                             (2 * arch.rope_freq_dim)))
        self.rope._buffers["inv_freq"] = inv_freq
fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.prepare_for_compile
prepare_for_compile() -> None

Pipeline hook, called once right before torch.compile wraps the blocks.

Resolve each loaded VSA compression gate eagerly. Generic and training compile retain their established attention dispatch; only the inference loader's separate prepare_for_regional_compile hook may preselect the inference-only sm_100a path.

The inference-only Triton fusions expose fake-backed custom operators, so Dynamo can keep them active as opaque nodes inside each fullgraph block instead of tracing into their launcher implementation.

Source code in fastvideo/models/dits/minimax_h3.py
def prepare_for_compile(self) -> None:
    """Pipeline hook, called once right before torch.compile wraps the blocks.

    Resolve each loaded VSA compression gate eagerly. Generic and training
    compile retain their established attention dispatch; only the
    inference loader's separate ``prepare_for_regional_compile`` hook may
    preselect the inference-only sm_100a path.

    The inference-only Triton fusions expose fake-backed custom operators,
    so Dynamo can keep them active as opaque nodes inside each fullgraph
    block instead of tracing into their launcher implementation.
    """
    gate_states: list[bool] = []
    for block in self.transformer_blocks:
        attention = block.attn
        if attention.to_gate_compress is not None:
            attention._resolve_gate_compress_for_compile()
            assert attention._gate_compress_active is not None
            gate_states.append(attention._gate_compress_active)
    if gate_states:
        logger.info(
            "Resolved MiniMax H3 VSA compression gates before torch.compile: %d active, %d inactive",
            sum(gate_states),
            len(gate_states) - sum(gate_states),
        )
    if self.enabled_fusions:
        logger.info(
            "MiniMax H3 inference fusions remain active under torch.compile through custom-op boundaries: %s",
            ",".join(sorted(self.enabled_fusions)),
        )
fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.prepare_for_regional_compile
prepare_for_regional_compile() -> str | None

Resolve state used only by inference regional fullgraph compile.

Source code in fastvideo/models/dits/minimax_h3.py
def prepare_for_regional_compile(self) -> str | None:
    """Resolve state used only by inference regional fullgraph compile."""
    self.prepare_for_compile()
    prepared_vsa_impls = 0
    unsupported_reasons: set[str] = set()
    for block in self.transformer_blocks:
        attention = block.attn
        prepare_vsa = getattr(attention.distributed_attention.attn_impl, "prepare_for_regional_compile", None)
        if not callable(prepare_vsa):
            continue
        # Post-load FP8 conversion may replace to_q.weight with packed
        # buffers. Either representation identifies the local device.
        query_state = next(attention.to_q.parameters(), None)
        if query_state is None:
            query_state = next(attention.to_q.buffers(), None)
        if query_state is None:
            raise RuntimeError("MiniMax H3 to_q has no materialized parameter or buffer for compile setup.")
        unsupported = prepare_vsa(query_state.device)
        if unsupported:
            unsupported_reasons.add(str(unsupported))
        prepared_vsa_impls += 1
    if prepared_vsa_impls:
        logger.info("Prepared %d MiniMax H3 VSA attention implementations for regional torch.compile",
                    prepared_vsa_impls)
    if unsupported_reasons:
        return "; ".join(sorted(unsupported_reasons))
    return None

fastvideo.models.dits.minimax_h3.MiniMaxH3TransformerBlock

MiniMaxH3TransformerBlock(hidden_size: int, num_attention_heads: int, attention_head_dim: int, ffn_dim: int, time_embed_dim: int, norm_eps: float, qk_norm_eps: float, supported_attention_backends: tuple[AttentionBackendEnum, ...], quant_config: QuantizationConfig | None, prefix: str, adaln_apply_silu: bool = True, fuse_modulate: bool = False, fuse_qknorm_rope: bool = False, fuse_swiglu: bool = False, fa4_packed_varlen: bool = False)

Bases: Module

Packed self-attention and feed-forward branches with row-indexed AdaLN.

Source code in fastvideo/models/dits/minimax_h3.py
def __init__(
    self,
    hidden_size: int,
    num_attention_heads: int,
    attention_head_dim: int,
    ffn_dim: int,
    time_embed_dim: int,
    norm_eps: float,
    qk_norm_eps: float,
    supported_attention_backends: tuple[AttentionBackendEnum, ...],
    quant_config: QuantizationConfig | None,
    prefix: str,
    adaln_apply_silu: bool = True,
    fuse_modulate: bool = False,
    fuse_qknorm_rope: bool = False,
    fuse_swiglu: bool = False,
    fa4_packed_varlen: bool = False,
) -> None:
    super().__init__()
    self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
    self.attn = MiniMaxH3Attention(
        hidden_size,
        num_attention_heads,
        attention_head_dim,
        qk_norm_eps,
        supported_attention_backends,
        quant_config,
        prefix=f"{prefix}.attn",
        fuse_qknorm_rope=fuse_qknorm_rope,
        fa4_packed_varlen=fa4_packed_varlen,
    )
    self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
    self.ff = MiniMaxH3FeedForward(
        hidden_size,
        ffn_dim,
        quant_config=quant_config,
        prefix=f"{prefix}.ff",
        fuse_swiglu=fuse_swiglu,
    )
    self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
        time_embed_dim,
        hidden_size,
        quant_config=quant_config,
        prefix=f"{prefix}.adaln_proj",
        apply_silu=adaln_apply_silu,
    )
    self.fuse_modulate = fuse_modulate

Functions: