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.