gdn_decode_f16_strided_norm

Function gdn_decode_f16_strided_norm 

Source
pub fn gdn_decode_f16_strided_norm(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    h_state: DevicePtr,
    query: DevicePtr,
    key: DevicePtr,
    value: DevicePtr,
    gate: DevicePtr,
    beta: DevicePtr,
    z_gate: DevicePtr,
    norm_weight: 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,
    z_stride: u32,
    out_stride: u32,
    h_seq_stride: u64,
    eps: f32,
    stream: u64,
) -> Result<()>
Expand description

FP16 h-state twin of gdn_decode_f32_strided_norm (ATLAS_SSM_H_FP16).

The only signature difference is h_seq_stride: the per-sequence stride of the h-state pool in __half elements. Stage 1 keeps the pool FP32-sized, so slots are h_state_bytes apart while the dense FP16 footprint is half that — the stride must be passed, not inferred from the head dims.