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",
    )

Methods:

fastvideo.models.dits.minimax_h3.MiniMaxH3AdaLayerNormModulation.enable_host_cache
enable_host_cache(table: dict | None = None) -> None

Keep the projection in pinned host memory and cache its output per timestep set.

The modulation is a pure function of the timestep embedding, and few-step checkpoints sample a fixed timestep ladder, so each block's output is a small constant table. The weights (the largest bf16 tensors in the DiT) then never occupy device memory: a cache miss copies them in for one matmul.

Source code in fastvideo/models/dits/minimax_h3.py
def enable_host_cache(self, table: dict | None = None) -> None:
    """Keep the projection in pinned host memory and cache its output per timestep set.

    The modulation is a pure function of the timestep embedding, and few-step checkpoints sample a fixed
    timestep ladder, so each block's output is a small constant table. The weights (the largest bf16 tensors
    in the DiT) then never occupy device memory: a cache miss copies them in for one matmul.
    """
    weight, bias = self.linear.weight, self.linear.bias
    if table is not None and (weight is None or bias is None):
        # Table-only load skipped these tensors entirely.
        self._host_weight = self._host_bias = None
        self._modulation_cache = dict(table)
        self._cache_key = None
        return
    weight, bias = weight.data, bias.data
    if table is None:
        self._host_weight = weight.to("cpu").pin_memory()
        self._host_bias = bias.to("cpu").pin_memory()
    else:
        # Precomputed modulation for a fixed timestep ladder: the projection weights are not needed at all.
        self._host_weight = self._host_bias = None
    self.linear.weight.data = torch.empty(0, dtype=weight.dtype)
    self.linear.bias.data = torch.empty(0, dtype=bias.dtype)
    self._modulation_cache: dict[Any, torch.Tensor] = dict(table or {})
    self._cache_key: Any = None

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, exact_rope: 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,
    exact_rope: 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
    self.exact_rope = exact_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,
    )
    # Opt-in inference route: VSA-H3 selection on the block-sparse FP4
    # kernel (see minimax_h3_vsa_fp4); grad and compile keep the generic path.
    self._layer_idx = layer_idx_from_prefix(prefix, default=-1)
    self._vsa_fp4 = use_vsa and vsa_fp4_requested()
    self._vsa_tile_first = use_vsa and envs.FASTVIDEO_H3_VSA_TILE_FIRST.get()
    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, exact_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,
    exact_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
    self.exact_swiglu = exact_swiglu
    self.use_mxfp8 = isinstance(self.fc_in.quant_method, MXFP8QuantizeMethod) and isinstance(
        self.fc_out.quant_method, MXFP8QuantizeMethod)
    # Inference-only token chunking: the 2 * ffn_dim intermediate is ~5.3x the block input
    # (4.5 GiB at 78k tokens), so chunks bound the activation peak on 24-32 GB GPUs.
    self.chunk_tokens = envs.FASTVIDEO_H3_FFN_CHUNK_TOKENS.get()

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)))
    self.enabled_exact_kernels = _enabled_minimax_h3_exact_kernels()
    if self.enabled_exact_kernels:
        if exact_kernels.HAVE_TRITON:
            logger.info(
                "MiniMax H3 bit-exact kernels enabled: %s (CUDA BF16 inference-only; other inputs and any "
                "op whose FASTVIDEO_MINIMAX_H3_FUSIONS fusion is on keep their existing path).",
                ",".join(sorted(self.enabled_exact_kernels)))
        else:
            logger.warning(
                "FASTVIDEO_MINIMAX_H3_EXACT_KERNELS requested %s but Triton is unavailable; "
                "every forward stays on the eager path.", ",".join(sorted(self.enabled_exact_kernels)))
    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.get(),
            exact_kernels=self.enabled_exact_kernels,
        ) 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,
    )
    # Parallel Decoding Distillation (PDD) students widen both output
    # projections to ``pdd_steps`` heads, one per interval of a fixed fine
    # time grid. The sampler fuses one block of heads per forward
    # (``fuse_pdd_block``), so the fused output keeps the released width.
    self.pdd_steps: int | None = arch.pdd_steps
    self.proj_out = self._output_projection(arch.hidden_size, video_patch_dim, config, "proj_out")
    self.audio_proj_out = self._output_projection(arch.hidden_size, arch.audio_in_channels, config,
                                                  "audio_proj_out")
    self.__post_init__()

Attributes

fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.pdd_linears property
pdd_linears: dict[str, PDDReplicatedLinear]

The widened {"video": proj_out, "audio": audio_proj_out} heads.

Methods:

fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.attach_step_splice
attach_step_splice(late: Module, from_step: int) -> None

Hand denoising steps from_step onward to late (same architecture, other weights).

Early DMD steps fix layout and object count, late ones texture and detail, so two checkpoints can split the trajectory. late is kept out of this module's children: it is placed, offloaded and checkpointed on its own.

Source code in fastvideo/models/dits/minimax_h3.py
def attach_step_splice(self, late: nn.Module, from_step: int) -> None:
    """Hand denoising steps ``from_step`` onward to ``late`` (same architecture, other weights).

    Early DMD steps fix layout and object count, late ones texture and detail, so two checkpoints
    can split the trajectory. ``late`` is kept out of this module's children: it is placed, offloaded
    and checkpointed on its own.
    """
    object.__setattr__(self, "_splice_late", late)
    self._splice_from_step = int(from_step)
fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.enable_adaln_host_cache
enable_adaln_host_cache(table_path: str | None = None) -> None

Move every block's AdaLN projection to pinned host memory behind a per-timestep cache.

With table_path (written by FASTVIDEO_H3_ADALN_DUMP), the cache is prefilled from precomputed modulation tables and the projection weights are dropped entirely.

Source code in fastvideo/models/dits/minimax_h3.py
def enable_adaln_host_cache(self, table_path: str | None = None) -> None:
    """Move every block's AdaLN projection to pinned host memory behind a per-timestep cache.

    With ``table_path`` (written by FASTVIDEO_H3_ADALN_DUMP), the cache is prefilled from precomputed
    modulation tables and the projection weights are dropped entirely.
    """
    tables = None
    if table_path:
        import ast
        raw = torch.load(table_path, map_location="cpu")
        tables = {int(i): {ast.literal_eval(k): v for k, v in blk.items()} for i, blk in raw.items()}
    for index, block in enumerate(self.transformer_blocks):
        block.adaln_proj.enable_host_cache(None if tables is None else tables[index])
    self._adaln_host_cache = True
    self._adaln_dumped_entries = -1
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."""
    late = self.__dict__.get("_splice_late")
    if late is not None and get_forward_context().current_timestep >= self._splice_from_step:
        return late(hidden_states, audio_hidden_states, encoder_hidden_states, timestep, timestep_indices,
                    token_tags, position_ids, video_indices, audio_indices, text_indices)
    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)

    if getattr(self, "_adaln_host_cache", False):
        # One host read per forward keys every block's modulation cache by the timestep values.
        key = (tuple(timestep.reshape(-1).tolist()), tuple(temb.shape), str(temb.dtype))
        for block in self.transformer_blocks:
            block.adaln_proj._cache_key = key
        self._move_adaln_tables(temb.device)

    # 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):
        if STAGES.enabled:
            logger.info("H3_MEMORY_BLOCK %d allocated=%.3f GiB reserved=%.3f GiB", block_index,
                        torch.cuda.memory_allocated() / 2**30, torch.cuda.memory_reserved() / 2**30)
        with nvtx_range(f"minimax_h3.transformer_block.{block_index}"), STAGES.span("block_total"):
            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)

    if getattr(self, "_adaln_host_cache", False):
        self._maybe_dump_adaln_tables()
    if STAGES.enabled:
        stages = STAGES.flush()
        if not model_parallel_is_initialized() or get_sp_group().rank_in_group == 0:
            logger.info("H3_STAGE_MS %s", json.dumps(stages))
    return video_output, audio_output
fastvideo.models.dits.minimax_h3.MiniMaxH3Transformer3DModel.fuse_pdd_block
fuse_pdd_block(start: int, end: int, integration_weights: Mapping[str, Tensor], precision_decoding: dtype) -> Iterator[None]

Fuse fine-grid block [start, end) on both widened heads for the enclosed forwards.

Source code in fastvideo/models/dits/minimax_h3.py
@contextlib.contextmanager
def fuse_pdd_block(
    self,
    start: int,
    end: int,
    integration_weights: Mapping[str, torch.Tensor],
    precision_decoding: torch.dtype,
) -> Iterator[None]:
    """Fuse fine-grid block ``[start, end)`` on both widened heads for the enclosed forwards."""
    with fuse_pdd_heads(self.pdd_linears, start, end, integration_weights, precision_decoding):
        yield
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 and tensorize its layer identity so repeated blocks share one Dynamo graph. 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 and tensorize its
    layer identity so repeated blocks share one Dynamo graph. 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] = []
    prepared_vsa_impls = 0
    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)
        prepare_vsa = getattr(attention.distributed_attention.attn_impl, "prepare_for_compile", None)
        if callable(prepare_vsa):
            prepare_vsa(self._compile_setup_device(attention))
            prepared_vsa_impls += 1
    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 prepared_vsa_impls:
        logger.info("Prepared %d MiniMax H3 VSA layer indices for torch.compile", prepared_vsa_impls)
    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
        unsupported = prepare_vsa(self._compile_setup_device(attention))
        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, exact_kernels: frozenset[str] = frozenset())

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,
    exact_kernels: frozenset[str] = frozenset(),
) -> 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,
        exact_rope="rope" in exact_kernels,
    )
    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,
        exact_swiglu="swiglu" in exact_kernels,
    )
    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
    self.exact_modulate = "modulate" in exact_kernels

Functions: