moe_bf16_grouped_gemm

Function moe_bf16_grouped_gemm 

Source
pub fn moe_bf16_grouped_gemm(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    input: DevicePtr,
    weight_ptrs: DevicePtr,
    output: DevicePtr,
    expert_offsets: DevicePtr,
    sorted_token_ids: DevicePtr,
    num_experts: u32,
    n: u32,
    k: u32,
    max_m_tiles: u32,
    stream: u64,
) -> Result<()>
Expand description

BF16 grouped GEMM for sorted MoE prefill (FP8-dequant-on-load path).

BF16 activations × BF16 expert weights via pointer table. No scale. Used when expert weights have been dequanted from FP8 to BF16 at load time (ATLAS_FP8_DEQUANT_MOE_TO_BF16=1). Eliminates the per-layer 0.989 cosine ceiling that comes from FP8 quantization itself.

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