gdn_decode

Function gdn_decode 

Source
pub fn gdn_decode(
    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,
    stream: u64,
) -> Result<()>
Expand description

Gated delta rule decode (recurrent SSM update, supports batched sequences).

Kernel: gated_delta_rule_decode(h_state, query, key, value, gate, beta, output, batch_size, num_k_heads, num_v_heads, k_dim, v_dim) Grid: (num_v_heads, batch_size, 1) Block: (128, 1, 1)

For batch_size > 1, h_state layout: [batch, num_v_heads, k_dim, v_dim].