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) of cu_seqlens and every head h:

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 any tokens and 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 bfloat16 CUDA tensors of shape [tokens, heads, 128] (packed THD layout; MiniMax-H3 uses 56 heads).

  • k (torch.Tensor) – Contiguous bfloat16 CUDA tensors of shape [tokens, heads, 128] (packed THD layout; MiniMax-H3 uses 56 heads).

  • v (torch.Tensor) – Contiguous bfloat16 CUDA tensors of shape [tokens, heads, 128] (packed THD layout; MiniMax-H3 uses 56 heads).

  • cu_seqlens (torch.Tensor) – int32 tensor of shape [segments + 1] on the same device with cu_seqlens[0] = 0, non-decreasing entries and cu_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