flashinfer_gdn_prefill

Function flashinfer_gdn_prefill 

Source
pub fn flashinfer_gdn_prefill(
    gpu: &dyn GpuBackend,
    qkv: DevicePtr,
    gate_beta: DevicePtr,
    output: DevicePtr,
    h_state: DevicePtr,
    scale: f32,
    total: u32,
    nk: u32,
    nv: u32,
    kd: u32,
    vd: u32,
    conv_dim: u32,
    gb_stride: u32,
    num_seqs: u32,
    stream: u64,
) -> Result<()>
Expand description

Run one prefill GDN scan through the FlashInfer kernel on Atlas’s native buffers.

qkv: packed [Q(key_dim)|K(key_dim)|V(value_dim)] bf16, row stride conv_dim. gate_beta: interleaved [gate(nv)|beta(nv)] fp32, row stride gb_stride. output: contiguous [total, value_dim] bf16. h_state: [nv,kd,vd] fp32 (final state out). Single-stream only (num_seqs == 1); fresh prefill (zero init state).