Skip to content

exact

Bit-exact Triton replacements for MiniMax-H3's eager RoPE, AdaLN and SwiGLU.

Each kernel performs exactly the BF16-rounded operation sequence the eager PyTorch expression performs, in the same order, so its output equals the eager output bit for bit. That is the difference from the Sol-Engine fusions in this package (FASTVIDEO_MINIMAX_H3_FUSIONS), which keep the whole expression in FP32 and round once at the store: faster still, but not the eager bits.

The kernels gain by fusing the eager expression's separate elementwise launches (and the index_select gathers of the AdaLN tables) into one pass over the activations, not by changing the arithmetic. Two details keep the rounding exact:

  • every FP32 multiply, add and divide is issued as explicit round-to-nearest PTX (mul.rn/add.rn/div.rn), so the compiler cannot contract a multiply and an add into one FMA, or fold the FP32 result and the BF16 downcast into one BF16 instruction, either of which rounds once where eager rounds twice;
  • every intermediate eager materializes as a BF16 tensor is rounded to BF16 (round-to-nearest-even) at the same point.

Selected with FASTVIDEO_MINIMAX_H3_EXACT_KERNELS. Every entry point has a supports_* predicate; callers fall back to the eager expression for inputs outside it.

Functions:

fastvideo.models.dits.minimax_h3_fusions.exact.gate_residual

gate_residual(hidden: Tensor, gate: Tensor, update: Tensor, indices: Tensor) -> Tensor

hidden + gate.index_select(0, indices) * update.

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def gate_residual(hidden: torch.Tensor, gate: torch.Tensor, update: torch.Tensor,
                  indices: torch.Tensor) -> torch.Tensor:
    """``hidden + gate.index_select(0, indices) * update``."""
    h, y = _rows(hidden), _rows(update)
    indices = indices.contiguous()
    out = torch.empty_like(h)
    cols = h.shape[1]
    _gate_residual_kernel[(h.shape[0], triton.cdiv(cols, _BLOCK_C))](h,
                                                                    gate,
                                                                    y,
                                                                    indices,
                                                                    out,
                                                                    cols,
                                                                    h.stride(0),
                                                                    y.stride(0),
                                                                    gate.stride(0),
                                                                    out.stride(0),
                                                                    BLOCK_C=_BLOCK_C)
    return out.view(hidden.shape)

fastvideo.models.dits.minimax_h3_fusions.exact.modulate

modulate(normed: Tensor, scale: Tensor, shift: Tensor, indices: Tensor) -> Tensor

normed * (1.0 + scale.index_select(0, indices)) + shift.index_select(0, indices).

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def modulate(normed: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor, indices: torch.Tensor) -> torch.Tensor:
    """``normed * (1.0 + scale.index_select(0, indices)) + shift.index_select(0, indices)``."""
    n = _rows(normed)
    indices = indices.contiguous()
    out = torch.empty_like(n)
    cols = n.shape[1]
    _modulate_kernel[(n.shape[0], triton.cdiv(cols, _BLOCK_C))](n,
                                                               scale,
                                                               shift,
                                                               indices,
                                                               out,
                                                               cols,
                                                               n.stride(0),
                                                               scale.stride(0),
                                                               shift.stride(0),
                                                               out.stride(0),
                                                               BLOCK_C=_BLOCK_C)
    return out.view(normed.shape)

fastvideo.models.dits.minimax_h3_fusions.exact.rope_prefix

rope_prefix(hidden_states: Tensor, rotary_emb: tuple[Tensor, Tensor]) -> Tensor

Eager MiniMaxH3Attention._apply_rotary_emb for contiguous BF16 [1, S, H, D].

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def rope_prefix(hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor:
    """Eager ``MiniMaxH3Attention._apply_rotary_emb`` for contiguous BF16 ``[1, S, H, D]``."""
    cos, sin = rotary_emb
    cos = cos.to(hidden_states.dtype).contiguous()
    sin = sin.to(hidden_states.dtype).contiguous()
    _, seq, heads, dim = hidden_states.shape
    rotary = cos.shape[-1]
    out = torch.empty_like(hidden_states)
    block_h = 4
    _rope_prefix_kernel[(seq, triton.cdiv(heads, block_h))](
        hidden_states,
        cos,
        sin,
        out,
        heads,
        hidden_states.stride(1),
        hidden_states.stride(2),
        cos.stride(0),
        D=dim,
        R=rotary,
        RP=triton.next_power_of_2(rotary),
        TAIL=triton.next_power_of_2(max(dim - rotary, 1)),
        BLOCK_H=block_h,
    )
    return out

fastvideo.models.dits.minimax_h3_fusions.exact.supports_rope

supports_rope(hidden_states: Tensor, rotary_emb: tuple[Tensor, Tensor] | None) -> bool

Whether rope_prefix reproduces eager _apply_rotary_emb on these inputs.

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def supports_rope(hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor] | None) -> bool:
    """Whether ``rope_prefix`` reproduces eager ``_apply_rotary_emb`` on these inputs."""
    if rotary_emb is None or not HAVE_TRITON or torch.is_grad_enabled() or torch.compiler.is_compiling():
        return False
    cos, sin = rotary_emb
    dim = hidden_states.shape[-1]
    return (hidden_states.is_cuda and hidden_states.dtype == torch.bfloat16 and hidden_states.dim() == 4
            and hidden_states.shape[0] == 1 and hidden_states.is_contiguous() and dim & (dim - 1) == 0
            and cos.dim() == 2 and cos.shape == sin.shape and cos.shape[0] == hidden_states.shape[1]
            and cos.shape[-1] % 2 == 0 and cos.shape[-1] <= dim)

fastvideo.models.dits.minimax_h3_fusions.exact.supports_rowwise

supports_rowwise(*tensors: Tensor) -> bool

Whether modulate/gate_residual/swiglu reproduce eager on these tensors.

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def supports_rowwise(*tensors: torch.Tensor) -> bool:
    """Whether ``modulate``/``gate_residual``/``swiglu`` reproduce eager on these tensors."""
    return (HAVE_TRITON and not torch.is_grad_enabled() and not torch.compiler.is_compiling()
            and all(t.is_cuda and t.dtype == torch.bfloat16 and t.stride(-1) == 1 and _flattens_to_rows(t)
                    for t in tensors))

fastvideo.models.dits.minimax_h3_fusions.exact.swiglu

swiglu(packed: Tensor) -> Tensor

value, gate = packed.chunk(2, -1); value * F.silu(gate).

Source code in fastvideo/models/dits/minimax_h3_fusions/exact.py
def swiglu(packed: torch.Tensor) -> torch.Tensor:
    """``value, gate = packed.chunk(2, -1); value * F.silu(gate)``."""
    x = _rows(packed)
    cols = x.shape[1] // 2
    out = torch.empty((x.shape[0], cols), device=x.device, dtype=x.dtype)
    _swiglu_kernel[(x.shape[0], triton.cdiv(cols, _BLOCK_C))](x,
                                                             out,
                                                             cols,
                                                             x.stride(0),
                                                             out.stride(0),
                                                             BLOCK_C=_BLOCK_C)
    return out.view(*packed.shape[:-1], cols)