pub fn psa_extract_v_backward<T: Triton, D: Float, const KEY_DIM: i32>(
d_v_ptr: T::Pointer<D>,
d_qkv_ptr: T::Pointer<D>,
qkv_h: i32,
c: i32,
H: i32,
W: i32,
num_heads: i32,
)Expand description
Backward of psa_extract_v_nchw: scatters d_v back to V_lo / V_hi
channels of d_qkv.
Grid / block matches the forward: [BH * N, 1, 1], block [KEY_DIM, 1, 1].
Uses atomic_add because V_lo / V_hi channels of d_qkv are also updated
by psa_pack_qkv_backward.