fp8_gemm_n128

Function fp8_gemm_n128 

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

Pre-dequanted FP8 GEMM (prefill): C = A @ B_fp8.

A: [M, K] BF16, B_fp8: [N, K] FP8 E4M3 (pre-dequanted from NVFP4), C: [M, N] BF16. Eliminates runtime NVFP4→FP8 dequant — only LOAD + FP8 MMA per K step.

Grid: (ceil(N/128), ceil(M/64), 1) Block: (128, 1, 1)