fp8_gemm_n128_row_scaled

Function fp8_gemm_n128_row_scaled 

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

Row-scaled FP8 GEMM: C[M, N] = A[M, K] @ (dequant(B_fp8[N, K]) * row_scale[N]).

Same tiling and FP8 MMA as fp8_gemm_n128 (BF16 × FP8 → BF16), with a per-column scale multiply before the BF16 write-out. Consumes the Fp8DenseWeight produced by crate::weight_map::DenseWeight::quantize_to_fp8 — the per-row scale on Fp8DenseWeight matches the kernel’s row_scale parameter.

Phase G (DFlash drafter FP8) hot-path GEMM. Replaces dense_gemm on the seven dense-GEMM call sites in forward_block_layer_pre_attn / _post_attn when self.quant == DflashQuantization::Fp8Weights.

Kernel: fp8_gemm_t_row_scaled(A, B_fp8, row_scale, C, M, N, K)kernels/gb10/qwen3.6-27b/nvfp4/w4a16_gemm.cu. Grid: (ceil(N/128), ceil(M/64), 1) Block: (128, 1, 1)