flashinfer.diffusion_ops.minimax_h3_nvfp4_pre_attention¶
- flashinfer.diffusion_ops.minimax_h3_nvfp4_pre_attention(x: Tensor, x_norm_weight: Tensor, adaln_scale: Tensor, adaln_shift: Tensor, adaln_index: Tensor, qkv_weight_q: Tensor, qkv_weight_sf: Tensor, qkv_weight_global_scale: float | Tensor, act_global_scale: Tensor, q_norm_weight: Tensor, k_norm_weight: Tensor, rope_cos_sin: Tensor, *, eps: float = 1e-05, out_mode: str = 'bf16', alpha: float | None = None, 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_sf: Tensor | None = None) MiniMaxH3PreAttentionOutput¶
Fused NVFP4 MiniMax-H3 pre-attention for SM120 (RTX 5090 / RTX PRO 6000 Blackwell).
Same pre-norm, AdaLN, GEMM epilogue and output formats as
minimax_h3_fp8_pre_attention(), with NVFP4 operands in FlashInfer’s conventions: the normalized activationais quantized per token likefp4_quantize(a, act_global_scale, sf_vec_size=16, is_sf_swizzled_layout=False)(E2M1 codesuint8 [M, 2688]+ UE4M3 block scalesuint8 [M, 336]), the weight comes fromquantize_minimax_h3_qkv_weight_nvfp4()(128x4 swizzled scales) and the accumulator is rescaled byalpha = 1 / (act_global_scale * qkv_weight_global_scale).- 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) – Outputs of
quantize_minimax_h3_qkv_weight_nvfp4().qkv_weight_sf (torch.Tensor) – Outputs of
quantize_minimax_h3_qkv_weight_nvfp4().qkv_weight_global_scale (Scalar) – NVFP4 weight global scale as a float or single-element tensor.
act_global_scale (torch.Tensor) – FP32
[1]CUDA tensor448 * 6 / amaxof the calibrated normalized activation (nvfp4_global_scale_from_amax()); read on the device, no host synchronization.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]tensor storing[cos(48), sin(48)]per token.eps (float) – RMSNorm epsilon for both the pre-norm and the Q/K norms.
out_mode (str) – Output format:
"bf16","e4m3", or"nvfp4". Shapes and quantization conventions matchminimax_h3_fp8_pre_attention().alpha (Optional[float]) –
1 / (act_global_scale * qkv_weight_global_scale). Pass it explicitly to avoid the host synchronization needed to read the global scales; derived from them when omitted.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 NVFP4 scale outputs; used only when
out_mode == "nvfp4"and ignored for other output modes. Allocated when omitted in NVFP4 mode.k_sf (Optional[torch.Tensor]) – Optional pre-allocated NVFP4 scale outputs; used only when
out_mode == "nvfp4"and ignored for other output modes. Allocated when omitted in NVFP4 mode.v_sf (Optional[torch.Tensor]) – Optional pre-allocated NVFP4 scale outputs; used only when
out_mode == "nvfp4"and ignored for other output modes. Allocated when omitted in NVFP4 mode.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 (packed E2M1
uint8 [M, 2688]and UE4M3 scalesuint8 [M, 336]); allocated when omitted.act_sf (Optional[torch.Tensor]) – Optional stage-1 workspaces (packed E2M1
uint8 [M, 2688]and UE4M3 scalesuint8 [M, 336]); allocated when omitted.
- Returns:
(q, k, v, q_sf, k_sf, v_sf); the scale entries areNoneunlessout_mode == "nvfp4".- Return type:
MiniMaxH3PreAttentionOutput