flashinfer.diffusion_ops.minimax_h3_fp8_pre_attention¶
- flashinfer.diffusion_ops.minimax_h3_fp8_pre_attention(x: Tensor, x_norm_weight: Tensor, adaln_scale: Tensor, adaln_shift: Tensor, adaln_index: Tensor, qkv_weight_q: Tensor, qkv_weight_scale: Tensor, q_norm_weight: Tensor, k_norm_weight: Tensor, rope_cos_sin: Tensor, *, eps: float = 1e-05, out_mode: str = 'bf16', q: Tensor | None = None, k: Tensor | None = None, v: Tensor | None = None, q_sf: Tensor | None = None, k_sf: Tensor | None = None, v_sf: Tensor | None = None, q_descale: Tensor | float | None = None, k_descale: Tensor | float | None = None, v_descale: Tensor | float | None = None, q_global_scale: Tensor | float | None = None, k_global_scale: Tensor | float | None = None, v_global_scale: Tensor | float | None = None, act_q: Tensor | None = None, act_scale: Tensor | None = None) MiniMaxH3PreAttentionOutput¶
Fused FP8 (W8A8) MiniMax-H3 pre-attention for SM120 (RTX 5090 / RTX PRO 6000 Blackwell).
Computes, for each of the
Mtokens (batch 1, no sequence parallelism):n = bf16(rmsnorm(x, eps) * x_norm_weight) a = bf16(adaln_shift[i] + n * bf16(1 + adaln_scale[i])) # i = adaln_index[t] a_q = e4m3(a / s_t), s_t = amax_t(|a|) / 448 # per-token activation scale y = bf16(a_q @ qkv_weight_q^T * s_t * qkv_weight_scale) # [M, 56 heads x (Q, K, V) x 128] q = rope(bf16(rmsnorm(y_q, eps) * q_norm_weight)), k likewise with k_norm_weight, v = y_v
RoPE rotates channel pairs
(d, d + 48)of the first 96 channels withrope_cos_sin[t] = [cos(48), sin(48)]and passes channels 96..127 through. Both launches (norm/AdaLN/quantization and the fused persistent GEMM) run on the current stream.- Parameters:
x (torch.Tensor) – BF16
[M, 5376]hidden states.x_norm_weight (torch.Tensor) – BF16
[5376]pre-norm weight.adaln_scale (torch.Tensor) – BF16
[rows, 5376]AdaLN modulation tables.adaln_shift (torch.Tensor) – BF16
[rows, 5376]AdaLN modulation tables.adaln_index (torch.Tensor) – Int32
[M]indices selecting one AdaLN modulation row per token.qkv_weight_q (torch.Tensor) –
float8_e4m3fn [21504, 5376]and FP32[21504]fromquantize_minimax_h3_qkv_weight_fp8(). Output columnh * 384 + kind * 128 + dis channeldof headhforkind0 = Q, 1 = K, 2 = V.qkv_weight_scale (torch.Tensor) –
float8_e4m3fn [21504, 5376]and FP32[21504]fromquantize_minimax_h3_qkv_weight_fp8(). Output columnh * 384 + kind * 128 + dis channeldof headhforkind0 = Q, 1 = K, 2 = V.q_norm_weight (torch.Tensor) – BF16
[128]per-head RMSNorm weights.k_norm_weight (torch.Tensor) – BF16
[128]per-head RMSNorm weights.rope_cos_sin (torch.Tensor) – BF16
[M, 96]=[cos(48), sin(48)]per token.eps (float) – RMSNorm epsilon for both the pre-norm and the Q/K norms.
out_mode (str) –
"bf16": Q/K/V BF16[M, 56, 128]."e4m3":float8_e4m3fn [M, 56, 128]storingRN(value / descale)with the caller’s per-tensorq_descale,k_descale,v_descale."nvfp4": packed E2M1uint8 [M, 56, 64]plus UE4M3 block-16 scalesuint8 [M, 56, 8](row-major) with the caller’s per-tensor*_global_scale(FlashInferfp4_quantizesemantics).q (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
k (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
v (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
q_sf (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
k_sf (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
v_sf (Optional[torch.Tensor]) – Optional pre-allocated outputs; allocated when omitted.
q_descale (Optional[Scalar]) – Per-tensor E4M3 dequantization scales required when
out_mode == "e4m3".k_descale (Optional[Scalar]) – Per-tensor E4M3 dequantization scales required when
out_mode == "e4m3".v_descale (Optional[Scalar]) – Per-tensor E4M3 dequantization scales required when
out_mode == "e4m3".q_global_scale (Optional[Scalar]) – NVFP4 global scales required when
out_mode == "nvfp4".k_global_scale (Optional[Scalar]) – NVFP4 global scales required when
out_mode == "nvfp4".v_global_scale (Optional[Scalar]) – NVFP4 global scales required when
out_mode == "nvfp4".act_q (Optional[torch.Tensor]) – Optional stage-1 workspaces (
float8_e4m3fn [M, 5376], FP32[M]); allocated when omitted.act_scale (Optional[torch.Tensor]) – Optional stage-1 workspaces (
float8_e4m3fn [M, 5376], FP32[M]); allocated when omitted.
- Returns:
(q, k, v, q_sf, k_sf, v_sf); the scale entries areNoneunlessout_mode == "nvfp4".- Return type:
MiniMaxH3PreAttentionOutput