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 themma.syncGEMM[M, 7168] x [7168, 5376]whose epilogue applies the per-token x per-channel dequant scales, rounds the projectionoto BF16 once and writesout = BF16(residual + BF16(gate[gate_index[m]] * o)). This is the SGLangindexed_gate_bf16round-point convention. Rows whosegate_indexlies outside[0, 9)contributegate = 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 (seequantize_minimax_h3_o_weight_fp8()).o_weight_scale (torch.Tensor) – FP32
[5376]dequant multipliers.gate (torch.Tensor) – BF16
[9, 5376]per-indexgate_msarows 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 whenNone).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