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.