compute_gdn_gates

Function compute_gdn_gates 

Source
pub fn compute_gdn_gates(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    ba_interleaved: DevicePtr,
    a_log: DevicePtr,
    dt_bias: DevicePtr,
    gate_out: DevicePtr,
    beta_out: DevicePtr,
    num_tokens: u32,
    num_v_heads: u32,
    num_groups: u32,
    vheads_per_group: u32,
    ba_stride: u32,
    stream: u64,
) -> Result<()>
Expand description

Compute GDN gates from interleaved BA projection + learned A_log/dt_bias.

Outputs FP32 gate (decay) and beta (write gate) for each value head.

Kernel: compute_gdn_gates(ba_interleaved, A_log, dt_bias, gate_out, beta_out, num_v_heads, num_groups, vheads_per_group) Grid: (1, 1, 1) Block: (num_v_heads, 1, 1)