1use crate::pool::Pool;
27use crate::qtensor::QTensor;
28
29pub struct VmfPhaseWeights {
31 pub thq: QTensor,
33 pub thk: QTensor,
35 pub v_proj: QTensor,
37 pub out_proj: QTensor,
39 pub decay: Vec<f64>,
41 pub k_gate: Option<(QTensor, Vec<f32>)>,
48}
49
50#[derive(Clone, Copy)]
51pub struct VmfPhaseCfg {
52 pub num_heads: usize,
53 pub nphase: usize,
54 pub value_head_dim: usize,
55 pub hidden_size: usize,
56 pub phase_mass: f32,
65}
66
67impl VmfPhaseCfg {
68 pub fn state_len(&self) -> usize {
69 self.num_heads * 2 * self.nphase * self.value_head_dim
70 }
71}
72
73fn phase_step(
76 thq: &[f32],
77 thk: &[f32],
78 v: &[f32],
79 decay: &[f64],
80 kap: Option<&[f32]>,
81 cfg: &VmfPhaseCfg,
82 state: &mut [f64],
83 out: &mut [f32],
84) {
85 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
86 let mscale = 1.0f64 / (1.0 + cfg.phase_mass as f64);
88 let p2 = 2 * nph;
89 for h in 0..nh {
90 let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
91 let thk_h = &thk[h * nph..(h + 1) * nph];
92 let thq_h = &thq[h * nph..(h + 1) * nph];
93 let vt = &v[h * dv..(h + 1) * dv];
94 let ot = &mut out[h * dv..(h + 1) * dv];
95 let dec = &decay[h * p2..(h + 1) * p2];
96 let kh = kap.map_or(1.0f64, |k| k[h] as f64);
98 for f in 0..p2 {
99 let (fk, fq) = if f < nph {
101 ((thk_h[f] as f64 * mscale).cos(), (thq_h[f] as f64 * mscale).cos())
102 } else {
103 (
104 (thk_h[f - nph] as f64 * mscale).sin(),
105 (thq_h[f - nph] as f64 * mscale).sin(),
106 )
107 };
108 let fkw = fk * kh;
109 let row = &mut s[f * dv..(f + 1) * dv];
110 let dcf = dec[f];
111 for d in 0..dv {
112 row[d] = dcf * row[d] + fkw * vt[d] as f64; ot[d] += (fq * row[d]) as f32; }
115 }
116 }
117}
118
119fn kappa_of(x: &[f32], w: &VmfPhaseWeights, nh: usize, pool: Option<&Pool>) -> Option<Vec<f32>> {
122 let (kw, kb) = w.k_gate.as_ref()?;
123 let mut k = vec![0.0f32; nh];
124 kw.matvec(x, &mut k, pool);
125 for (v, b) in k.iter_mut().zip(kb) {
126 *v = 1.0 / (1.0 + (-(*v + b)).exp());
127 }
128 Some(k)
129}
130
131pub fn vmf_phase_forward(
133 x: &[f32],
134 w: &VmfPhaseWeights,
135 cfg: &VmfPhaseCfg,
136 state: &mut Vec<f64>,
137 pool: Option<&Pool>,
138) -> Vec<f32> {
139 if state.len() != cfg.state_len() {
140 *state = vec![0f64; cfg.state_len()];
141 }
142 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
143
144 let mut thq = vec![0.0f32; nh * nph];
145 w.thq.matvec(x, &mut thq, pool);
146 let mut thk = vec![0.0f32; nh * nph];
147 w.thk.matvec(x, &mut thk, pool);
148 let mut v = vec![0.0f32; nh * dv];
149 w.v_proj.matvec(x, &mut v, pool);
150
151 let kap = kappa_of(x, w, nh, pool);
152 let mut o = vec![0.0f32; nh * dv];
153 phase_step(&thq, &thk, &v, &w.decay, kap.as_deref(), cfg, state, &mut o);
154
155 let mut out = vec![0.0f32; cfg.hidden_size];
156 w.out_proj.matvec(&o, &mut out, pool);
157 out
158}
159
160#[allow(clippy::too_many_arguments)]
165pub fn vmf_phase_pair(
166 x1: &[f32],
167 x2: &[f32],
168 w: &VmfPhaseWeights,
169 cfg: &VmfPhaseCfg,
170 state: &mut Vec<f64>,
171 scratch: &mut Vec<f64>,
172 pool: Option<&Pool>,
173) -> (Vec<f32>, Vec<f32>) {
174 if state.len() != cfg.state_len() {
175 *state = vec![0f64; cfg.state_len()];
176 }
177 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
178
179 let mut thq1 = vec![0.0f32; nh * nph];
180 let mut thq2 = vec![0.0f32; nh * nph];
181 w.thq.matvec2(x1, x2, &mut thq1, &mut thq2, pool);
182 let mut thk1 = vec![0.0f32; nh * nph];
183 let mut thk2 = vec![0.0f32; nh * nph];
184 w.thk.matvec2(x1, x2, &mut thk1, &mut thk2, pool);
185 let mut v1 = vec![0.0f32; nh * dv];
186 let mut v2 = vec![0.0f32; nh * dv];
187 w.v_proj.matvec2(x1, x2, &mut v1, &mut v2, pool);
188
189 let kap1 = kappa_of(x1, w, nh, pool);
191 let mut o1 = vec![0.0f32; nh * dv];
192 phase_step(&thq1, &thk1, &v1, &w.decay, kap1.as_deref(), cfg, state, &mut o1);
193
194 let kap2 = kappa_of(x2, w, nh, pool);
196 scratch.clear();
197 scratch.extend_from_slice(state);
198 let mut o2 = vec![0.0f32; nh * dv];
199 phase_step(&thq2, &thk2, &v2, &w.decay, kap2.as_deref(), cfg, scratch, &mut o2);
200
201 let mut out1 = vec![0.0f32; cfg.hidden_size];
202 let mut out2 = vec![0.0f32; cfg.hidden_size];
203 w.out_proj.matvec2(&o1, &o2, &mut out1, &mut out2, pool);
204 (out1, out2)
205}
206
207pub struct GdnWeights {
212 pub in_proj_qkv: QTensor,
214 pub in_proj_z: QTensor,
216 pub in_proj_a: QTensor,
218 pub in_proj_b: QTensor,
220 pub conv1d: Vec<f32>,
222 pub a_log: Vec<f32>,
224 pub dt_bias: Vec<f32>,
226 pub norm: Vec<f32>,
228 pub out_proj: QTensor,
230}
231
232#[derive(Clone, Copy)]
233pub struct GdnCfg {
234 pub num_v_heads: usize,
235 pub num_k_heads: usize,
236 pub key_head_dim: usize,
237 pub value_head_dim: usize,
238 pub conv_kernel: usize,
239 pub hidden_size: usize,
240 pub rms_eps: f64,
241}
242
243impl GdnCfg {
244 pub fn conv_dim(&self) -> usize {
245 2 * self.num_k_heads * self.key_head_dim + self.num_v_heads * self.value_head_dim
246 }
247
248 pub fn state_len(&self) -> usize {
251 (self.conv_kernel - 1) * self.conv_dim()
252 + self.num_v_heads * self.key_head_dim * self.value_head_dim
253 }
254}
255
256fn softplus(x: f64) -> f64 {
257 if x > 20.0 {
258 x
259 } else {
260 x.exp().ln_1p()
261 }
262}
263
264fn sigmoid(x: f64) -> f64 {
265 1.0 / (1.0 + (-x).exp())
266}
267
268fn silu(x: f64) -> f64 {
269 x / (1.0 + (-x).exp())
270}
271
272#[allow(clippy::too_many_arguments)]
276fn gdn_step(qkv: &[f32], z: &[f32], a: &[f32], b: &[f32], w: &GdnWeights, cfg: &GdnCfg, state: &mut [f64], of: &mut [f32]) {
277 let (nv, nk, dk, dv, kk) = (
278 cfg.num_v_heads,
279 cfg.num_k_heads,
280 cfg.key_head_dim,
281 cfg.value_head_dim,
282 cfg.conv_kernel,
283 );
284 let c_dim = cfg.conv_dim();
285 let (kd, rep) = (nk * dk, nv / nk);
286 let (ring, s_all) = state.split_at_mut((kk - 1) * c_dim);
287
288 let mut cq = vec![0f64; c_dim];
291 for c in 0..c_dim {
292 let taps = &w.conv1d[c * kk..(c + 1) * kk];
293 let mut acc = qkv[c] as f64 * taps[kk - 1] as f64;
294 for j in 0..kk - 1 {
295 acc += ring[j * c_dim + c] * taps[j] as f64;
296 }
297 cq[c] = silu(acc);
298 }
299 if kk > 1 {
301 ring.copy_within(c_dim.., 0);
302 let tail = (kk - 2) * c_dim;
303 for c in 0..c_dim {
304 ring[tail + c] = qkv[c] as f64;
305 }
306 }
307
308 for h in 0..nv {
309 let ko = h / rep; let (qs, ks) = (ko * dk, kd + ko * dk);
311 let (mut nq, mut nkn) = (0f64, 0f64);
313 for d in 0..dk {
314 nq += cq[qs + d] * cq[qs + d];
315 nkn += cq[ks + d] * cq[ks + d];
316 }
317 let invq = 1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt());
318 let invk = 1.0 / (nkn + 1e-6).sqrt();
319
320 let g = (-(w.a_log[h] as f64).exp() * softplus(a[h] as f64 + w.dt_bias[h] as f64)).exp();
321 let beta = sigmoid(b[h] as f64);
322
323 let s = &mut s_all[h * dk * dv..(h + 1) * dk * dv];
324 let vt = &cq[2 * kd + h * dv..2 * kd + (h + 1) * dv];
325 let mut kv = vec![0f64; dv];
327 for di in 0..dk {
328 let kf = cq[ks + di] * invk;
329 let row = &mut s[di * dv..(di + 1) * dv];
330 for dj in 0..dv {
331 row[dj] *= g;
332 kv[dj] += row[dj] * kf;
333 }
334 }
335 let mut o = vec![0f64; dv];
336 for di in 0..dk {
337 let kf = cq[ks + di] * invk;
338 let qf = cq[qs + di] * invq;
339 let row = &mut s[di * dv..(di + 1) * dv];
340 for dj in 0..dv {
341 row[dj] += kf * (vt[dj] - kv[dj]) * beta;
342 o[dj] += qf * row[dj];
343 }
344 }
345 let ss: f64 = o.iter().map(|v| v * v).sum();
347 let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
348 for dj in 0..dv {
349 of[h * dv + dj] =
350 ((o[dj] * inv) * w.norm[dj] as f64 * silu(z[h * dv + dj] as f64)) as f32;
351 }
352 }
353}
354
355pub fn gdn_forward(
357 x: &[f32],
358 w: &GdnWeights,
359 cfg: &GdnCfg,
360 state: &mut Vec<f64>,
361 pool: Option<&Pool>,
362) -> Vec<f32> {
363 if state.len() != cfg.state_len() {
364 *state = vec![0f64; cfg.state_len()];
365 }
366 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
367
368 let mut qkv = vec![0.0f32; c_dim];
369 let mut z = vec![0.0f32; vd];
370 if !gdn_projs_gpu(w, x, &mut qkv, &mut z) {
374 w.in_proj_qkv.matvec(x, &mut qkv, pool);
375 w.in_proj_z.matvec(x, &mut z, pool);
376 }
377 let mut a = vec![0.0f32; cfg.num_v_heads];
378 w.in_proj_a.matvec(x, &mut a, pool);
379 let mut b = vec![0.0f32; cfg.num_v_heads];
380 w.in_proj_b.matvec(x, &mut b, pool);
381
382 let mut of = vec![0.0f32; vd];
383 gdn_step(&qkv, &z, &a, &b, w, cfg, state, &mut of);
384
385 let mut out = vec![0.0f32; cfg.hidden_size];
386 w.out_proj.matvec(&of, &mut out, pool);
387 out
388}
389
390pub fn gdn_forward_batch(
395 xs: &[f32],
396 b: usize,
397 w: &GdnWeights,
398 cfg: &GdnCfg,
399 state: &mut Vec<f64>,
400 pool: Option<&Pool>,
401) -> Vec<f32> {
402 if state.len() != cfg.state_len() {
403 *state = vec![0f64; cfg.state_len()];
404 }
405 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
406 let nv = cfg.num_v_heads;
407
408 let mut qkv = vec![0.0f32; b * c_dim];
409 w.in_proj_qkv.matmat(xs, b, &mut qkv, pool);
410 let mut z = vec![0.0f32; b * vd];
411 w.in_proj_z.matmat(xs, b, &mut z, pool);
412 let mut a = vec![0.0f32; b * nv];
413 w.in_proj_a.matmat(xs, b, &mut a, pool);
414 let mut bb = vec![0.0f32; b * nv];
415 w.in_proj_b.matmat(xs, b, &mut bb, pool);
416
417 let mut of = vec![0.0f32; b * vd];
418 for bi in 0..b {
419 gdn_step(
420 &qkv[bi * c_dim..(bi + 1) * c_dim],
421 &z[bi * vd..(bi + 1) * vd],
422 &a[bi * nv..(bi + 1) * nv],
423 &bb[bi * nv..(bi + 1) * nv],
424 w,
425 cfg,
426 state,
427 &mut of[bi * vd..(bi + 1) * vd],
428 );
429 }
430 let mut out = vec![0.0f32; b * cfg.hidden_size];
431 w.out_proj.matmat(&of, b, &mut out, pool);
432 out
433}
434
435fn gdn_projs_gpu(w: &GdnWeights, x: &[f32], qkv: &mut [f32], z: &mut [f32]) -> bool {
437 use crate::gpu::matvec_batch;
438 use crate::qtensor::QTensor;
439 if !crate::gpu::enabled_here()
443 || !std::env::var("CMF_GPU_GDN").map(|v| v == "1").unwrap_or(false)
444 {
445 return false;
446 }
447 fn part<'a>(
448 t: &'a QTensor,
449 x: &[f32],
450 ) -> Option<(std::sync::Arc<cortiq_core::CmfModel>, crate::gpu::BatchJob<'a>)> {
451 use crate::gpu::BatchJob;
452 use crate::qtensor::prescale;
453 use cortiq_core::TensorDtype;
454 match t {
455 QTensor::Mapped {
456 model,
457 idx,
458 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
459 rows,
460 cols,
461 row_scale,
462 col_field,
463 ..
464 } => Some((
465 model.clone(),
466 BatchJob {
467 idx: *idx,
468 rows: *rows,
469 cols: *cols,
470 row_scale,
471 xs: prescale(x, col_field, *dt).into_owned(),
472 },
473 )),
474 _ => None,
475 }
476 }
477 let Some((model, jq)) = part(&w.in_proj_qkv, x) else { return false };
478 let Some((_, jz)) = part(&w.in_proj_z, x) else { return false };
479 matvec_batch(&model, &[jq, jz], &mut [qkv, z])
480}
481
482#[allow(clippy::too_many_arguments)]
485pub fn gdn_pair(
486 x1: &[f32],
487 x2: &[f32],
488 w: &GdnWeights,
489 cfg: &GdnCfg,
490 state: &mut Vec<f64>,
491 scratch: &mut Vec<f64>,
492 pool: Option<&Pool>,
493) -> (Vec<f32>, Vec<f32>) {
494 if state.len() != cfg.state_len() {
495 *state = vec![0f64; cfg.state_len()];
496 }
497 let (c_dim, vd, nv) = (
498 cfg.conv_dim(),
499 cfg.num_v_heads * cfg.value_head_dim,
500 cfg.num_v_heads,
501 );
502
503 let mut qkv1 = vec![0.0f32; c_dim];
504 let mut qkv2 = vec![0.0f32; c_dim];
505 w.in_proj_qkv.matvec2(x1, x2, &mut qkv1, &mut qkv2, pool);
506 let mut z1 = vec![0.0f32; vd];
507 let mut z2 = vec![0.0f32; vd];
508 w.in_proj_z.matvec2(x1, x2, &mut z1, &mut z2, pool);
509 let mut a1 = vec![0.0f32; nv];
510 let mut a2 = vec![0.0f32; nv];
511 w.in_proj_a.matvec2(x1, x2, &mut a1, &mut a2, pool);
512 let mut b1 = vec![0.0f32; nv];
513 let mut b2 = vec![0.0f32; nv];
514 w.in_proj_b.matvec2(x1, x2, &mut b1, &mut b2, pool);
515
516 let mut of1 = vec![0.0f32; vd];
517 gdn_step(&qkv1, &z1, &a1, &b1, w, cfg, state, &mut of1);
518
519 scratch.clear();
520 scratch.extend_from_slice(state);
521 let mut of2 = vec![0.0f32; vd];
522 gdn_step(&qkv2, &z2, &a2, &b2, w, cfg, scratch, &mut of2);
523
524 let mut out1 = vec![0.0f32; cfg.hidden_size];
525 let mut out2 = vec![0.0f32; cfg.hidden_size];
526 w.out_proj.matvec2(&of1, &of2, &mut out1, &mut out2, pool);
527 (out1, out2)
528}
529
530#[cfg(test)]
531mod tests {
532 use super::*;
533
534 fn tiny() -> (VmfPhaseWeights, VmfPhaseCfg) {
535 let cfg = VmfPhaseCfg {
536 num_heads: 2,
537 nphase: 3,
538 value_head_dim: 4,
539 hidden_size: 8,
540 phase_mass: 0.0,
541 };
542 let synth = |rows: usize, cols: usize, salt: usize| {
543 QTensor::from_f32(
544 (0..rows * cols)
545 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
546 .collect(),
547 rows,
548 cols,
549 )
550 };
551 let w = VmfPhaseWeights {
552 thq: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 1),
553 thk: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 2),
554 v_proj: synth(cfg.num_heads * cfg.value_head_dim, cfg.hidden_size, 3),
555 out_proj: synth(cfg.hidden_size, cfg.num_heads * cfg.value_head_dim, 4),
556 decay: (0..cfg.num_heads * 2 * cfg.nphase)
557 .map(|i| 0.9 + 0.005 * (i % 10) as f64)
558 .collect(),
559 k_gate: None,
560 };
561 (w, cfg)
562 }
563
564 #[test]
565 fn state_persists_and_changes_output() {
566 let (w, cfg) = tiny();
567 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
568 let mut state = Vec::new();
569 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
570 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
571 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
573 assert_eq!(state.len(), cfg.state_len());
574 }
575
576 #[test]
580 fn phase_mass_zero_is_noop_and_positive_shifts() {
581 let (w, cfg0) = tiny();
582 let mut cfg_m = cfg0.clone();
583 cfg_m.phase_mass = 1.0;
584 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.4).sin()).collect();
585
586 let mut s0 = Vec::new();
587 let base = vmf_phase_forward(&x, &w, &cfg0, &mut s0, None);
588 let mut s0b = Vec::new();
590 let base2 = vmf_phase_forward(&x, &w, &cfg0, &mut s0b, None);
591 assert_eq!(base, base2, "mass=0 must be deterministic/no-op");
592 let mut sm = Vec::new();
594 let massed = vmf_phase_forward(&x, &w, &cfg_m, &mut sm, None);
595 assert!(
596 base.iter().zip(&massed).any(|(a, b)| (a - b).abs() > 1e-5),
597 "mass>0 must change the output"
598 );
599 assert!(massed.iter().all(|v| v.is_finite()));
600 }
601
602 #[test]
607 fn kappa_gate_open_matches_none_and_closed_writes_nothing() {
608 let (mut w, cfg) = tiny();
609 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
610
611 let mut s_none = Vec::new();
612 let base1 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
613 let base2 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
614
615 w.k_gate = Some((
617 QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
618 vec![20.0; cfg.num_heads],
619 ));
620 let mut s_open = Vec::new();
621 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
622 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
623 for (a, b) in base1.iter().zip(&o1).chain(base2.iter().zip(&o2)) {
624 assert!((a - b).abs() < 1e-5, "open κ must match gateless: {a} vs {b}");
625 }
626
627 w.k_gate = Some((
629 QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
630 vec![-20.0; cfg.num_heads],
631 ));
632 let mut s_closed = Vec::new();
633 let oc = vmf_phase_forward(&x, &w, &cfg, &mut s_closed, None);
634 assert!(s_closed.iter().all(|&v| v.abs() < 1e-7), "closed κ: state must stay empty");
635 assert!(oc.iter().all(|&v| v.abs() < 1e-6), "closed κ: empty-condensate readout");
636 }
637
638 #[test]
639 fn pair_matches_two_singles_bitexact() {
640 let (w, cfg) = tiny();
641 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
642 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
643
644 let mut s_ref = Vec::new();
646 let r1 = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
647 let r2 = vmf_phase_forward(&x2, &w, &cfg, &mut s_ref, None);
648
649 let mut s = Vec::new();
651 let mut scratch = Vec::new();
652 let (p1, p2) = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
653 assert_eq!(r1, p1, "lane 1 must be bit-identical");
654 assert_eq!(r2, p2, "lane 2 must be bit-identical");
655 std::mem::swap(&mut s, &mut scratch);
657 assert_eq!(s, s_ref, "accepted state must equal sequential state");
658 }
659
660 #[test]
661 fn rejected_draft_leaves_state_at_lane1() {
662 let (w, cfg) = tiny();
663 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
664 let x2 = vec![0.5f32; 8];
665
666 let mut s_ref = Vec::new();
667 let _ = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
668
669 let mut s = Vec::new();
670 let mut scratch = Vec::new();
671 let _ = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
672 assert_eq!(s, s_ref);
674 }
675
676 fn tiny_gdn() -> (GdnWeights, GdnCfg) {
679 let cfg = GdnCfg {
680 num_v_heads: 4,
681 num_k_heads: 2,
682 key_head_dim: 3,
683 value_head_dim: 5,
684 conv_kernel: 4,
685 hidden_size: 8,
686 rms_eps: 1e-6,
687 };
688 let c_dim = cfg.conv_dim();
689 let vd = cfg.num_v_heads * cfg.value_head_dim;
690 let synth = |rows: usize, cols: usize, salt: usize| {
691 QTensor::from_f32(
692 (0..rows * cols)
693 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
694 .collect(),
695 rows,
696 cols,
697 )
698 };
699 let vecf = |n: usize, salt: usize| -> Vec<f32> {
700 (0..n)
701 .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.6)
702 .collect()
703 };
704 let w = GdnWeights {
705 in_proj_qkv: synth(c_dim, cfg.hidden_size, 1),
706 in_proj_z: synth(vd, cfg.hidden_size, 2),
707 in_proj_a: synth(cfg.num_v_heads, cfg.hidden_size, 3),
708 in_proj_b: synth(cfg.num_v_heads, cfg.hidden_size, 4),
709 conv1d: vecf(c_dim * cfg.conv_kernel, 5),
710 a_log: (0..cfg.num_v_heads).map(|i| 0.2 + 0.3 * i as f32).collect(),
711 dt_bias: vecf(cfg.num_v_heads, 6),
712 norm: vec![1.0; cfg.value_head_dim],
713 out_proj: synth(cfg.hidden_size, vd, 7),
714 };
715 (w, cfg)
716 }
717
718 #[test]
719 fn gdn_state_persists_and_changes_output() {
720 let (w, cfg) = tiny_gdn();
721 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
722 let mut state = Vec::new();
723 let o1 = gdn_forward(&x, &w, &cfg, &mut state, None);
724 let o2 = gdn_forward(&x, &w, &cfg, &mut state, None);
725 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
726 assert_eq!(state.len(), cfg.state_len());
727 }
728
729 #[test]
730 fn gdn_pair_matches_two_singles_bitexact() {
731 let (w, cfg) = tiny_gdn();
732 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
733 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
734
735 let mut s_ref = Vec::new();
736 let r1 = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
737 let r2 = gdn_forward(&x2, &w, &cfg, &mut s_ref, None);
738
739 let mut s = Vec::new();
740 let mut scratch = Vec::new();
741 let (p1, p2) = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
742 assert_eq!(r1, p1, "lane 1 must be bit-identical");
743 assert_eq!(r2, p2, "lane 2 must be bit-identical");
744 std::mem::swap(&mut s, &mut scratch);
745 assert_eq!(s, s_ref, "accepted state must equal sequential state");
746 }
747
748 #[test]
749 fn gdn_rejected_draft_leaves_state_at_lane1() {
750 let (w, cfg) = tiny_gdn();
751 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
752 let x2 = vec![0.5f32; 8];
753
754 let mut s_ref = Vec::new();
755 let _ = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
756
757 let mut s = Vec::new();
758 let mut scratch = Vec::new();
759 let _ = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
760 assert_eq!(s, s_ref);
761 }
762
763 #[test]
767 fn gdn_conv_ring_matches_explicit_causal_conv() {
768 let (w, cfg) = tiny_gdn();
769 let seq: Vec<Vec<f32>> = (0..6)
770 .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.17).sin()).collect())
771 .collect();
772
773 let mut s_inc = Vec::new();
776 for (t, x) in seq.iter().enumerate() {
777 let inc = gdn_forward(x, &w, &cfg, &mut s_inc, None);
778 let mut s_replay = Vec::new();
779 let mut replay = Vec::new();
780 for xr in &seq[..=t] {
781 replay = gdn_forward(xr, &w, &cfg, &mut s_replay, None);
782 }
783 assert_eq!(inc, replay, "position {t}: ring must equal replay");
784 }
785 }
786}