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.