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 activation a is quantized per token like fp4_quantize(a, act_global_scale, sf_vec_size=16, is_sf_swizzled_layout=False) (E2M1 codes uint8 [M, 2688] + UE4M3 block scales uint8 [M, 336]), the weight comes from quantize_minimax_h3_qkv_weight_nvfp4() (128x4 swizzled scales) and the accumulator is rescaled by alpha = 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 tensor 448 * 6 / amax of 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 match minimax_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 scales uint8 [M, 336]); allocated when omitted.

  • act_sf (Optional[torch.Tensor]) – Optional stage-1 workspaces (packed E2M1 uint8 [M, 2688] and UE4M3 scales uint8 [M, 336]); allocated when omitted.

Returns:

(q, k, v, q_sf, k_sf, v_sf); the scale entries are None unless out_mode == "nvfp4".

Return type:

MiniMaxH3PreAttentionOutput