Skip to content

minimax_h3_sparse_int8

Experimental sm89 tile-64 VSA with BF16 or INT8 QK and FP32 accumulation.

Retains each query tile's original key selection and masks partial key tiles. Unlike a 128-query adapter, it adds no attention blocks. Q/K use per-token scales; K centering is a softmax-invariant shift. V uses one scale per head and channel, so its dequantization can be applied once in the epilogue. PV stays in BF16: FP8 PV had excessive error on real H3 inputs. Numerical validation and same-seed clip review are required before enabling.

Functions:

fastvideo.attention.backends.minimax_h3_sparse_int8.sparse_sm89_attention

sparse_sm89_attention(q: Tensor, k: Tensor, v: Tensor, mask: Tensor, vbs: Tensor, *, int8_qk: bool = True) -> Tensor

Forward-only [B,H,S,128] BF16 attention on sm89, with 64-token tiles.

Source code in fastvideo/attention/backends/minimax_h3_sparse_int8.py
def sparse_sm89_attention(q: torch.Tensor,
                          k: torch.Tensor,
                          v: torch.Tensor,
                          mask: torch.Tensor,
                          vbs: torch.Tensor,
                          *,
                          int8_qk: bool = True) -> torch.Tensor:
    """Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles."""
    if torch.is_grad_enabled():
        raise ValueError("Sparse INT8/FP8 attention is inference-only")
    if not q.is_cuda or torch.cuda.get_device_capability(q.device) != (8, 9):
        raise ValueError("Sparse INT8/FP8 attention requires sm89 CUDA")
    if q.dtype != torch.bfloat16 or q.shape[-1] != 128 or q.shape != k.shape or q.shape != v.shape:
        raise ValueError("Sparse INT8/FP8 attention requires matching BF16 Q/K/V with head dimension 128")
    b, h, length, dim = q.shape
    if length != vbs.numel() * 64 or mask.shape != (b, h, length // 64, length // 64):
        raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector")
    from fastvideo_kernel.triton_kernels.index import map_to_index

    # The production INT8-QK/BF16-PV route reads BSHD-backed views directly.
    # Quantized Q/K and the output remain contiguous BHSD. Other ablations
    # retain their established layout and arithmetic.
    if not int8_qk:
        q, k, v = q.contiguous(), k.contiguous(), v.contiguous()
    vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous()
    grid = (triton.cdiv(length, 16), b * h)
    qi, ki = q, k
    qs, ks = q, k  # unused pointers in the BF16-QK ablation
    if int8_qk:
        # Tile pads are zero by contract; avoid a full FP32 copy for the reduction.
        # Preserve the exact reduction used by the old contiguous adapter;
        # its temporary copy dies before Q/K quantization and fine attention.
        mean = k.contiguous().sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1)
        qi = torch.empty(q.shape, device=q.device, dtype=torch.int8)
        ki = torch.empty(k.shape, device=k.device, dtype=torch.int8)
        qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32)
        ks = torch.empty_like(qs)
        _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, h, *q.stride(), CENTER=False, ROWS=16, num_warps=4)
        _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, h, *k.stride(), CENTER=True, ROWS=16, num_warps=4)
    index, count = map_to_index(mask.contiguous())
    out = torch.empty(q.shape, device=q.device, dtype=q.dtype)
    _sparse_int8_fp8[(length // 64, b * h)](qi,
                                            ki,
                                            v,
                                            qs,
                                            ks,
                                            index,
                                            count,
                                            vbs,
                                            out,
                                            length,
                                            dim,
                                            h,
                                            *v.stride(),
                                            INT8_QK=int8_qk)
    return out