flashinfer.diffusion_ops.minimax_h3_nvfp4_out_proj¶
- flashinfer.diffusion_ops.minimax_h3_nvfp4_out_proj(attn_out: Tensor, o_weight_q: Tensor, o_weight_sf: Tensor, o_weight_global_scale: float | Tensor, act_global_scale: Tensor, gate: Tensor, gate_index: Tensor, residual: Tensor, *, out: Tensor | None = None, act_q: Tensor | None = None, act_sf: Tensor | None = None) Tensor¶
MiniMax-H3 NVFP4 (W4A4) attention output projection with the fused gated residual for SM120 (GB202: RTX 5090 / RTX PRO 6000 Blackwell).
One launch: FlashInfer-convention block-16 NVFP4 quantization of
attn_out(sf = UE4M3(block_amax * act_global_scale / 6),code = RN(x * act_global_scale / sf)) run by the spare warps of the persistent kernel’s producer warpgroup and overlapped, through per-tile ready flags, with thekind::mxf4nvf4block-scaledmma.syncGEMM with the fusedout = BF16(residual + BF16(gate[gate_index[m]] * o))epilogue. Rows whosegate_indexlies outside[0, 9)contributegate = 0(out = residual).- Parameters:
attn_out (torch.Tensor) – BF16
[M, 7168]packed attention output.o_weight_q (torch.Tensor) – uint8
[5376, 3584]E2M1x2 weight (seequantize_minimax_h3_o_weight_nvfp4()).o_weight_sf (torch.Tensor) – uint8 UE4M3 block-16 weight scales in the FlashInfer 128x4 swizzled layout.
o_weight_global_scale (float or torch.Tensor) – FP32 per-tensor weight global scale
448 * 6 / amax(W).act_global_scale (torch.Tensor) – FP32
[1]CUDA tensor, the calibrated activation global scale448 * 6 / amax.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 E2M1x2 activation workspace (uint8
[M, 3584]).act_sf (Optional[torch.Tensor]) – Optional caller-owned activation block-scale workspace (uint8
[M, 448]).
- Returns:
The BF16
[M, 5376]post-attention hidden state.- Return type:
torch.Tensor