deinterleave_qkvz

Function deinterleave_qkvz 

Source
pub fn deinterleave_qkvz(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    interleaved: DevicePtr,
    output: DevicePtr,
    num_tokens: u32,
    num_groups: u32,
    head_k_dim: u32,
    vheads_per_group: u32,
    head_v_dim: u32,
    stream: u64,
) -> Result<()>
Expand description

Deinterleave QKVZ projection output from GQA-grouped to sequential layout.

Input: interleaved [num_groups × (kd + kd + vpgvd + vpgvd)] Output: sequential [Q_total | K_total | V_total | Z_total]

Kernel: deinterleave_qkvz(interleaved, output, num_groups, head_k_dim, vheads_per_group, head_v_dim) Grid: (ceil(total/256), 1, 1) Block: (256, 1, 1)