flashinfer.diffusion_ops.minimax_h3_sm120_varlen_attention_fp8¶
- flashinfer.diffusion_ops.minimax_h3_sm120_varlen_attention_fp8(q: Tensor, k: Tensor, v: Tensor, cu_seqlens: Tensor, out: Tensor | None = None, *, cu_seqlens_host: Sequence[int] | None = None, softmax_scale: float | None = None) Tensor¶
FP8 (E4M3) non-causal packed-varlen self-attention for MiniMax-H3 on SM120 (GB202).
Computes, for every segment
[a, b)ofcu_seqlensand every headh:out[a:b, h] = softmax(q[a:b, h] @ k[a:b, h]^T * softmax_scale) @ v[a:b, h]
with FP8 E4M3 tensor-core operands and an FP32 softmax: Q is quantized per token, K per 128-key block after subtracting its segment’s per-channel mean (an exact softmax shift), the probabilities are stored as E4M3 with a 2^8 exponent bias and V per (segment, channel). Both QK^T and PV run
mma.sync.m16n8k32 kind::f8f6f4; the output is rounded to BF16 once. One runtime-variable kernel set serves anytokensand any segment lengths (including empty segments and lengths below one tile); the quantized operands live in a grow-only per-device workspace (workspace_bytes).- Parameters:
q (torch.Tensor) – Contiguous
bfloat16CUDA tensors of shape[tokens, heads, 128](packed THD layout; MiniMax-H3 uses 56 heads).k (torch.Tensor) – Contiguous
bfloat16CUDA tensors of shape[tokens, heads, 128](packed THD layout; MiniMax-H3 uses 56 heads).v (torch.Tensor) – Contiguous
bfloat16CUDA tensors of shape[tokens, heads, 128](packed THD layout; MiniMax-H3 uses 56 heads).cu_seqlens (torch.Tensor) –
int32tensor of shape[segments + 1]on the same device withcu_seqlens[0] = 0, non-decreasing entries andcu_seqlens[-1] = tokens.out (Optional[torch.Tensor]) – Optional pre-allocated output of the same shape/dtype as
q; allocated when omitted.cu_seqlens_host (Optional[Sequence[int]]) – Host copy of
cu_seqlens(avoids a synchronizing device-to-host copy). The segment plan is cached per(cu_seqlens, heads, device).softmax_scale (Optional[float]) – Softmax scale; defaults to
1 / sqrt(128).
- Returns:
bfloat16[tokens, heads, 128]attention output.- Return type:
torch.Tensor