flashinfer.diffusion_ops.minimax_h3_fp8_out_proj

flashinfer.diffusion_ops.minimax_h3_fp8_out_proj(attn_out: Tensor, o_weight_q: Tensor, o_weight_scale: Tensor, gate: Tensor, gate_index: Tensor, residual: Tensor, *, out: Tensor | None = None, act_q: Tensor | None = None, act_scale: Tensor | None = None) → Tensor

MiniMax-H3 FP8 (W8A8) attention output projection with the fused gated residual for SM120 (GB202: RTX 5090 / RTX PRO 6000 Blackwell).

One launch: a per-token E4M3 quantization of attn_out (scale = RN(amax / 448)) run by the spare warps of the persistent kernel’s producer warpgroup and overlapped, through per-tile ready flags, with the mma.sync GEMM [M, 7168] x [7168, 5376] whose epilogue applies the per-token x per-channel dequant scales, rounds the projection o to BF16 once and writes out = BF16(residual + BF16(gate[gate_index[m]] * o)). This is the SGLang indexed_gate_bf16 round-point convention. Rows whose gate_index lies outside [0, 9) contribute gate = 0 (out = residual).

Parameters:
  • attn_out (torch.Tensor) – BF16 [M, 7168] packed attention output (the [M, 56, 128] NHD output viewed row-major).

  • o_weight_q (torch.Tensor) – E4M3 [5376, 7168] per-output-channel quantized weight (see quantize_minimax_h3_o_weight_fp8()).

  • o_weight_scale (torch.Tensor) – FP32 [5376] dequant multipliers.

  • gate (torch.Tensor) – BF16 [9, 5376] per-index gate_msa rows of the AdaLN plan.

  • gate_index (torch.Tensor) – int32 [M] per-row table index.

  • residual (torch.Tensor) – BF16 [M, 5376] residual stream.

  • out (Optional[torch.Tensor]) – BF16 [M, 5376] output (allocated when None).

  • act_q (Optional[torch.Tensor]) – Optional caller-owned E4M3 activation workspace ([M, 7168]).

  • act_scale (Optional[torch.Tensor]) – Optional caller-owned FP32 per-row activation scale workspace ([M]).

Returns:

The BF16 [M, 5376] post-attention hidden state.

Return type:

torch.Tensor