1use ferrox_core::matmul::{rms_norm, silu};
67use ferrox_core::weight_matrix::WeightMatrix;
68
69#[derive(Debug, Clone, Copy)]
79pub struct GdnConfig {
80 pub hidden_dim: usize,
81 pub num_key_heads: usize,
83 pub num_value_heads: usize,
86 pub key_head_dim: usize,
88 pub value_head_dim: usize,
90 pub conv_kernel_size: usize,
91 pub rms_norm_eps: f32,
92}
93
94impl GdnConfig {
95 pub fn key_dim(&self) -> usize {
98 self.num_key_heads * self.key_head_dim
99 }
100
101 pub fn value_dim(&self) -> usize {
104 self.num_value_heads * self.value_head_dim
105 }
106
107 pub fn qkv_dim(&self) -> usize {
115 2 * self.key_dim() + self.value_dim()
116 }
117
118 pub fn heads_per_key_group(&self) -> usize {
122 self.num_value_heads / self.num_key_heads
123 }
124}
125
126pub struct GdnWeights {
128 pub attn_qkv: WeightMatrix, pub attn_gate: WeightMatrix, pub ssm_conv1d: Vec<f32>,
132 pub ssm_dt: Vec<f32>, pub ssm_a: Vec<f32>, pub ssm_beta: WeightMatrix, pub ssm_alpha: WeightMatrix, pub ssm_norm: Vec<f32>, pub ssm_out: WeightMatrix, }
139
140pub struct GdnState {
142 conv_hist: Vec<f32>,
143 recurrent: Vec<f32>,
150}
151
152impl GdnState {
153 pub fn new(cfg: &GdnConfig) -> Self {
154 Self {
155 conv_hist: Vec::new(),
156 recurrent: vec![0f32; cfg.num_value_heads * cfg.value_head_dim * cfg.key_head_dim],
157 }
158 }
159}
160
161fn softplus(x: f32) -> f32 {
162 if x > 20.0 {
163 x
164 } else {
165 (1.0 + x.exp()).ln()
166 }
167}
168
169fn sigmoid(x: f32) -> f32 {
170 1.0 / (1.0 + (-x).exp())
171}
172
173fn l2_normalize(v: &mut [f32], eps: f32) {
174 let norm_sq: f32 = v.iter().map(|x| x * x).sum();
175 let scale = 1.0 / (norm_sq + eps).sqrt();
176 for x in v.iter_mut() {
177 *x *= scale;
178 }
179}
180
181fn causal_conv_step(
183 weight: &[f32],
184 history: &mut Vec<f32>,
185 current: &[f32],
186 kernel_size: usize,
187 dim: usize,
188) -> Vec<f32> {
189 let hist_len = history.len() / dim.max(1);
190 let missing = (kernel_size - 1).saturating_sub(hist_len);
191
192 let mut y = vec![0f32; dim];
193 for j in 0..kernel_size {
194 if j < missing {
195 continue;
196 }
197 let src: &[f32] = if j == kernel_size - 1 {
198 current
199 } else {
200 let hist_idx = j - missing;
201 &history[hist_idx * dim..(hist_idx + 1) * dim]
202 };
203 for d in 0..dim {
204 y[d] += weight[d * kernel_size + j] * src[d];
205 }
206 }
207 for v in y.iter_mut() {
208 *v = silu(*v);
209 }
210
211 history.extend_from_slice(current);
212 let max_hist_len = (kernel_size - 1) * dim;
213 if history.len() > max_hist_len {
214 let excess = history.len() - max_hist_len;
215 history.drain(0..excess);
216 }
217 y
218}
219
220pub fn gdn_forward_token(
230 weights: &GdnWeights,
231 cfg: &GdnConfig,
232 hidden: &[f32],
233 state: &mut GdnState,
234) -> Vec<f32> {
235 assert_eq!(hidden.len(), cfg.hidden_dim);
236 assert!(
237 cfg.num_key_heads > 0 && cfg.num_value_heads.is_multiple_of(cfg.num_key_heads),
238 "GDN needs num_value_heads ({}) to be a positive multiple of num_key_heads ({}); \
239 otherwise repeat_interleave has no whole replication factor and some V heads would \
240 silently read a K head that never fed them",
241 cfg.num_value_heads,
242 cfg.num_key_heads
243 );
244
245 let qkv_dim = cfg.qkv_dim();
246 let key_dim = cfg.key_dim();
247 let value_dim = cfg.value_dim();
248 let key_head_dim = cfg.key_head_dim;
249 let value_head_dim = cfg.value_head_dim;
250 let rep = cfg.heads_per_key_group();
251
252 let qkv_lin = weights.attn_qkv.apply(hidden);
253 let z = weights.attn_gate.apply(hidden);
254 let beta_raw = weights.ssm_beta.apply(hidden);
255 let alpha_raw = weights.ssm_alpha.apply(hidden);
256
257 let qkv = causal_conv_step(
258 &weights.ssm_conv1d,
259 &mut state.conv_hist,
260 &qkv_lin,
261 cfg.conv_kernel_size,
262 qkv_dim,
263 );
264
265 let (q_all, rest) = qkv.split_at(key_dim);
267 let (k_all, v_all) = rest.split_at(key_dim);
268
269 let scale = 1.0 / (key_head_dim as f32).sqrt();
270 let mut y_flat = vec![0f32; value_dim];
271
272 #[allow(clippy::needless_range_loop)]
273 for h in 0..cfg.num_value_heads {
274 let k_base = (h / rep) * key_head_dim;
276 let v_base = h * value_head_dim;
277 let mut q_h = q_all[k_base..k_base + key_head_dim].to_vec();
278 let mut k_h = k_all[k_base..k_base + key_head_dim].to_vec();
279 let v_h = &v_all[v_base..v_base + value_head_dim];
280
281 l2_normalize(&mut q_h, 1e-6);
282 l2_normalize(&mut k_h, 1e-6);
283 for x in q_h.iter_mut() {
284 *x *= scale;
285 }
286
287 let gate = softplus(alpha_raw[h] + weights.ssm_dt[h]) * weights.ssm_a[h];
289 let decay = gate.exp();
290 let beta = sigmoid(beta_raw[h]);
291
292 let block = value_head_dim * key_head_dim;
293 let s_base = h * block;
294 let s = &mut state.recurrent[s_base..s_base + block];
295
296 for cell in s.iter_mut() {
298 *cell *= decay;
299 }
300
301 let mut kv_mem = vec![0f32; value_head_dim];
303 for v_idx in 0..value_head_dim {
304 let mut acc = 0f32;
305 for k_idx in 0..key_head_dim {
306 acc += s[v_idx * key_head_dim + k_idx] * k_h[k_idx];
307 }
308 kv_mem[v_idx] = acc;
309 }
310
311 for v_idx in 0..value_head_dim {
313 let delta = (v_h[v_idx] - kv_mem[v_idx]) * beta;
314 for k_idx in 0..key_head_dim {
315 s[v_idx * key_head_dim + k_idx] += delta * k_h[k_idx];
316 }
317 }
318
319 for v_idx in 0..value_head_dim {
321 let mut acc = 0f32;
322 for k_idx in 0..key_head_dim {
323 acc += s[v_idx * key_head_dim + k_idx] * q_h[k_idx];
324 }
325 y_flat[v_base + v_idx] = acc;
326 }
327 }
328
329 let mut gated = vec![0f32; value_dim];
331 for h in 0..cfg.num_value_heads {
332 let base = h * value_head_dim;
333 let normed = rms_norm(
334 &y_flat[base..base + value_head_dim],
335 &weights.ssm_norm,
336 cfg.rms_norm_eps,
337 );
338 for i in 0..value_head_dim {
339 gated[base + i] = silu(z[base + i]) * normed[i];
340 }
341 }
342
343 weights.ssm_out.apply(&gated)
344}
345
346#[cfg(test)]
347mod tests {
348 use super::*;
349 use ferrox_core::tensor::Tensor;
350
351 const HIDDEN: usize = 4;
352 const N_HEADS: usize = 2;
353 const HEAD_DIM: usize = 2;
354 const CONV_K: usize = 2;
355 const QKV_DIM: usize = 3 * N_HEADS * HEAD_DIM; const V_DIM: usize = N_HEADS * HEAD_DIM; fn wm(data: &[f32], rows: usize, cols: usize) -> WeightMatrix {
359 assert_eq!(data.len(), rows * cols);
360 WeightMatrix::F32(Tensor::new(data.to_vec(), vec![rows, cols]))
361 }
362
363 fn cfg() -> GdnConfig {
364 GdnConfig {
365 hidden_dim: HIDDEN,
366 num_key_heads: N_HEADS,
367 num_value_heads: N_HEADS,
368 key_head_dim: HEAD_DIM,
369 value_head_dim: HEAD_DIM,
370 conv_kernel_size: CONV_K,
371 rms_norm_eps: 1e-5,
372 }
373 }
374
375 fn make_weights() -> GdnWeights {
376 let mut qkv = Vec::with_capacity(QKV_DIM * HIDDEN);
378 for i in 0..QKV_DIM * HIDDEN {
379 qkv.push(((i % 7) as f32 - 3.0) * 0.1);
380 }
381 let mut gate = Vec::with_capacity(V_DIM * HIDDEN);
382 for i in 0..V_DIM * HIDDEN {
383 gate.push(((i % 5) as f32 - 2.0) * 0.08);
384 }
385 let mut conv = Vec::with_capacity(QKV_DIM * CONV_K);
386 for i in 0..QKV_DIM * CONV_K {
387 conv.push(if i % CONV_K == CONV_K - 1 { 1.0 } else { 0.1 });
388 }
389 let mut beta = Vec::with_capacity(N_HEADS * HIDDEN);
390 let mut alpha = Vec::with_capacity(N_HEADS * HIDDEN);
391 for i in 0..N_HEADS * HIDDEN {
392 beta.push(((i % 3) as f32 - 1.0) * 0.2);
393 alpha.push(((i % 4) as f32 - 1.5) * 0.15);
394 }
395 let mut out = Vec::with_capacity(HIDDEN * V_DIM);
396 for i in 0..HIDDEN * V_DIM {
397 out.push(((i % 6) as f32 - 2.5) * 0.12);
398 }
399 GdnWeights {
400 attn_qkv: wm(&qkv, QKV_DIM, HIDDEN),
401 attn_gate: wm(&gate, V_DIM, HIDDEN),
402 ssm_conv1d: conv,
403 ssm_dt: vec![0.1, -0.05],
404 ssm_a: vec![-0.5, -0.75],
406 ssm_beta: wm(&beta, N_HEADS, HIDDEN),
407 ssm_alpha: wm(&alpha, N_HEADS, HIDDEN),
408 ssm_norm: vec![1.0, 1.0],
409 ssm_out: wm(&out, HIDDEN, V_DIM),
410 }
411 }
412
413 struct RawGdn {
417 qkv: Vec<f32>, gate: Vec<f32>, conv: Vec<f32>, dt: Vec<f32>, a: Vec<f32>, beta: Vec<f32>, alpha: Vec<f32>, norm: Vec<f32>, out: Vec<f32>, }
427
428 impl RawGdn {
429 fn to_weights(&self, cfg: &GdnConfig) -> GdnWeights {
430 let h = cfg.hidden_dim;
431 GdnWeights {
432 attn_qkv: wm(&self.qkv, cfg.qkv_dim(), h),
433 attn_gate: wm(&self.gate, cfg.value_dim(), h),
434 ssm_conv1d: self.conv.clone(),
435 ssm_dt: self.dt.clone(),
436 ssm_a: self.a.clone(),
437 ssm_beta: wm(&self.beta, cfg.num_value_heads, h),
438 ssm_alpha: wm(&self.alpha, cfg.num_value_heads, h),
439 ssm_norm: self.norm.clone(),
440 ssm_out: wm(&self.out, h, cfg.value_dim()),
441 }
442 }
443 }
444
445 fn fill(n: usize, seed: usize) -> Vec<f32> {
449 (0..n)
450 .map(|i| {
451 let k = (i * 37 + seed * 101) % 23;
452 (k as f32 - 11.0)
453 * 0.043
454 * if (i + seed).is_multiple_of(2) {
455 1.0
456 } else {
457 -1.0
458 }
459 })
460 .collect()
461 }
462
463 fn matvec(rows_data: &[f32], rows: usize, cols: usize, x: &[f32]) -> Vec<f32> {
464 assert_eq!(rows_data.len(), rows * cols);
465 assert_eq!(x.len(), cols);
466 (0..rows)
467 .map(|r| (0..cols).map(|c| rows_data[r * cols + c] * x[c]).sum())
468 .collect()
469 }
470
471 fn ref_sigmoid(x: f32) -> f32 {
472 1.0 / (1.0 + (-x).exp())
473 }
474
475 fn ref_softplus(x: f32) -> f32 {
476 (1.0 + x.exp()).ln()
477 }
478
479 fn ref_silu(x: f32) -> f32 {
480 x * ref_sigmoid(x)
481 }
482
483 fn ref_l2norm(v: &[f32]) -> Vec<f32> {
484 let sum_sq: f32 = v.iter().map(|x| x * x).sum();
485 let inv = 1.0 / (sum_sq + 1e-6).sqrt();
486 v.iter().map(|x| x * inv).collect()
487 }
488
489 fn reference_forward(raw: &RawGdn, cfg: &GdnConfig, tokens: &[Vec<f32>]) -> Vec<Vec<f32>> {
503 let hidden = cfg.hidden_dim;
504 let qkv_dim = cfg.qkv_dim();
505 let key_dim = cfg.key_dim();
506 let value_dim = cfg.value_dim();
507 let dk = cfg.key_head_dim;
508 let dv = cfg.value_head_dim;
509 let kernel = cfg.conv_kernel_size;
510 let rep = cfg.num_value_heads / cfg.num_key_heads;
511 let scale = 1.0 / (dk as f32).sqrt();
512
513 let mixed: Vec<Vec<f32>> = tokens
516 .iter()
517 .map(|h| matvec(&raw.qkv, qkv_dim, hidden, h))
518 .collect();
519 let mut conved = vec![vec![0f32; qkv_dim]; tokens.len()];
520 for (t, conved_t) in conved.iter_mut().enumerate() {
521 for (d, out_d) in conved_t.iter_mut().enumerate() {
522 let mut acc = 0f32;
523 for j in 0..kernel {
524 let src = t as isize - (kernel as isize - 1) + j as isize;
525 if src < 0 {
526 continue;
527 }
528 acc += raw.conv[d * kernel + j] * mixed[src as usize][d];
529 }
530 *out_d = ref_silu(acc);
531 }
532 }
533
534 let mut state = vec![vec![vec![0f32; dv]; dk]; cfg.num_value_heads];
536 let mut outputs = Vec::with_capacity(tokens.len());
537
538 for (t, token) in tokens.iter().enumerate() {
539 let z = matvec(&raw.gate, value_dim, hidden, token);
540 let a_raw = matvec(&raw.alpha, cfg.num_value_heads, hidden, token);
541 let b_raw = matvec(&raw.beta, cfg.num_value_heads, hidden, token);
542
543 let q_slice = &conved[t][0..key_dim];
544 let k_slice = &conved[t][key_dim..2 * key_dim];
545 let v_slice = &conved[t][2 * key_dim..];
546
547 let mut q_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
549 let mut k_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
550 for kh in 0..cfg.num_key_heads {
551 for _ in 0..rep {
552 q_heads.push(q_slice[kh * dk..(kh + 1) * dk].to_vec());
553 k_heads.push(k_slice[kh * dk..(kh + 1) * dk].to_vec());
554 }
555 }
556
557 let mut core = vec![0f32; value_dim];
558 for h in 0..cfg.num_value_heads {
559 let q = ref_l2norm(&q_heads[h]);
560 let k = ref_l2norm(&k_heads[h]);
561 let v = &v_slice[h * dv..(h + 1) * dv];
562
563 let decay = (ref_softplus(a_raw[h] + raw.dt[h]) * raw.a[h]).exp();
564 let beta = ref_sigmoid(b_raw[h]);
565
566 for row in state[h].iter_mut() {
567 for cell in row.iter_mut() {
568 *cell *= decay;
569 }
570 }
571 let mut kv_mem = vec![0f32; dv];
573 for (k_idx, row) in state[h].iter().enumerate() {
574 for (v_idx, cell) in row.iter().enumerate() {
575 kv_mem[v_idx] += cell * k[k_idx];
576 }
577 }
578 let delta: Vec<f32> = (0..dv).map(|i| (v[i] - kv_mem[i]) * beta).collect();
580 for (k_idx, row) in state[h].iter_mut().enumerate() {
581 for (v_idx, cell) in row.iter_mut().enumerate() {
582 *cell += k[k_idx] * delta[v_idx];
583 }
584 }
585 for (k_idx, row) in state[h].iter().enumerate() {
587 for (v_idx, cell) in row.iter().enumerate() {
588 core[h * dv + v_idx] += cell * q[k_idx] * scale;
589 }
590 }
591 }
592
593 let mut gated = vec![0f32; value_dim];
595 for h in 0..cfg.num_value_heads {
596 let base = h * dv;
597 let mean_sq = core[base..base + dv].iter().map(|x| x * x).sum::<f32>() / dv as f32;
598 let inv = 1.0 / (mean_sq + cfg.rms_norm_eps).sqrt();
599 for i in 0..dv {
600 gated[base + i] = core[base + i] * inv * raw.norm[i] * ref_silu(z[base + i]);
601 }
602 }
603 outputs.push(matvec(&raw.out, hidden, value_dim, &gated));
604 }
605 outputs
606 }
607
608 fn unequal_cfg() -> GdnConfig {
612 GdnConfig {
613 hidden_dim: 3,
614 num_key_heads: 1,
615 num_value_heads: 2,
616 key_head_dim: 2,
617 value_head_dim: 3,
618 conv_kernel_size: 3,
619 rms_norm_eps: 1e-5,
620 }
621 }
622
623 fn unequal_raw(cfg: &GdnConfig) -> RawGdn {
624 RawGdn {
625 qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
626 gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
627 conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
628 dt: vec![0.1, -0.05],
629 a: vec![-0.5, -0.75],
630 beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
631 alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
632 norm: vec![1.1, 0.9, 1.3],
633 out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
634 }
635 }
636
637 #[test]
638 fn gdn_forward_token_tiny_dims_finite_and_shaped() {
639 let weights = make_weights();
640 let cfg = cfg();
641 let mut state = GdnState::new(&cfg);
642 let hidden = [0.2f32, -0.1, 0.3, -0.4];
643
644 let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
645 assert_eq!(out0.len(), HIDDEN);
646 assert!(out0.iter().all(|x| x.is_finite()));
647
648 let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
649 assert_eq!(out1.len(), HIDDEN);
650 assert!(out1.iter().all(|x| x.is_finite()));
651 assert!(
653 out0.iter()
654 .zip(out1.iter())
655 .any(|(a, b)| (a - b).abs() > 1e-6),
656 "recurrent state should change the second token"
657 );
658 }
659
660 #[test]
661 fn softplus_matches_closed_form_at_zero() {
662 assert!((softplus(0.0) - (2.0f32).ln()).abs() < 1e-6);
663 }
664
665 #[test]
680 fn unequal_head_geometry_matches_the_reference_split_and_replication() {
681 let cfg = unequal_cfg();
682 assert_eq!(cfg.key_dim(), 2, "key_dim = num_key_heads * key_head_dim");
684 assert_eq!(
685 cfg.value_dim(),
686 6,
687 "value_dim = num_value_heads * value_head_dim"
688 );
689 assert_eq!(
690 cfg.qkv_dim(),
691 10,
692 "conv_dim = 2*key_dim + value_dim; the equal-head formula would say 12"
693 );
694 assert_eq!(cfg.heads_per_key_group(), 2);
695
696 let raw = unequal_raw(&cfg);
697 let tokens = vec![
698 vec![0.2f32, -0.1, 0.3],
699 vec![-0.4f32, 0.25, 0.05],
700 vec![0.15f32, 0.35, -0.2],
701 ];
702 let expected = reference_forward(&raw, &cfg, &tokens);
703
704 let weights = raw.to_weights(&cfg);
705 let mut state = GdnState::new(&cfg);
706 for (t, token) in tokens.iter().enumerate() {
707 let got = gdn_forward_token(&weights, &cfg, token, &mut state);
708 assert_eq!(got.len(), cfg.hidden_dim);
709 for (i, (g, e)) in got.iter().zip(expected[t].iter()).enumerate() {
710 assert!(
711 (g - e).abs() <= 1e-6 + 1e-5 * e.abs(),
712 "token {t} dim {i}: got {g}, reference {e}"
713 );
714 }
715 }
716 }
717
718 #[test]
725 fn replicating_one_key_head_equals_a_checkpoint_with_duplicated_key_rows() {
726 let shared = unequal_cfg(); let mut duplicated = shared;
728 duplicated.num_key_heads = 2; let raw_shared = unequal_raw(&shared);
731 let hidden = shared.hidden_dim;
732 let dk = shared.key_head_dim;
733 let kernel = shared.conv_kernel_size;
734
735 let mut qkv_dup = Vec::with_capacity(duplicated.qkv_dim() * hidden);
738 let mut conv_dup = Vec::with_capacity(duplicated.qkv_dim() * kernel);
739 for part in 0..2 {
740 let w_src = part * shared.key_dim() * hidden;
742 let c_src = part * shared.key_dim() * kernel;
743 for _ in 0..2 {
744 qkv_dup.extend_from_slice(&raw_shared.qkv[w_src..w_src + dk * hidden]);
745 conv_dup.extend_from_slice(&raw_shared.conv[c_src..c_src + dk * kernel]);
746 }
747 }
748 qkv_dup.extend_from_slice(&raw_shared.qkv[2 * shared.key_dim() * hidden..]);
749 conv_dup.extend_from_slice(&raw_shared.conv[2 * shared.key_dim() * kernel..]);
750
751 let raw_dup = RawGdn {
752 qkv: qkv_dup,
753 conv: conv_dup,
754 gate: raw_shared.gate.clone(),
755 dt: raw_shared.dt.clone(),
756 a: raw_shared.a.clone(),
757 beta: raw_shared.beta.clone(),
758 alpha: raw_shared.alpha.clone(),
759 norm: raw_shared.norm.clone(),
760 out: raw_shared.out.clone(),
761 };
762
763 let w_shared = raw_shared.to_weights(&shared);
764 let w_dup = raw_dup.to_weights(&duplicated);
765 let mut s_shared = GdnState::new(&shared);
766 let mut s_dup = GdnState::new(&duplicated);
767 for token in [
768 vec![0.2f32, -0.1, 0.3],
769 vec![-0.4f32, 0.25, 0.05],
770 vec![0.15f32, 0.35, -0.2],
771 ] {
772 let a = gdn_forward_token(&w_shared, &shared, &token, &mut s_shared);
773 let b = gdn_forward_token(&w_dup, &duplicated, &token, &mut s_dup);
774 for (x, y) in a.iter().zip(b.iter()) {
775 assert!((x - y).abs() < 1e-6, "{a:?} vs {b:?}");
776 }
777 }
778 }
779
780 #[test]
790 fn equal_head_geometry_stays_bit_identical_to_the_pre_generalization_output() {
791 const GOLDEN_STEP0: [u32; HIDDEN] = [1006633802, 998763940, 3163995192, 1006633802];
792 const GOLDEN_STEP1: [u32; HIDDEN] = [1006151492, 1000425832, 3164698040, 1006151492];
793
794 let weights = make_weights();
795 let cfg = cfg();
796 assert_eq!(cfg.qkv_dim(), QKV_DIM, "equal heads keep 3 * n_heads * dim");
797 assert_eq!(cfg.value_dim(), V_DIM);
798 assert_eq!(cfg.heads_per_key_group(), 1, "no replication when K == V");
799
800 let mut state = GdnState::new(&cfg);
801 let hidden = [0.2f32, -0.1, 0.3, -0.4];
802 let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
803 let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
804
805 for (i, (got, want)) in out0.iter().zip(GOLDEN_STEP0.iter()).enumerate() {
806 assert_eq!(got.to_bits(), *want, "step 0 dim {i}: {got}");
807 }
808 for (i, (got, want)) in out1.iter().zip(GOLDEN_STEP1.iter()).enumerate() {
809 assert_eq!(got.to_bits(), *want, "step 1 dim {i}: {got}");
810 }
811 }
812
813 #[test]
818 fn recurrent_state_is_rectangular_when_key_and_value_head_dims_differ() {
819 let cfg = unequal_cfg();
820 let state = GdnState::new(&cfg);
821 assert_eq!(state.recurrent.len(), 2 * 3 * 2);
822 assert!(state.recurrent.iter().all(|x| *x == 0.0));
823 }
824
825 #[test]
826 #[should_panic(expected = "positive multiple")]
827 fn value_heads_not_a_multiple_of_key_heads_is_rejected_not_floored() {
828 let cfg = GdnConfig {
829 hidden_dim: 3,
830 num_key_heads: 3,
831 num_value_heads: 4,
832 key_head_dim: 2,
833 value_head_dim: 2,
834 conv_kernel_size: 2,
835 rms_norm_eps: 1e-5,
836 };
837 let raw = RawGdn {
838 qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
839 gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
840 conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
841 dt: vec![0.0; cfg.num_value_heads],
842 a: vec![-0.5; cfg.num_value_heads],
843 beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
844 alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
845 norm: vec![1.0; cfg.value_head_dim],
846 out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
847 };
848 let weights = raw.to_weights(&cfg);
849 let mut state = GdnState::new(&cfg);
850 gdn_forward_token(&weights, &cfg, &[0.1, 0.2, 0.3], &mut state);
851 }
852}