pub fn psa_extract_v_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
qkv_ptr: T::Pointer<D>,
v_ptr: T::Pointer<D>,
qkv_h: i32,
c: i32,
H: i32,
W: i32,
num_heads: i32,
)Expand description
Extracts V channels from QKV NCHW into V NCHW.
Input: qkv_ptr — [B, qkv_h, H, W] NCHW
Output: v_ptr — [B, c, H, W] NCHW (c = num_heads * 2 * KEY_DIM)
Grid: [BH * N, 1, 1]. Block: [KEY_DIM, 1, 1].