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 M tokens (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 with rope_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] from quantize_minimax_h3_qkv_weight_fp8(). Output column h * 384 + kind * 128 + d is channel d of head h for kind 0 = Q, 1 = K, 2 = V.

  • qkv_weight_scale (torch.Tensor) – float8_e4m3fn [21504, 5376] and FP32 [21504] from quantize_minimax_h3_qkv_weight_fp8(). Output column h * 384 + kind * 128 + d is channel d of head h for kind 0 = 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] storing RN(value / descale) with the caller’s per-tensor q_descale, k_descale, v_descale. "nvfp4": packed E2M1 uint8 [M, 56, 64] plus UE4M3 block-16 scales uint8 [M, 56, 8] (row-major) with the caller’s per-tensor *_global_scale (FlashInfer fp4_quantize semantics).

  • 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 are None unless out_mode == "nvfp4".

Return type:

MiniMaxH3PreAttentionOutput