flashinfer.diffusion_ops.minimax_h3_fc1_swiglu_fp8

flashinfer.diffusion_ops.minimax_h3_fc1_swiglu_fp8(x: Tensor, x_norm_weight: Tensor, adaln_scale: Tensor, adaln_shift: Tensor, adaln_index: Tensor, fc1_weight_q: Tensor, fc1_weight_scale: Tensor, *, out: Tensor | None = None, workspace_q: Tensor | None = None, workspace_scale: Tensor | None = None, eps: float = 1e-05) → Tensor

Fused FP8 (W8A8) RMSNorm + indexed AdaLN + FC1 GEMM + SwiGLU of the MiniMax-H3 video DiT block for SM120 (RTX 5090 / RTX PRO 6000 Blackwell).

Computes, for each of the M rows (batch 1, no sequence parallelism):

n   = bf16(rmsnorm(x, eps) * x_norm_weight)                       # FP32 sum of squares, rsqrt
a   = bf16(adaln_shift[i] + n * bf16(1 + adaln_scale[i]))        # i = adaln_index[row]
a[i outside [0, 9)] = 0                                            # device-side guard
a_q = e4m3(a / s_row),  s_row = RN(amax_row(|a|) / 448)           # per-token activation scale
h   = bf16(a_q @ fc1_weight_q^T * s_row * fc1_weight_scale)       # FP32 accumulation, [M, 28672]
y   = bf16(bf16(silu(h_gate)) * h_up)                              # [M, 14336]

Two kernels run on the current stream: a one-CTA-per-row norm/AdaLN/quantization kernel that writes a_q and s_row into the workspaces, and a persistent mma.sync GEMM whose epilogue applies the dequantization scales and the SwiGLU.

Parameters:
  • x (torch.Tensor) – Contiguous bfloat16 [M, 5376] hidden states, 1 <= M <= 2**24.

  • x_norm_weight (torch.Tensor) – bfloat16 [5376] RMSNorm weight.

  • adaln_scale (torch.Tensor) – bfloat16 [9, 5376] AdaLN tables.

  • adaln_shift (torch.Tensor) – bfloat16 [9, 5376] AdaLN tables.

  • adaln_index (torch.Tensor) – int32 [M] table row per activation row.

  • fc1_weight_q (torch.Tensor) – float8_e4m3fn [28672, 5376] and float32 [28672] from prepare_minimax_h3_fc1_weight_fp8() (SM120 prepacked row order).

  • fc1_weight_scale (torch.Tensor) – float8_e4m3fn [28672, 5376] and float32 [28672] from prepare_minimax_h3_fc1_weight_fp8() (SM120 prepacked row order).

  • out (Optional[torch.Tensor]) – Optional bfloat16 [M, 14336] output (allocated when omitted).

  • workspace_q (Optional[torch.Tensor]) – Optional caller-owned float8_e4m3fn [M, 5376] buffer that receives a_q.

  • workspace_scale (Optional[torch.Tensor]) – Optional caller-owned float32 [M] buffer that receives the per-token scales.

  • eps (float) – RMSNorm epsilon (positive; the operator is validated at 1e-5).

Returns:

bfloat16 [M, 14336].

Return type:

torch.Tensor