flashinfer.diffusion_ops.minimax_h3_bf16_pre_attention¶
- flashinfer.diffusion_ops.minimax_h3_bf16_pre_attention(x: Tensor, x_norm_weight: Tensor, adaln_scale: Tensor, adaln_shift: Tensor, adaln_index: Tensor, qkv_weight: Tensor, q_norm_weight: Tensor, k_norm_weight: Tensor, rope_cos_sin: Tensor, *, ulysses_degree: int, out: Tensor, eps: float = 1e-05) Tensor¶
Run the fused BF16 pre-attention projection for MiniMax-H3.
The operation applies input RMSNorm, indexed AdaLN, a BF16 QKV projection, per-head Q/K RMSNorm, partial 3-D split-half NeoX RoPE, and a destination-major output pack. The collective that consumes
outis outside this operation.- Parameters:
x (torch.Tensor) – Contiguous BF16 input with shape
[M, 5376].x_norm_weight (torch.Tensor) – Contiguous BF16 input RMSNorm weight with shape
[5376].adaln_scale (torch.Tensor) – Contiguous BF16 AdaLN tables with shape
[9, 5376].adaln_shift (torch.Tensor) – Contiguous BF16 AdaLN tables with shape
[9, 5376].adaln_index (torch.Tensor) – Contiguous int32 row indices with shape
[M]and values in[0, 8]. A malformed index is guarded in the CUDA kernel and makes its corresponding output row all-zero instead of addressing outside the AdaLN tables.qkv_weight (torch.Tensor) – Contiguous BF16 checkpoint weight with physical shape
[21504, 5376]and row order[head, qkv_kind, head_dim].q_norm_weight (torch.Tensor) – Contiguous BF16 per-head RMSNorm weights with shape
[128].k_norm_weight (torch.Tensor) – Contiguous BF16 per-head RMSNorm weights with shape
[128].rope_cos_sin (torch.Tensor) – Contiguous BF16 cache with shape
[M, 96]. Columns[0, 48)hold frame/height/width cosine values and columns[48, 96)hold the corresponding sine values. RoPE transforms Q/K dimensions[0, 96); dimensions[96, 128)pass through.ulysses_degree (int) – Destination count, one of
1,2,4, or8.out (torch.Tensor) – Caller-owned contiguous BF16 destination with shape
[P, M, 56 // P, 3, 128].eps (float) – RMSNorm epsilon. This kernel supports
1e-5.
- Returns:
The same tensor passed as
out.- Return type:
torch.Tensor
Notes
The CUDA kernel independently guards its AdaLN table loads, so malformed indices cannot form an out-of-bounds address. Valid indices preserve the MiniMax-H3 checkpoint semantics without a synchronizing host reduction.
This is a direct kernel entry point for all supported destination counts. On SM103a, the measured performance promotion range is
P in {2, 4, 8}. Callers that dispatch byPshould retain their segmented fallback forP=1.