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) 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 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 term qm K^T is 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 run mma.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 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_nvfp4) shared with minimax_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’s 0.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 mxf4nvf4 MMA 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 precomputed qm K^T compensation 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 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