pub fn psa_pack_qkv_backward<T: Triton, D: Float, const KEY_DIM: i32>(
d_packed_ptr: T::Pointer<D>,
d_qkv_ptr: T::Pointer<D>,
qkv_h: i32,
H: i32,
W: i32,
B: i32,
num_heads: i32,
)Expand description
Backward of psa_pack_qkv: scatters d_packed back to d_qkv.
Grid / block matches the forward pass: [4 * BH * N, 1, 1], block [KEY_DIM, 1, 1].
Uses atomic_add because both psa_pack_qkv_backward (all four sections)
and psa_extract_v_backward (V_lo/V_hi sections) write to the same
d_qkv buffer.