pub fn gdn_decode_f32_strided_norm_snap(
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,
h_inter: DevicePtr,
h_inter_seq_stride: u64,
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,
eps: f32,
stream: u64,
) -> Result<()>Expand description
super::gdn_decode_f32_strided_norm + inline h-state snapshot, for the
batched-verify arm at batch_size = n sequences.
h_inter is the snapshot base for THIS token position; sequences are
h_inter_seq_stride FP32 elements apart (the ssm-pool per-slot
intermediate stride — passed, not inferred, because pool slots are
num_intermediates snapshots wide while H itself is dense). NULL skips.