pub fn psa_pack_qkv<T: Triton, D: Float, const KEY_DIM: i32>(
qkv_ptr: T::Pointer<D>,
out_ptr: T::Pointer<D>,
qkv_h: i32,
H: i32,
W: i32,
B: i32,
num_heads: i32,
)Expand description
Rearranges the NCHW QKV tensor into packed FA2 format.
Input: qkv_ptr — [B, qkv_h, H, W] NCHW, qkv_h = num_heads * 4 * KEY_DIM
Output: out_ptr — flat [4, BH, N, KEY_DIM] buffer
- Section 0: Q, Section 1: K, Section 2: V_lo, Section 3: V_hi
Grid: [4 * BH * N, 1, 1] — one CTA per (section, bh, n) triple.
Block: [KEY_DIM, 1, 1]