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
Mrows (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_qands_rowinto the workspaces, and a persistentmma.syncGEMM 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]andfloat32[28672]fromprepare_minimax_h3_fc1_weight_fp8()(SM120 prepacked row order).fc1_weight_scale (torch.Tensor) –
float8_e4m3fn[28672, 5376]andfloat32[28672]fromprepare_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 receivesa_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