mla_q_rope_scatter

Function mla_q_rope_scatter 

Source
pub fn mla_q_rope_scatter(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    q_full: DevicePtr,
    q_absorbed_buf: DevicePtr,
    q_rope_contiguous: DevicePtr,
    nq: u32,
    hd: u32,
    nope: u32,
    rope: u32,
    kv_lora: u32,
    mla_cache_dim: u32,
    stream: u64,
) -> Result<()>
Expand description

MLA Q_rope scatter: copy rope portion from q_full to strided q_absorbed_buf. 1 kernel replaces 32 D2D copies.