Skip to content

minimax_h3_vsa_fp4

Inference fast path: VSA-H3 attention on the block-sparse FP4 kernel.

Opt-in with FASTVIDEO_H3_VSA_FP4=1 (no-grad, single sequence-parallel rank, fastvideo-kernel built with attn_qat_infer). The selection is VSA-H3's own: tile pooling, top-k block mask and the gated compression branch are unchanged; only the block-sparse attention itself runs on SageAttention3's FP4 kernel (BF16 Triton otherwise), with 64-token tiles carried by quadrant masks on the kernel's 128x128 blocks.

The attention input is gathered into tile order once per block (one hidden_size-wide pass; pad rows stay zero, so q/k/v pad rows are exactly zero through the bias-free projections, RMSNorm and RoPE). That replaces the generic path's concat, four tile scatters and three transposes, and lets q, k and v share one activation quantization. The output returns to packed order with one gather before to_out.

Functions:

fastvideo.models.dits.minimax_h3_vsa_fp4.vsa_fp4_attention

vsa_fp4_attention(attn: Any, hidden_states: Tensor, rotary_emb: tuple[Tensor, Tensor], meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> Tensor

Attention core for MiniMaxH3Attention; returns the pre-to_out [B, L, H*D].

Source code in fastvideo/models/dits/minimax_h3_vsa_fp4.py
def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor],
                      meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor:
    """Attention core for ``MiniMaxH3Attention``; returns the pre-``to_out`` ``[B, L, H*D]``."""
    api = _api()
    layout = _layout_for(meta, rotary_emb)
    heads, dim = attn.num_attention_heads, attn.attention_head_dim
    with STAGES.span("qkv_proj_rope"):
        x_tiles = layout.gather_in(hidden_states)
        query, key, value = (t.unflatten(-1, (heads, dim))
                             for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), x_tiles))
        if use_fused_rope:
            from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope
            cos, sin = layout.cos.to(query.dtype), layout.sin.to(query.dtype)
            query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps)
            key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps)
        else:
            rope = (layout.cos, layout.sin)
            query = attn._apply_rotary_emb(attn.norm_q(query), rope)
            key = attn._apply_rotary_emb(attn.norm_k(key), rope)

    sim_fp8 = envs.FASTVIDEO_H3_SIM_SP_FP8.get()
    if sim_fp8:
        query, key, value = (_fp8_roundtrip(t) for t in (query, key, value))

    vbs = meta.variable_block_sizes
    logical = layout.n_tiles * layout.tile
    with STAGES.span("select_mask"):
        q_pooled = _pool_tiles(query[:, :logical], vbs, layout.tile)
        k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile)
        scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (dim**0.5)
        sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity
        mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, meta.video_tile_spans,
                                 meta.span_sparsities)
        q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs)
    with STAGES.span("fp4_attention"):
        # Lists come from vsa_tile_mask_to_fp4_blocks and are in range; skip the per-call host sync.
        out = _sparse_fp4_attention(api, query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad)
    with STAGES.span("out_untile"):
        out = out.transpose(1, 2).index_select(1, layout.untile)  # [B, L, H, D], packed order
    if sim_fp8:
        out = _fp8_roundtrip(out)

    if attn.to_gate_compress is not None and attn._gate_active():
        with STAGES.span("gate_compress"):
            gate, _ = attn.to_gate_compress(hidden_states)
            v_pooled = _pool_tiles(value[:, :logical], vbs, layout.tile)
            out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled).permute(0, 2, 1, 3).to(out.dtype)
            out = out.addcmul_(out_c.index_select(1, layout.row_tile), gate.unflatten(-1, (heads, dim)))
    return out.flatten(2, 3)

fastvideo.models.dits.minimax_h3_vsa_fp4.vsa_fp4_attention_sp

vsa_fp4_attention_sp(attn: Any, hidden_states: Tensor, rotary_emb: tuple[Tensor, Tensor], meta: MiniMaxH3VSAMetadata, use_fused_rope: bool, sp_group: Any) -> Tensor

Ulysses-SP attention core on local sequence rows [1, rows, C]; returns pre-to_out [1, rows, H*D].

Source code in fastvideo/models/dits/minimax_h3_vsa_fp4.py
def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor],
                         meta: MiniMaxH3VSAMetadata, use_fused_rope: bool, sp_group: Any) -> torch.Tensor:
    """Ulysses-SP attention core on local sequence rows ``[1, rows, C]``; returns pre-``to_out`` ``[1, rows, H*D]``."""
    import torch.distributed as dist

    api = _api()
    world, rank = sp_group.world_size, sp_group.rank_in_group
    heads, dim = attn.num_attention_heads, attn.attention_head_dim
    local_rows = hidden_states.shape[1]
    layout = getattr(meta, "_h3_fp4_sp_layout", None)
    if layout is None:
        layout = _SPTileLayout(meta, rank, local_rows)
        meta._h3_fp4_sp_layout = layout  # type: ignore[attr-defined]

    with STAGES.span("qkv_proj_rope"):
        query, key, value = (t.unflatten(-1, (heads, dim))
                             for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), hidden_states))
        if use_fused_rope:
            from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope
            cos, sin = rotary_emb[0].to(query.dtype), rotary_emb[1].to(query.dtype)
            query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps)
            key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps)
        else:
            query = attn._apply_rotary_emb(attn.norm_q(query), rotary_emb)
            key = attn._apply_rotary_emb(attn.norm_k(key), rotary_emb)

    with STAGES.span("qkv_pack"):
        payload, scale = _pack_heads_fp8(query[0], key[0], value[0], world)
    with STAGES.span("qkv_all_to_all"):
        payload, scale = _all_to_all(payload, scale, sp_group.device_group)
    with STAGES.span("qkv_unpack_tile"):
        qkv = layout.tiles_from(_unpack_seq_fp8(payload, scale, layout.seq_len))  # [3, R, Hs, D]
    q_t, k_t, v_t = qkv[0:1], qkv[1:2], qkv[2:3]

    vbs = meta.variable_block_sizes
    logical = layout.n_tiles * layout.tile
    with STAGES.span("select_mask"):
        scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile),
                              _pool_tiles(k_t[:, :logical], vbs, layout.tile).transpose(-2, -1)) / (dim**0.5)
        sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity
        mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, meta.video_tile_spans,
                                 meta.span_sparsities)
        q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs)
    with STAGES.span("fp4_attention"):
        out_bhsd = _sparse_fp4_attention(api, q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad)

    with STAGES.span("out_pack"):
        payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows)
    with STAGES.span("out_all_to_all"):
        payload, scale = _all_to_all(payload, scale, sp_group.device_group)
    with STAGES.span("out_unpack"):
        out = _unpack_heads_fp8(payload, scale)  # [rows, H, D]

    if attn.to_gate_compress is not None and attn._gate_active():
        with STAGES.span("gate_compress"):
            v_pooled = _pool_tiles(v_t[:, :logical], vbs, layout.tile)
            out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)[0].to(out.dtype)  # [Hs, n_tiles, D]
            gathered = torch.empty((world, *out_c.shape), dtype=out_c.dtype, device=out_c.device)
        with STAGES.span("gate_all_gather"):
            dist.all_gather_into_tensor(gathered, out_c.contiguous(), group=sp_group.device_group)
        with STAGES.span("gate_apply"):
            out_c_all = gathered.flatten(0, 1).transpose(0, 1)  # [n_tiles, H, D]
            gate, _ = attn.to_gate_compress(hidden_states)
            out = _apply_gate(out, out_c_all, layout.local_row_tile, gate[0].unflatten(-1, (heads, dim)))
    return out.flatten(1, 2).unsqueeze(0)

fastvideo.models.dits.minimax_h3_vsa_fp4.vsa_tile_first_attention

vsa_tile_first_attention(attn: Any, hidden_states: Tensor, rotary_emb: tuple[Tensor, Tensor], meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> Tensor

Single-rank BF16 VSA with one input scatter instead of a Q/K/V/gate stack.

The existing backend computes the same tile-64 mask, valid-key handling, fine attention and compression branch. Bias-free projections keep pad rows zero. This path is inference-only and keeps the checkpoint layout.

Source code in fastvideo/models/dits/minimax_h3_vsa_fp4.py
def vsa_tile_first_attention(attn: Any, hidden_states: torch.Tensor,
                             rotary_emb: tuple[torch.Tensor, torch.Tensor],
                             meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor:
    """Single-rank BF16 VSA with one input scatter instead of a Q/K/V/gate stack.

    The existing backend computes the same tile-64 mask, valid-key handling,
    fine attention and compression branch. Bias-free projections keep pad
    rows zero. This path is inference-only and keeps the checkpoint layout.
    """
    layout = _layout_for(meta, rotary_emb)
    logical = layout.n_tiles * layout.tile
    heads, dim = attn.num_attention_heads, attn.attention_head_dim
    with STAGES.span("tile_input"):
        x_tiles = layout.gather_in(hidden_states)[:, :logical]
    with STAGES.span("qkv_proj"):
        query, key, value = (t.unflatten(-1, (heads, dim))
                             for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), x_tiles))
    with STAGES.span("qknorm_rope"):
        cos, sin = layout.cos[:logical], layout.sin[:logical]
        if use_fused_rope:
            from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope
            query = fused_qknorm_rope(query, attn.norm_q.weight, cos.to(query.dtype), sin.to(query.dtype), attn.norm_q.eps)
            key = fused_qknorm_rope(key, attn.norm_k.weight, cos.to(key.dtype), sin.to(key.dtype), attn.norm_k.eps)
        else:
            query = attn._apply_rotary_emb(attn.norm_q(query), (cos, sin))
            key = attn._apply_rotary_emb(attn.norm_k(key), (cos, sin))
    gate = None
    if attn.to_gate_compress is not None and attn._gate_active():
        with STAGES.span("gate_proj"):
            gate, _ = attn.to_gate_compress(x_tiles)
            gate = gate.unflatten(-1, (heads, dim))
    capture_root = envs.FASTVIDEO_H3_CAPTURE_QKV.get()
    if capture_root and attn._layer_idx in (0, 20, 41):
        from pathlib import Path
        root = Path(capture_root)
        root.mkdir(parents=True, exist_ok=True)
        capture = root / f"layer-{attn._layer_idx}.pt"
        if not capture.exists():
            q_pooled = _pool_tiles(query, meta.variable_block_sizes, meta.tile_elems)
            k_pooled = _pool_tiles(key, meta.variable_block_sizes, meta.tile_elems)
            scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / dim**0.5
            sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity
            mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt,
                                     meta.video_tile_spans, meta.span_sparsities)
            # Two heads keep the artifact small while retaining all real keys,
            # query rows, per-tile selections and partial-tile validity.
            torch.save({"q": query[:, :, :2].transpose(1, 2).contiguous().cpu(),
                        "k": key[:, :, :2].transpose(1, 2).contiguous().cpu(),
                        "v": value[:, :, :2].transpose(1, 2).contiguous().cpu(),
                        "mask": mask[:, :2].cpu(), "vbs": meta.variable_block_sizes.cpu(),
                        "untile": meta.untile_combined_index.cpu()}, capture)
            del q_pooled, k_pooled, scores, mask
    with STAGES.span("attention"):
        out = attn.distributed_attention.attn_impl.forward(query, key, value, gate, meta)
    with STAGES.span("untile_output"):
        return out.index_select(1, layout.untile).flatten(2, 3)