gdn_decode_f32_strided

Function gdn_decode_f32_strided 

Source
pub fn gdn_decode_f32_strided(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    h_state: DevicePtr,
    query: DevicePtr,
    key: DevicePtr,
    value: DevicePtr,
    gate: DevicePtr,
    beta: DevicePtr,
    output: DevicePtr,
    batch_size: u32,
    num_k_heads: u32,
    num_v_heads: u32,
    k_dim: u32,
    v_dim: u32,
    qk_stride: u32,
    v_stride: u32,
    gb_stride: u32,
    out_stride: u32,
    stream: u64,
) -> Result<()>
Expand description

Strided FP32 GDN decode for concurrent sequence decode.

Q/K/V are read from strided rows, typically the FP32 conv output laid out as [batch, Q | K | V]. Gate/beta and output are also strided by batch row.