flashinfer.diffusion_ops.minimax_h3_sm120_varlen_attention_nvfp4¶
- flashinfer.diffusion_ops.minimax_h3_sm120_varlen_attention_nvfp4(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¶
Experimental NVFP4 non-causal packed-varlen self-attention for MiniMax-H3 on SM120 (GB202), following the SageAttention3 FP4 recipe.
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 block-scaled NVFP4 tensor-core operands for both the scores and the value product and an FP32 softmax, as in SageAttention3 (
sageattention3_blackwell): K rows are centred by the segment’s per-channel mean and Q rows by their 128-row block’s mean (the block-mean termqm K^Tis added back inside the kernel, so the softmax is unchanged), Q / K / V are stored as E2M1 codes with one UE4M3 scale per 16 elements (block amax / 6; V transposed with per-16-key scales), QK^T and PV both runmma.sync.m16n8k64 kind::mxf4nvf4 block_scale scale_vec::4X, and the probabilities use Sage3’s two-level scheme (row maximum brought to 448 * 6, one UE4M3 scale per 16 keys, E2M1 codes); 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_nvfp4) shared withminimax_h3_sm120_varlen_attention_fp8().The E2M1 scores and probabilities carry a larger quantization error than the FP8 operator (about 4x the relative L2 error on Gaussian inputs, the same as SageAttention3’s own); this route is an experimental precision/latency trade-off and is validated against the FP32 oracle with the FP4 block-scaled tolerance
atol = 1.0, rtol = 0.1(see the tests), not the FP8 operator’s0.1.Performance ceiling of the attention launch (measured on RTX 5090 and RTX PRO 6000 Blackwell, both GB202): the kernel needs 246 registers per thread, so one 8-warp CTA is resident per SM (two warps per SM sub-partition) and each warp alternates a ~1400-cycle
mxf4nvf4MMA burst with a ~2700-cycle FP32 softmax / P-quantization phase that only the other warp’s burst can overlap. The tensor pipe is therefore active ~58 % of the time and the launch runs at ~60 % of the tensor-pipe floor on production plans (~55 % on 4096-token plans), against 84-88 % for the MMA-only skeleton; L2 and DRAM are idle (L2 hit > 99 %). Ping-pong barrier placement, softmax emission order, FMA-pipe exp2, power-of-two P block scales, packed f16x2 P conversions, a precomputedqm K^Tcompensation table, K/V multicast / DSM sharing, a 64-key 3-CTA geometry and split-KV were each measured or bounded on both SKUs and none is faster within the FP4 error budget, so the attention kernel is unchanged; the softmax dependency chain under the two-warps-per-sub-partition register budget is the binding resource.- 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