gdn_decode_wy4

Function gdn_decode_wy4 

Source
pub fn gdn_decode_wy4(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    h_state: DevicePtr,
    query: DevicePtr,
    key: DevicePtr,
    value: DevicePtr,
    gate: DevicePtr,
    beta: DevicePtr,
    output: DevicePtr,
    h_state_inter0: DevicePtr,
    h_state_inter1: DevicePtr,
    h_state_inter2: 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,
    state_is_table: bool,
    stream: u64,
) -> Result<()>
Expand description

WY-chunkwise 4-token GDN decode (2-pass algorithm).

All 4 H^T @ k_t dot products computed in a single pass, then WY correction derives v_new values. Second pass applies all 4 state updates + outputs. 2 passes vs 5, reducing memory traffic by 60%.

Grid: (num_v_heads, batch, 1) Block: (128, 1, 1)