batched_embed_fp8

Function batched_embed_fp8 

Source
pub fn batched_embed_fp8(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    token_ids_dev: DevicePtr,
    embed_table: DevicePtr,
    row_scale: DevicePtr,
    output: DevicePtr,
    num_tokens: u32,
    hidden_size: u32,
    stream: u64,
) -> Result<()>
Expand description

FP8-table variant of batched_embed: rows are FP8 E4M3 bytes with a per-row f32 dequant scale (the quantize_bf16_to_fp8 layout); the kernel dequantizes on read and writes BF16 rows.

Kernel: batched_embed_fp8(token_ids, table, row_scale, output, hidden) Grid: (num_tokens, 1, 1) Block: (256, 1, 1)