fp8_gemm_n128_row_scaled_m16

Function fp8_gemm_n128_row_scaled_m16 

Source
pub fn fp8_gemm_n128_row_scaled_m16(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    input: DevicePtr,
    weight: &Fp8DenseWeight,
    output: DevicePtr,
    m: u32,
    n: u32,
    k: u32,
    stream: u64,
) -> Result<()>
Expand description

Small-M row-scaled FP8 GEMM (M ≤ 16) — single warp per CTA variant.

Same math as fp8_gemm_n128_row_scaled but M_TILE=16 instead of 64, so all M rows are valid (no wasted MMA cycles on bounds-checked rows). Uses 32 threads per CTA (1 warp) instead of 128, so 4× fewer threads for the same useful work. Critical for the DFlash drafter lm_head where M=γ=16 vs N=vocab_size=248320.

Kernel: fp8_gemm_t_row_scaled_m16(A, B_fp8, row_scale, C, M, N, K). Grid: (ceil(N/128), 1, 1) Block: (32, 1, 1)