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 out is 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, or 8.

  • 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 by P should retain their segmented fallback for P=1.