mla_batched_gemv

Function mla_batched_gemv 

Source
pub fn mla_batched_gemv(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    input: DevicePtr,
    weight: DevicePtr,
    output: DevicePtr,
    n_out: u32,
    k: u32,
    num_heads: u32,
    input_stride: u32,
    output_stride: u32,
    stream: u64,
) -> Result<()>
Expand description

Paged decode attention (FP8 KV cache, single/multi sequence).

Kernel: paged_decode_attn_fp8(Q, K_cache, V_cache, O, block_tables, seq_lens, max_blocks_per_seq, num_q_heads, num_kv_heads, head_dim, block_size, inv_sqrt_d, k_scale, v_scale, q_stride, cache_stride) Grid: (num_q_heads, num_seqs, 1) Block: (256, 1, 1)

block_tables: device ptr to i32[num_seqs * max_blocks_per_seq] seq_lens: device ptr to i32[num_seqs] cache_stride is in elements (u64). BF16 paged decode attention — no FP8 quantization, direct BF16 KV cache. MLA batched GEMV: output[head, n] = sum_k(weight[head, n, k] * input[head, k]) Replaces 32 sequential dense_gemv calls with a single kernel launch.