1use crate::pool::Pool;
32use crate::qtensor::QTensor;
33use cortiq_core::TensorDtype;
34use std::sync::OnceLock;
35use std::sync::atomic::{AtomicU64, Ordering};
36
37static PERF_GDN_FORWARD_CALLS: AtomicU64 = AtomicU64::new(0);
38static PERF_GDN_FORWARD_NS: AtomicU64 = AtomicU64::new(0);
39static PERF_GDN_BATCH_CALLS: AtomicU64 = AtomicU64::new(0);
40static PERF_GDN_BATCH_NS: AtomicU64 = AtomicU64::new(0);
41static PERF_GDN_STEP_CALLS: AtomicU64 = AtomicU64::new(0);
42static PERF_GDN_STEP_NS: AtomicU64 = AtomicU64::new(0);
43
44fn perf_enabled() -> bool {
45 static ON: OnceLock<bool> = OnceLock::new();
46 *ON.get_or_init(|| std::env::var("CMF_PERF_PROFILE").as_deref() == Ok("1"))
47}
48
49pub struct VmfPhaseWeights {
51 pub thq: QTensor,
53 pub thk: QTensor,
55 pub v_proj: QTensor,
57 pub out_proj: QTensor,
59 pub decay: Vec<f64>,
61 pub conv: Option<Vec<f32>>,
67 pub k_gate: Option<(QTensor, Vec<f32>)>,
74 pub phase_delta: bool,
79}
80
81#[derive(Clone, Copy)]
82pub struct VmfPhaseCfg {
83 pub num_heads: usize,
84 pub nphase: usize,
85 pub value_head_dim: usize,
86 pub hidden_size: usize,
87 pub phase_mass: f32,
96}
97
98impl VmfPhaseCfg {
99 pub fn state_len(&self) -> usize {
100 self.num_heads * 2 * self.nphase * self.value_head_dim
101 }
102}
103
104fn phase_delta_step_f32(
110 thq: &[f32],
111 thk: &[f32],
112 v: &[f32],
113 decay: &[f64],
114 kap: Option<&[f32]>,
115 cfg: &VmfPhaseCfg,
116 state: &mut [f32],
117 out: &mut [f32],
118) {
119 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
120 let p2 = 2 * nph;
121 let scale = 1.0f64 / (nph as f64).sqrt();
122 debug_assert_eq!(thq.len(), nh * nph);
123 debug_assert_eq!(thk.len(), nh * nph);
124 debug_assert_eq!(v.len(), nh * dv);
125 debug_assert_eq!(decay.len(), nh * p2);
126 debug_assert!(state.len() >= nh * p2 * dv);
127 debug_assert!(out.len() >= nh * dv);
128
129 let mut r = vec![0.0f64; dv];
133 let mut key = vec![0.0f64; p2];
134 for h in 0..nh {
135 let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
136 let thk_h = &thk[h * nph..(h + 1) * nph];
137 let thq_h = &thq[h * nph..(h + 1) * nph];
138 let vt = &v[h * dv..(h + 1) * dv];
139 let ot = &mut out[h * dv..(h + 1) * dv];
140 let dec = &decay[h * p2..(h + 1) * p2];
141 let beta = kap.map_or(1.0f64, |k| k[h] as f64);
142
143 r.fill(0.0);
147 for f in 0..p2 {
148 let kf = if f < nph {
149 scale * (thk_h[f] as f64).cos()
150 } else {
151 scale * (thk_h[f - nph] as f64).sin()
152 };
153 key[f] = kf;
154 let row = &s[f * dv..(f + 1) * dv];
155 for d in 0..dv {
156 r[d] += kf * (dec[f] * row[d] as f64);
157 }
158 }
159
160 for f in 0..p2 {
168 let kf = key[f];
169 #[cfg(target_arch = "aarch64")]
170 let qf = if f < nph {
171 scale * (thq_h[f] as f64).cos()
172 } else {
173 scale * (thq_h[f - nph] as f64).sin()
174 };
175 let row = &mut s[f * dv..(f + 1) * dv];
176 for d in 0..dv {
177 let p = dec[f] * row[d] as f64;
178 row[d] = (p + beta * kf * (vt[d] as f64 - r[d])) as f32;
179 #[cfg(target_arch = "aarch64")]
180 {
181 ot[d] += (qf * row[d] as f64) as f32;
182 }
183 }
184 }
185 #[cfg(not(target_arch = "aarch64"))]
186 for f in 0..p2 {
187 let qf = if f < nph {
188 scale * (thq_h[f] as f64).cos()
189 } else {
190 scale * (thq_h[f - nph] as f64).sin()
191 };
192 let row = &s[f * dv..(f + 1) * dv];
193 for d in 0..dv {
194 ot[d] += (qf * row[d] as f64) as f32;
195 }
196 }
197 }
198}
199
200fn phase_step(
204 thq: &[f32],
205 thk: &[f32],
206 v: &[f32],
207 decay: &[f64],
208 kap: Option<&[f32]>,
209 cfg: &VmfPhaseCfg,
210 state: &mut [f32],
211 out: &mut [f32],
212 phase_delta: bool,
213) {
214 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
215 if phase_delta {
216 phase_delta_step_f32(thq, thk, v, decay, kap, cfg, state, out);
217 return;
218 }
219 let mscale = 1.0f64 / (1.0 + cfg.phase_mass as f64);
221 let p2 = 2 * nph;
222 for h in 0..nh {
223 let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
224 let thk_h = &thk[h * nph..(h + 1) * nph];
225 let thq_h = &thq[h * nph..(h + 1) * nph];
226 let vt = &v[h * dv..(h + 1) * dv];
227 let ot = &mut out[h * dv..(h + 1) * dv];
228 let dec = &decay[h * p2..(h + 1) * p2];
229 let kh = kap.map_or(1.0f64, |k| k[h] as f64);
231 for f in 0..p2 {
232 let (fk, fq) = if f < nph {
234 (
235 (thk_h[f] as f64 * mscale).cos(),
236 (thq_h[f] as f64 * mscale).cos(),
237 )
238 } else {
239 (
240 (thk_h[f - nph] as f64 * mscale).sin(),
241 (thq_h[f - nph] as f64 * mscale).sin(),
242 )
243 };
244 let fkw = fk * kh;
245 let row = &mut s[f * dv..(f + 1) * dv];
246 let dcf = dec[f];
247 for d in 0..dv {
248 let cell = dcf * row[d] as f64 + fkw * vt[d] as f64;
250 row[d] = cell as f32;
251 ot[d] += (fq * cell) as f32; }
253 }
254 }
255}
256
257fn kappa_of(x: &[f32], w: &VmfPhaseWeights, nh: usize, pool: Option<&Pool>) -> Option<Vec<f32>> {
260 let (kw, kb) = w.k_gate.as_ref()?;
261 let mut k = vec![0.0f32; nh];
262 kw.matvec(x, &mut k, pool);
263 for (v, b) in k.iter_mut().zip(kb) {
264 *v = 1.0 / (1.0 + (-(*v + b)).exp());
265 }
266 Some(k)
267}
268
269fn conv_in(
273 x: &[f32],
274 w: &VmfPhaseWeights,
275 cfg: &VmfPhaseCfg,
276 state: &mut Vec<f32>,
277) -> Option<Vec<f32>> {
278 let taps = w.conv.as_ref()?;
279 let h = cfg.hidden_size;
280 let k = taps.len() / h.max(1);
281 if k < 2 || taps.len() != h * k {
282 return None;
283 }
284 let ring = (k - 1) * h;
285 let base = cfg.state_len();
286 if state.len() != base + ring {
287 let mut ns = vec![0f32; base + ring];
289 let n = state.len().min(base);
290 ns[..n].copy_from_slice(&state[..n]);
291 *state = ns;
292 }
293 let mut y = vec![0.0f32; h];
294 for c in 0..h {
295 let mut acc = taps[c * k + k - 1] * x[c];
297 for j in 0..k - 1 {
298 acc += taps[c * k + j] * state[base + j * h + c];
299 }
300 y[c] = acc;
301 }
302 conv_ring_push(x, h, base, state);
303 Some(y)
304}
305
306fn conv_ring_push(x: &[f32], h: usize, base: usize, state: &mut [f32]) {
308 let ring = state.len() - base;
309 state.copy_within(base + h.., base);
310 let at = base + ring - h;
311 state[at..at + h].copy_from_slice(&x[..h]);
312}
313
314pub fn vmf_phase_forward(
316 x: &[f32],
317 w: &VmfPhaseWeights,
318 cfg: &VmfPhaseCfg,
319 state: &mut Vec<f32>,
320 pool: Option<&Pool>,
321) -> Vec<f32> {
322 if w.conv.is_none() && state.len() != cfg.state_len() {
323 *state = vec![0f32; cfg.state_len()];
324 }
325 let xc = conv_in(x, w, cfg, state);
326 let x = xc.as_deref().unwrap_or(x);
327 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
328
329 let mut thq = vec![0.0f32; nh * nph];
330 w.thq.matvec(x, &mut thq, pool);
331 let mut thk = vec![0.0f32; nh * nph];
332 w.thk.matvec(x, &mut thk, pool);
333 let mut v = vec![0.0f32; nh * dv];
334 w.v_proj.matvec(x, &mut v, pool);
335
336 let kap = kappa_of(x, w, nh, pool);
337 let mut o = vec![0.0f32; nh * dv];
338 phase_step(
339 &thq,
340 &thk,
341 &v,
342 &w.decay,
343 kap.as_deref(),
344 cfg,
345 state,
346 &mut o,
347 w.phase_delta,
348 );
349
350 let mut out = vec![0.0f32; cfg.hidden_size];
351 w.out_proj.matvec(&o, &mut out, pool);
352 out
353}
354
355#[allow(clippy::too_many_arguments)]
360pub fn vmf_phase_pair(
361 x1: &[f32],
362 x2: &[f32],
363 w: &VmfPhaseWeights,
364 cfg: &VmfPhaseCfg,
365 state: &mut Vec<f32>,
366 scratch: &mut Vec<f32>,
367 pool: Option<&Pool>,
368) -> (Vec<f32>, Vec<f32>) {
369 if w.conv.is_none() && state.len() != cfg.state_len() {
370 *state = vec![0f32; cfg.state_len()];
371 }
372 let xc1 = conv_in(x1, w, cfg, state);
375 let x1 = xc1.as_deref().unwrap_or(x1);
376 let (xc2, x2raw) = if w.conv.is_some() {
377 let mut tmp = state.clone();
378 (conv_in(x2, w, cfg, &mut tmp), Some(x2))
379 } else {
380 (None, None)
381 };
382 let x2 = xc2.as_deref().unwrap_or(x2);
383 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
384
385 let mut thq1 = vec![0.0f32; nh * nph];
386 let mut thq2 = vec![0.0f32; nh * nph];
387 w.thq.matvec2(x1, x2, &mut thq1, &mut thq2, pool);
388 let mut thk1 = vec![0.0f32; nh * nph];
389 let mut thk2 = vec![0.0f32; nh * nph];
390 w.thk.matvec2(x1, x2, &mut thk1, &mut thk2, pool);
391 let mut v1 = vec![0.0f32; nh * dv];
392 let mut v2 = vec![0.0f32; nh * dv];
393 w.v_proj.matvec2(x1, x2, &mut v1, &mut v2, pool);
394
395 let kap1 = kappa_of(x1, w, nh, pool);
397 let mut o1 = vec![0.0f32; nh * dv];
398 phase_step(
399 &thq1,
400 &thk1,
401 &v1,
402 &w.decay,
403 kap1.as_deref(),
404 cfg,
405 state,
406 &mut o1,
407 w.phase_delta,
408 );
409
410 let kap2 = kappa_of(x2, w, nh, pool);
412 scratch.clear();
413 scratch.extend_from_slice(state);
414 let mut o2 = vec![0.0f32; nh * dv];
415 phase_step(
416 &thq2,
417 &thk2,
418 &v2,
419 &w.decay,
420 kap2.as_deref(),
421 cfg,
422 scratch,
423 &mut o2,
424 w.phase_delta,
425 );
426 if let Some(xr) = x2raw {
429 conv_ring_push(xr, cfg.hidden_size, cfg.state_len(), scratch);
430 }
431
432 let mut out1 = vec![0.0f32; cfg.hidden_size];
433 let mut out2 = vec![0.0f32; cfg.hidden_size];
434 w.out_proj.matvec2(&o1, &o2, &mut out1, &mut out2, pool);
435 (out1, out2)
436}
437
438pub struct GdnWeights {
443 pub in_proj_qkv: QTensor,
445 pub in_proj_z: QTensor,
447 pub in_proj_a: QTensor,
449 pub in_proj_b: QTensor,
451 pub conv1d: Vec<f32>,
453 pub a_log: Vec<f32>,
455 pub dt_bias: Vec<f32>,
457 pub norm: Vec<f32>,
459 pub out_proj: QTensor,
461}
462
463#[derive(Clone, Copy)]
464pub struct GdnCfg {
465 pub num_v_heads: usize,
466 pub num_k_heads: usize,
467 pub key_head_dim: usize,
468 pub value_head_dim: usize,
469 pub conv_kernel: usize,
470 pub hidden_size: usize,
471 pub rms_eps: f64,
472 pub output_gate_sigmoid: bool,
475}
476
477impl GdnCfg {
478 pub fn conv_dim(&self) -> usize {
479 2 * self.num_k_heads * self.key_head_dim + self.num_v_heads * self.value_head_dim
480 }
481
482 pub fn state_len(&self) -> usize {
485 (self.conv_kernel - 1) * self.conv_dim()
486 + self.num_v_heads * self.key_head_dim * self.value_head_dim
487 }
488}
489
490fn softplus(x: f64) -> f64 {
491 if x > 20.0 { x } else { x.exp().ln_1p() }
492}
493
494fn sigmoid(x: f64) -> f64 {
495 1.0 / (1.0 + (-x).exp())
496}
497
498fn silu(x: f64) -> f64 {
499 x / (1.0 + (-x).exp())
500}
501
502#[derive(Clone, Copy)]
505struct SendMutF32(*mut f32);
506unsafe impl Send for SendMutF32 {}
507unsafe impl Sync for SendMutF32 {}
508
509#[allow(clippy::too_many_arguments)]
522fn gdn_step(
523 qkv: &[f32],
524 z: &[f32],
525 a: &[f32],
526 b: &[f32],
527 w: &GdnWeights,
528 cfg: &GdnCfg,
529 state: &mut [f32],
530 of: &mut [f32],
531 pool: Option<&Pool>,
532) {
533 let perf_t0 = perf_enabled().then(std::time::Instant::now);
534 let (nv, nk, dk, dv, kk) = (
535 cfg.num_v_heads,
536 cfg.num_k_heads,
537 cfg.key_head_dim,
538 cfg.value_head_dim,
539 cfg.conv_kernel,
540 );
541 let c_dim = cfg.conv_dim();
542 let (kd, rep) = (nk * dk, nv / nk);
543 let (ring, s_all) = state.split_at_mut((kk - 1) * c_dim);
544
545 let mut cq = vec![0f32; c_dim];
549 for c in 0..c_dim {
550 let taps = &w.conv1d[c * kk..(c + 1) * kk];
551 let mut acc = qkv[c] as f64 * taps[kk - 1] as f64;
552 for j in 0..kk - 1 {
553 acc += ring[j * c_dim + c] as f64 * taps[j] as f64;
554 }
555 cq[c] = silu(acc) as f32;
556 }
557 if kk > 1 {
559 ring.copy_within(c_dim.., 0);
560 let tail = (kk - 2) * c_dim;
561 ring[tail..tail + c_dim].copy_from_slice(&qkv[..c_dim]);
562 }
563
564 let cq = &cq;
565 let s_ptr = SendMutF32(s_all.as_mut_ptr());
566 let of_ptr = SendMutF32(of.as_mut_ptr());
567 let head_range = |h0: usize, h1: usize| {
568 let (s_ptr, of_ptr) = (s_ptr, of_ptr);
571 let mut kv = crate::attention::take_buf(dv);
573 let mut delta = crate::attention::take_buf(dv);
574 let mut o = crate::attention::take_buf(dv);
575 let mut kf = crate::attention::take_buf(dk);
576 let mut qf = crate::attention::take_buf(dk);
577 for h in h0..h1 {
578 let ko = h / rep; let (qs, ks) = (ko * dk, kd + ko * dk);
580 let (mut nq, mut nkn) = (0f64, 0f64);
582 for d in 0..dk {
583 nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
584 nkn += (cq[ks + d] as f64) * (cq[ks + d] as f64);
585 }
586 let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
587 let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
588 for d in 0..dk {
589 qf[d] = cq[qs + d] * invq;
590 kf[d] = cq[ks + d] * invk;
591 }
592
593 let g = (-(w.a_log[h] as f64).exp() * softplus(a[h] as f64 + w.dt_bias[h] as f64)).exp()
594 as f32;
595 let beta = sigmoid(b[h] as f64) as f32;
596
597 let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
599 let oh = unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
600 let vt = &cq[2 * kd + h * dv..2 * kd + (h + 1) * dv];
601
602 kv[..dv].fill(0.0);
606 for di in 0..dk {
607 let kfd = kf[di];
608 let row = &s[di * dv..(di + 1) * dv];
609 for dj in 0..dv {
610 kv[dj] += row[dj] * kfd; }
612 }
613 for dj in 0..dv {
614 delta[dj] = (vt[dj] - g * kv[dj]) * beta;
615 }
616 o[..dv].fill(0.0);
617 for di in 0..dk {
618 let kfd = kf[di];
619 let qfd = qf[di];
620 let row = &mut s[di * dv..(di + 1) * dv];
621 for dj in 0..dv {
622 let cell = g * row[dj] + kfd * delta[dj];
623 row[dj] = cell;
624 o[dj] += qfd * cell; }
626 }
627 let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
629 let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
630 for dj in 0..dv {
631 let gate = if cfg.output_gate_sigmoid {
632 sigmoid(z[h * dv + dj] as f64)
633 } else {
634 silu(z[h * dv + dj] as f64)
635 };
636 oh[dj] = ((o[dj] as f64 * inv) * w.norm[dj] as f64 * gate) as f32;
637 }
638 }
639 crate::attention::recycle_buf(&mut kv);
640 crate::attention::recycle_buf(&mut delta);
641 crate::attention::recycle_buf(&mut o);
642 crate::attention::recycle_buf(&mut kf);
643 crate::attention::recycle_buf(&mut qf);
644 };
645 match pool {
646 Some(pool) if nv >= 4 => pool.run(&|widx, n| {
647 let chunk = nv.div_ceil(n);
648 let h0 = (widx * chunk).min(nv);
649 let h1 = (h0 + chunk).min(nv);
650 if h0 < h1 {
651 head_range(h0, h1);
652 }
653 }),
654 _ => head_range(0, nv),
655 }
656 if let Some(t0) = perf_t0 {
657 PERF_GDN_STEP_CALLS.fetch_add(1, Ordering::Relaxed);
658 PERF_GDN_STEP_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
659 }
660}
661
662pub fn gdn_forward(
664 x: &[f32],
665 w: &GdnWeights,
666 cfg: &GdnCfg,
667 state: &mut Vec<f32>,
668 pool: Option<&Pool>,
669) -> Vec<f32> {
670 let perf_t0 = perf_enabled().then(std::time::Instant::now);
671 if state.len() != cfg.state_len() {
672 *state = vec![0f32; cfg.state_len()];
673 }
674 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
675
676 let mut qkv = vec![0.0f32; c_dim];
677 let mut z = vec![0.0f32; vd];
678 let mut a = vec![0.0f32; cfg.num_v_heads];
679 let mut b = vec![0.0f32; cfg.num_v_heads];
680 let cpu_projs = |qkv: &mut Vec<f32>, z: &mut Vec<f32>, a: &mut Vec<f32>, b: &mut Vec<f32>| {
684 QTensor::matvec_many(
685 [&w.in_proj_qkv, &w.in_proj_z, &w.in_proj_a, &w.in_proj_b],
686 x,
687 [
688 qkv.as_mut_slice(),
689 z.as_mut_slice(),
690 a.as_mut_slice(),
691 b.as_mut_slice(),
692 ],
693 pool,
694 );
695 };
696 let mut done = false;
697 if crate::gpu::enabled_here() && gdn_projs_eligible(w) {
698 match crate::gpu::probe_arm(crate::gpu::OpClass::Batch) {
699 crate::gpu::ProbeArm::Gpu => {
700 let t0 = std::time::Instant::now();
701 if gdn_projs_gpu(w, x, &mut qkv, &mut z) {
702 crate::gpu::probe_record(crate::gpu::OpClass::Batch, true, t0.elapsed());
703 w.in_proj_a.matvec(x, &mut a, pool);
704 w.in_proj_b.matvec(x, &mut b, pool);
705 done = true;
706 } else {
707 crate::gpu::probe_note_decline(crate::gpu::OpClass::Batch);
708 }
709 }
710 crate::gpu::ProbeArm::CpuTimed => {
711 let t0 = std::time::Instant::now();
712 crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
713 crate::gpu::probe_record(crate::gpu::OpClass::Batch, false, t0.elapsed());
714 done = true;
715 }
716 crate::gpu::ProbeArm::Cpu => {
717 crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
718 done = true;
719 }
720 }
721 }
722 if !done {
723 cpu_projs(&mut qkv, &mut z, &mut a, &mut b);
724 }
725
726 let mut of = vec![0.0f32; vd];
727 gdn_step(&qkv, &z, &a, &b, w, cfg, state, &mut of, pool);
728
729 let mut out = vec![0.0f32; cfg.hidden_size];
730 w.out_proj.matvec(&of, &mut out, pool);
731 if let Some(t0) = perf_t0 {
732 PERF_GDN_FORWARD_CALLS.fetch_add(1, Ordering::Relaxed);
733 PERF_GDN_FORWARD_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
734 }
735 out
736}
737
738pub fn gdn_forward_batch(
743 xs: &[f32],
744 b: usize,
745 w: &GdnWeights,
746 cfg: &GdnCfg,
747 state: &mut Vec<f32>,
748 pool: Option<&Pool>,
749) -> Vec<f32> {
750 let perf_t0 = perf_enabled().then(std::time::Instant::now);
751 if state.len() != cfg.state_len() {
752 *state = vec![0f32; cfg.state_len()];
753 }
754 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
755 let nv = cfg.num_v_heads;
756
757 let mut qkv = vec![0.0f32; b * c_dim];
758 w.in_proj_qkv.matmat(xs, b, &mut qkv, pool);
759 let mut z = vec![0.0f32; b * vd];
760 w.in_proj_z.matmat(xs, b, &mut z, pool);
761 let mut a = vec![0.0f32; b * nv];
762 w.in_proj_a.matmat(xs, b, &mut a, pool);
763 let mut bb = vec![0.0f32; b * nv];
764 w.in_proj_b.matmat(xs, b, &mut bb, pool);
765
766 let mut of = vec![0.0f32; b * vd];
767 for bi in 0..b {
768 gdn_step(
769 &qkv[bi * c_dim..(bi + 1) * c_dim],
770 &z[bi * vd..(bi + 1) * vd],
771 &a[bi * nv..(bi + 1) * nv],
772 &bb[bi * nv..(bi + 1) * nv],
773 w,
774 cfg,
775 state,
776 &mut of[bi * vd..(bi + 1) * vd],
777 pool,
778 );
779 }
780 let mut out = vec![0.0f32; b * cfg.hidden_size];
781 w.out_proj.matmat(&of, b, &mut out, pool);
782 if std::env::var("CMF_GDN_TRACE").is_ok() {
783 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
784 eprintln!(
785 "gdn-batch b={b}: |x0|={:.5} |qkv0|={:.5} |z0|={:.5} |a0|={:.5} |b0|={:.5} |of0|={:.5} |out0|={:.5} |state|={:.5}",
786 n(&xs[..cfg.hidden_size]),
787 n(&qkv[..c_dim]),
788 n(&z[..vd]),
789 n(&a[..nv]),
790 n(&bb[..nv]),
791 n(&of[..vd]),
792 n(&out[..cfg.hidden_size]),
793 n(state)
794 );
795 }
796 if let Some(t0) = perf_t0 {
797 PERF_GDN_BATCH_CALLS.fetch_add(1, Ordering::Relaxed);
798 PERF_GDN_BATCH_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
799 }
800 out
801}
802
803pub fn perf_report() {
805 if !perf_enabled() {
806 return;
807 }
808 let fc = PERF_GDN_FORWARD_CALLS.load(Ordering::Relaxed);
809 let bc = PERF_GDN_BATCH_CALLS.load(Ordering::Relaxed);
810 let sc = PERF_GDN_STEP_CALLS.load(Ordering::Relaxed);
811 eprintln!(
812 "[perf-gdn] forward_calls={} forward_ms={:.3} forward_ms_per_call={:.3} batch_calls={} batch_ms={:.3} step_calls={} step_ms={:.3} step_ms_per_call={:.3}",
813 fc,
814 PERF_GDN_FORWARD_NS.load(Ordering::Relaxed) as f64 / 1e6,
815 PERF_GDN_FORWARD_NS.load(Ordering::Relaxed) as f64 / 1e6 / fc.max(1) as f64,
816 bc,
817 PERF_GDN_BATCH_NS.load(Ordering::Relaxed) as f64 / 1e6,
818 sc,
819 PERF_GDN_STEP_NS.load(Ordering::Relaxed) as f64 / 1e6,
820 PERF_GDN_STEP_NS.load(Ordering::Relaxed) as f64 / 1e6 / sc.max(1) as f64,
821 );
822}
823
824fn gdn_projs_eligible(w: &GdnWeights) -> bool {
828 if std::env::var("CMF_GPU_GDN")
834 .map(|v| v == "0")
835 .unwrap_or(false)
836 {
837 return false;
838 }
839 w.in_proj_qkv.is_q1()
840 || w.in_proj_qkv.q4t_parts().is_some()
841 || w.in_proj_qkv.q4tp_parts().is_some()
842 || std::env::var("CMF_GPU_GDN")
843 .map(|v| v == "1")
844 .unwrap_or(false)
845}
846
847fn gdn_projs_gpu(w: &GdnWeights, x: &[f32], qkv: &mut [f32], z: &mut [f32]) -> bool {
849 use crate::gpu::matvec_batch;
850 use crate::qtensor::QTensor;
851 if !crate::gpu::enabled_here() {
852 return false;
853 }
854 fn part<'a>(
855 t: &'a QTensor,
856 x: &[f32],
857 ) -> Option<(
858 std::sync::Arc<cortiq_core::CmfModel>,
859 crate::gpu::BatchJob<'a>,
860 )> {
861 use crate::gpu::BatchJob;
862 use crate::qtensor::prescale;
863 use cortiq_core::TensorDtype;
864 match t {
865 QTensor::Mapped {
866 model,
867 idx,
868 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
869 rows,
870 cols,
871 row_scale,
872 col_field,
873 ..
874 } => Some((
875 model.clone(),
876 BatchJob {
877 idx: *idx,
878 rows: *rows,
879 cols: *cols,
880 row_scale,
881 xs: prescale(x, col_field, *dt).into_owned(),
882 layout: crate::gpu::BatchLayout::Q8,
883 },
884 )),
885 QTensor::Mapped {
886 model,
887 idx,
888 dtype: TensorDtype::Q1,
889 rows,
890 cols,
891 ..
892 } => Some((
893 model.clone(),
894 BatchJob {
895 idx: *idx,
896 rows: *rows,
897 cols: *cols,
898 row_scale: &[],
899 xs: x.to_vec(),
900 layout: crate::gpu::BatchLayout::Q1,
901 },
902 )),
903 QTensor::Mapped {
907 model,
908 idx,
909 dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
910 rows,
911 cols,
912 ..
913 } => Some((
914 model.clone(),
915 BatchJob {
916 idx: *idx,
917 rows: *rows,
918 cols: *cols,
919 row_scale: &[],
920 xs: x.to_vec(),
921 layout: if *dt == TensorDtype::Q4TiledP {
922 crate::gpu::BatchLayout::Q4tp
923 } else {
924 crate::gpu::BatchLayout::Q4t
925 },
926 },
927 )),
928 _ => None,
929 }
930 }
931 let Some((model, jq)) = part(&w.in_proj_qkv, x) else {
932 return false;
933 };
934 let Some((_, jz)) = part(&w.in_proj_z, x) else {
935 return false;
936 };
937 matvec_batch(&model, &[jq, jz], &mut [qkv, z])
938}
939
940#[allow(clippy::too_many_arguments)]
943pub fn gdn_pair(
944 x1: &[f32],
945 x2: &[f32],
946 w: &GdnWeights,
947 cfg: &GdnCfg,
948 state: &mut Vec<f32>,
949 scratch: &mut Vec<f32>,
950 pool: Option<&Pool>,
951) -> (Vec<f32>, Vec<f32>) {
952 if state.len() != cfg.state_len() {
953 *state = vec![0f32; cfg.state_len()];
954 }
955 let (c_dim, vd, nv) = (
956 cfg.conv_dim(),
957 cfg.num_v_heads * cfg.value_head_dim,
958 cfg.num_v_heads,
959 );
960
961 let mut qkv1 = vec![0.0f32; c_dim];
962 let mut qkv2 = vec![0.0f32; c_dim];
963 w.in_proj_qkv.matvec2(x1, x2, &mut qkv1, &mut qkv2, pool);
964 let mut z1 = vec![0.0f32; vd];
965 let mut z2 = vec![0.0f32; vd];
966 w.in_proj_z.matvec2(x1, x2, &mut z1, &mut z2, pool);
967 let mut a1 = vec![0.0f32; nv];
968 let mut a2 = vec![0.0f32; nv];
969 w.in_proj_a.matvec2(x1, x2, &mut a1, &mut a2, pool);
970 let mut b1 = vec![0.0f32; nv];
971 let mut b2 = vec![0.0f32; nv];
972 w.in_proj_b.matvec2(x1, x2, &mut b1, &mut b2, pool);
973
974 let mut of1 = vec![0.0f32; vd];
975 gdn_step(&qkv1, &z1, &a1, &b1, w, cfg, state, &mut of1, pool);
976
977 scratch.clear();
978 scratch.extend_from_slice(state);
979 let mut of2 = vec![0.0f32; vd];
980 gdn_step(&qkv2, &z2, &a2, &b2, w, cfg, scratch, &mut of2, pool);
981
982 let mut out1 = vec![0.0f32; cfg.hidden_size];
983 let mut out2 = vec![0.0f32; cfg.hidden_size];
984 w.out_proj.matvec2(&of1, &of2, &mut out1, &mut out2, pool);
985 (out1, out2)
986}
987
988pub struct ShortConvWeights {
995 pub in_proj: QTensor,
997 pub conv: Vec<f32>,
1001 pub out_proj: QTensor,
1003}
1004
1005#[derive(Clone, Copy)]
1006pub struct ShortConvCfg {
1007 pub hidden_size: usize,
1008 pub kernel: usize,
1010}
1011
1012impl ShortConvCfg {
1013 pub fn state_len(&self) -> usize {
1015 (self.kernel - 1) * self.hidden_size
1016 }
1017}
1018
1019fn short_conv_step(
1029 bcx: &[f32],
1030 conv: &[f32],
1031 cfg: &ShortConvCfg,
1032 ring_state: &mut [f32],
1033 y: &mut [f32],
1034) {
1035 let (h, k) = (cfg.hidden_size, cfg.kernel);
1036 let ring = k - 1;
1037 let (bg, cg, xg) = (&bcx[0..h], &bcx[h..2 * h], &bcx[2 * h..3 * h]);
1038 for c in 0..h {
1039 let bx = bg[c] * xg[c];
1040 let wc = &conv[c * k..(c + 1) * k];
1041 let mut acc = wc[k - 1] * bx;
1043 let rc = &mut ring_state[c * ring..c * ring + ring];
1044 for s in 0..ring {
1045 acc += wc[k - 2 - s] * rc[s];
1046 }
1047 y[c] = cg[c] * acc;
1048 for s in (1..ring).rev() {
1050 rc[s] = rc[s - 1];
1051 }
1052 if ring > 0 {
1053 rc[0] = bx;
1054 }
1055 }
1056}
1057
1058pub fn short_conv_forward(
1060 x: &[f32],
1061 w: &ShortConvWeights,
1062 cfg: &ShortConvCfg,
1063 state: &mut Vec<f32>,
1064 pool: Option<&Pool>,
1065) -> Vec<f32> {
1066 if state.len() != cfg.state_len() {
1067 *state = vec![0f32; cfg.state_len()];
1068 }
1069 let h = cfg.hidden_size;
1070 let mut bcx = vec![0.0f32; 3 * h];
1071 w.in_proj.matvec(x, &mut bcx, pool);
1072 let mut y = vec![0.0f32; h];
1073 short_conv_step(&bcx, &w.conv, cfg, state, &mut y);
1074 let mut out = vec![0.0f32; h];
1075 w.out_proj.matvec(&y, &mut out, pool);
1076 out
1077}
1078
1079pub fn short_conv_forward_batch(
1084 xs: &[f32],
1085 b: usize,
1086 w: &ShortConvWeights,
1087 cfg: &ShortConvCfg,
1088 state: &mut Vec<f32>,
1089 pool: Option<&Pool>,
1090) -> Vec<f32> {
1091 if state.len() != cfg.state_len() {
1092 *state = vec![0f32; cfg.state_len()];
1093 }
1094 let h = cfg.hidden_size;
1095 let mut bcx = vec![0.0f32; b * 3 * h];
1096 w.in_proj.matmat(xs, b, &mut bcx, pool);
1097 let mut y = vec![0.0f32; b * h];
1098 for bi in 0..b {
1099 short_conv_step(
1100 &bcx[bi * 3 * h..(bi + 1) * 3 * h],
1101 &w.conv,
1102 cfg,
1103 state,
1104 &mut y[bi * h..(bi + 1) * h],
1105 );
1106 }
1107 let mut out = vec![0.0f32; b * h];
1108 w.out_proj.matmat(&y, b, &mut out, pool);
1109 out
1110}
1111
1112#[allow(clippy::too_many_arguments)]
1117pub fn short_conv_pair(
1118 x1: &[f32],
1119 x2: &[f32],
1120 w: &ShortConvWeights,
1121 cfg: &ShortConvCfg,
1122 state: &mut Vec<f32>,
1123 scratch: &mut Vec<f32>,
1124 pool: Option<&Pool>,
1125) -> (Vec<f32>, Vec<f32>) {
1126 if state.len() != cfg.state_len() {
1127 *state = vec![0f32; cfg.state_len()];
1128 }
1129 let h = cfg.hidden_size;
1130 let mut bcx1 = vec![0.0f32; 3 * h];
1131 let mut bcx2 = vec![0.0f32; 3 * h];
1132 w.in_proj.matvec2(x1, x2, &mut bcx1, &mut bcx2, pool);
1133
1134 let mut y1 = vec![0.0f32; h];
1135 short_conv_step(&bcx1, &w.conv, cfg, state, &mut y1);
1136 scratch.clear();
1137 scratch.extend_from_slice(state);
1138 let mut y2 = vec![0.0f32; h];
1139 short_conv_step(&bcx2, &w.conv, cfg, scratch, &mut y2);
1140
1141 let mut out1 = vec![0.0f32; h];
1142 let mut out2 = vec![0.0f32; h];
1143 w.out_proj.matvec2(&y1, &y2, &mut out1, &mut out2, pool);
1144 (out1, out2)
1145}
1146
1147pub struct KdaWeights {
1158 pub q_proj: QTensor,
1160 pub k_proj: QTensor,
1162 pub v_proj: QTensor,
1164 pub conv_q: Vec<f32>,
1166 pub conv_k: Vec<f32>,
1167 pub conv_v: Vec<f32>,
1169 pub f_a: QTensor,
1171 pub f_b: QTensor,
1173 pub dt_bias: Vec<f32>,
1175 pub a_log: Vec<f32>,
1178 pub b_proj: QTensor,
1180 pub gate: KdaOutGate,
1182 pub o_norm: Vec<f32>,
1184 pub o_proj: QTensor,
1186 pub gate_lower_bound: Option<f32>,
1189}
1190
1191pub enum KdaOutGate {
1192 Full(QTensor),
1194 LowRank(QTensor, QTensor),
1196}
1197
1198#[derive(Clone, Copy)]
1199pub struct KdaCfg {
1200 pub num_heads: usize,
1201 pub head_k_dim: usize,
1202 pub head_v_dim: usize,
1203 pub conv_kernel: usize,
1204 pub hidden_size: usize,
1205 pub rms_eps: f64,
1206}
1207
1208impl KdaCfg {
1209 pub fn state_len(&self) -> usize {
1212 let (nh, dk, dv, kk) = (
1213 self.num_heads,
1214 self.head_k_dim,
1215 self.head_v_dim,
1216 self.conv_kernel,
1217 );
1218 (kk - 1) * (2 * nh * dk + nh * dv) + nh * dk * dv
1219 }
1220}
1221
1222fn kda_conv(raw: &[f32], taps: &[f32], ring: &mut [f32], kk: usize, out: &mut [f32]) {
1225 let c_dim = raw.len();
1226 for c in 0..c_dim {
1227 let t = &taps[c * kk..(c + 1) * kk];
1228 let mut acc = raw[c] as f64 * t[kk - 1] as f64;
1229 for j in 0..kk - 1 {
1230 acc += ring[j * c_dim + c] as f64 * t[j] as f64;
1231 }
1232 out[c] = silu(acc) as f32;
1233 }
1234 if kk > 1 {
1235 ring.copy_within(c_dim.., 0);
1236 let tail = (kk - 2) * c_dim;
1237 ring[tail..tail + c_dim].copy_from_slice(raw);
1238 }
1239}
1240
1241#[inline]
1244fn kda_log_decay(w: &KdaWeights, cfg: &KdaCfg, h: usize, d: usize, f: f32) -> f64 {
1245 let (nh, dk) = (cfg.num_heads, cfg.head_k_dim);
1246 let a = if w.a_log.len() == nh {
1247 w.a_log[h] as f64
1248 } else if w.a_log.len() == dk {
1249 w.a_log[d] as f64
1250 } else {
1251 w.a_log[h * dk + d] as f64
1252 };
1253 let raw = f as f64 + w.dt_bias[h * dk + d] as f64;
1254 match w.gate_lower_bound {
1255 Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
1256 None => -a.exp() * softplus(raw),
1257 }
1258}
1259
1260#[allow(clippy::too_many_arguments)]
1268fn kda_step(
1269 xq: &[f32],
1270 xk: &[f32],
1271 xv: &[f32],
1272 f: &[f32],
1273 b: &[f32],
1274 gate_out: &[f32],
1275 w: &KdaWeights,
1276 cfg: &KdaCfg,
1277 state: &mut [f32],
1278 of: &mut [f32],
1279 pool: Option<&Pool>,
1280) {
1281 let (nh, dk, dv, kk) = (
1282 cfg.num_heads,
1283 cfg.head_k_dim,
1284 cfg.head_v_dim,
1285 cfg.conv_kernel,
1286 );
1287 let (kd, vd) = (nh * dk, nh * dv);
1288 let ring_q_len = (kk - 1) * kd;
1289 let ring_v_len = (kk - 1) * vd;
1290 let (ring_q, rest) = state.split_at_mut(ring_q_len);
1291 let (ring_k, rest) = rest.split_at_mut(ring_q_len);
1292 let (ring_v, s_all) = rest.split_at_mut(ring_v_len);
1293
1294 let mut cq = vec![0f32; kd];
1295 let mut ck = vec![0f32; kd];
1296 let mut cv = vec![0f32; vd];
1297 kda_conv(xq, &w.conv_q, ring_q, kk, &mut cq);
1298 kda_conv(xk, &w.conv_k, ring_k, kk, &mut ck);
1299 kda_conv(xv, &w.conv_v, ring_v, kk, &mut cv);
1300
1301 let (cq, ck, cv) = (&cq, &ck, &cv);
1302 let s_ptr = SendMutF32(s_all.as_mut_ptr());
1303 let of_ptr = SendMutF32(of.as_mut_ptr());
1304 let head_range = |h0: usize, h1: usize| {
1305 let (s_ptr, of_ptr) = (s_ptr, of_ptr);
1306 let mut kv = crate::attention::take_buf(dv);
1307 let mut delta = crate::attention::take_buf(dv);
1308 let mut o = crate::attention::take_buf(dv);
1309 let mut kf = crate::attention::take_buf(dk);
1310 let mut qf = crate::attention::take_buf(dk);
1311 let mut gd = crate::attention::take_buf(dk);
1312 for h in h0..h1 {
1313 let qs = h * dk;
1314 let (mut nq, mut nkn) = (0f64, 0f64);
1316 for d in 0..dk {
1317 nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
1318 nkn += (ck[qs + d] as f64) * (ck[qs + d] as f64);
1319 }
1320 let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
1321 let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
1322 for d in 0..dk {
1323 qf[d] = cq[qs + d] * invq;
1324 kf[d] = ck[qs + d] * invk;
1325 gd[d] = kda_log_decay(w, cfg, h, d, f[qs + d]).exp() as f32;
1326 }
1327 let beta = sigmoid(b[h] as f64) as f32;
1328
1329 let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
1331 let oh = unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
1332 let vt = &cv[h * dv..(h + 1) * dv];
1333
1334 kv[..dv].fill(0.0);
1336 for di in 0..dk {
1337 let kg = kf[di] * gd[di];
1338 let row = &s[di * dv..(di + 1) * dv];
1339 for dj in 0..dv {
1340 kv[dj] += row[dj] * kg;
1341 }
1342 }
1343 for dj in 0..dv {
1344 delta[dj] = (vt[dj] - kv[dj]) * beta;
1345 }
1346 o[..dv].fill(0.0);
1348 for di in 0..dk {
1349 let (kfd, qfd, gdd) = (kf[di], qf[di], gd[di]);
1350 let row = &mut s[di * dv..(di + 1) * dv];
1351 for dj in 0..dv {
1352 let cell = gdd * row[dj] + kfd * delta[dj];
1353 row[dj] = cell;
1354 o[dj] += qfd * cell;
1355 }
1356 }
1357 let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
1359 let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
1360 for dj in 0..dv {
1361 oh[dj] = ((o[dj] as f64 * inv)
1362 * w.o_norm[dj] as f64
1363 * sigmoid(gate_out[h * dv + dj] as f64)) as f32;
1364 }
1365 }
1366 crate::attention::recycle_buf(&mut kv);
1367 crate::attention::recycle_buf(&mut delta);
1368 crate::attention::recycle_buf(&mut o);
1369 crate::attention::recycle_buf(&mut kf);
1370 crate::attention::recycle_buf(&mut qf);
1371 crate::attention::recycle_buf(&mut gd);
1372 };
1373 match pool {
1374 Some(pool) if nh >= 4 => pool.run(&|widx, n| {
1375 let chunk = nh.div_ceil(n);
1376 let h0 = (widx * chunk).min(nh);
1377 let h1 = (h0 + chunk).min(nh);
1378 if h0 < h1 {
1379 head_range(h0, h1);
1380 }
1381 }),
1382 _ => head_range(0, nh),
1383 }
1384}
1385
1386fn kda_gate_out(w: &KdaWeights, x: &[f32], vd: usize, pool: Option<&Pool>) -> Vec<f32> {
1389 let mut g = vec![0.0f32; vd];
1390 match &w.gate {
1391 KdaOutGate::Full(gp) => gp.matvec(x, &mut g, pool),
1392 KdaOutGate::LowRank(ga, gb) => {
1393 let mut low = vec![0.0f32; ga.rows()];
1394 ga.matvec(x, &mut low, pool);
1395 gb.matvec(&low, &mut g, pool);
1396 }
1397 }
1398 g
1399}
1400
1401pub fn kda_forward(
1403 x: &[f32],
1404 w: &KdaWeights,
1405 cfg: &KdaCfg,
1406 state: &mut Vec<f32>,
1407 pool: Option<&Pool>,
1408) -> Vec<f32> {
1409 if state.len() != cfg.state_len() {
1410 *state = vec![0f32; cfg.state_len()];
1411 }
1412 let (nh, dk, dv) = (cfg.num_heads, cfg.head_k_dim, cfg.head_v_dim);
1413 let (kd, vd) = (nh * dk, nh * dv);
1414
1415 let mut xq = vec![0.0f32; kd];
1416 let mut xk = vec![0.0f32; kd];
1417 let mut xv = vec![0.0f32; vd];
1418 let mut fl = vec![0.0f32; w.f_a.rows()];
1419 let mut b = vec![0.0f32; nh];
1420 let q4_group = [&w.q_proj, &w.k_proj, &w.v_proj, &w.f_a, &w.b_proj];
1425 let q4_uniform = q4_group.iter().all(|t| {
1426 matches!(
1427 t,
1428 QTensor::Mapped {
1429 dtype: TensorDtype::Q4TiledP,
1430 ..
1431 }
1432 )
1433 });
1434 if q4_uniform {
1435 QTensor::matvec_many(
1436 q4_group,
1437 x,
1438 [
1439 xq.as_mut_slice(),
1440 xk.as_mut_slice(),
1441 xv.as_mut_slice(),
1442 fl.as_mut_slice(),
1443 b.as_mut_slice(),
1444 ],
1445 pool,
1446 );
1447 } else {
1448 QTensor::matvec_many(
1449 [&w.q_proj, &w.k_proj, &w.v_proj, &w.f_a],
1450 x,
1451 [
1452 xq.as_mut_slice(),
1453 xk.as_mut_slice(),
1454 xv.as_mut_slice(),
1455 fl.as_mut_slice(),
1456 ],
1457 pool,
1458 );
1459 w.b_proj.matvec(x, &mut b, pool);
1460 }
1461 let mut f = vec![0.0f32; kd];
1462 w.f_b.matvec(&fl, &mut f, pool);
1463 let gate_out = kda_gate_out(w, x, vd, pool);
1464
1465 let mut of = vec![0.0f32; vd];
1466 kda_step(
1467 &xq, &xk, &xv, &f, &b, &gate_out, w, cfg, state, &mut of, pool,
1468 );
1469
1470 let mut out = vec![0.0f32; cfg.hidden_size];
1471 w.o_proj.matvec(&of, &mut out, pool);
1472 out
1473}
1474
1475pub fn kda_forward_batch(
1479 xs: &[f32],
1480 bsz: usize,
1481 w: &KdaWeights,
1482 cfg: &KdaCfg,
1483 state: &mut Vec<f32>,
1484 pool: Option<&Pool>,
1485) -> Vec<f32> {
1486 if state.len() != cfg.state_len() {
1487 *state = vec![0f32; cfg.state_len()];
1488 }
1489 let (nh, dk, dv, hs) = (
1490 cfg.num_heads,
1491 cfg.head_k_dim,
1492 cfg.head_v_dim,
1493 cfg.hidden_size,
1494 );
1495 let (kd, vd) = (nh * dk, nh * dv);
1496
1497 let mut xq = vec![0.0f32; bsz * kd];
1498 w.q_proj.matmat(xs, bsz, &mut xq, pool);
1499 let mut xk = vec![0.0f32; bsz * kd];
1500 w.k_proj.matmat(xs, bsz, &mut xk, pool);
1501 let mut xv = vec![0.0f32; bsz * vd];
1502 w.v_proj.matmat(xs, bsz, &mut xv, pool);
1503 let rank = w.f_a.rows();
1504 let mut fl = vec![0.0f32; bsz * rank];
1505 w.f_a.matmat(xs, bsz, &mut fl, pool);
1506 let mut f = vec![0.0f32; bsz * kd];
1507 w.f_b.matmat(&fl, bsz, &mut f, pool);
1508 let mut b = vec![0.0f32; bsz * nh];
1509 w.b_proj.matmat(xs, bsz, &mut b, pool);
1510 let mut gate_out = vec![0.0f32; bsz * vd];
1511 match &w.gate {
1512 KdaOutGate::Full(gp) => gp.matmat(xs, bsz, &mut gate_out, pool),
1513 KdaOutGate::LowRank(ga, gb) => {
1514 let mut low = vec![0.0f32; bsz * ga.rows()];
1515 ga.matmat(xs, bsz, &mut low, pool);
1516 gb.matmat(&low, bsz, &mut gate_out, pool);
1517 }
1518 }
1519
1520 let mut of = vec![0.0f32; bsz * vd];
1521 for bi in 0..bsz {
1522 let mut oh = vec![0.0f32; vd];
1523 kda_step(
1524 &xq[bi * kd..(bi + 1) * kd],
1525 &xk[bi * kd..(bi + 1) * kd],
1526 &xv[bi * vd..(bi + 1) * vd],
1527 &f[bi * kd..(bi + 1) * kd],
1528 &b[bi * nh..(bi + 1) * nh],
1529 &gate_out[bi * vd..(bi + 1) * vd],
1530 w,
1531 cfg,
1532 state,
1533 &mut oh,
1534 pool,
1535 );
1536 of[bi * vd..(bi + 1) * vd].copy_from_slice(&oh);
1537 }
1538
1539 let mut out = vec![0.0f32; bsz * hs];
1540 w.o_proj.matmat(&of, bsz, &mut out, pool);
1541 out
1542}
1543
1544#[cfg(test)]
1545mod tests {
1546 #[test]
1547 fn kda_forward_matches_naive_reference() {
1548 let (nh, dk, dv, kk, hs, rank) = (2usize, 4usize, 4usize, 3usize, 6usize, 3usize);
1554 let synth = |rows: usize, cols: usize, salt: usize| -> QTensor {
1555 QTensor::from_f32(
1556 (0..rows * cols)
1557 .map(|i| (((i * 31 + salt * 17) % 101) as f32 / 101.0 - 0.5) * 0.6)
1558 .collect(),
1559 rows,
1560 cols,
1561 )
1562 };
1563 let vecf = |n: usize, salt: usize| -> Vec<f32> {
1564 (0..n)
1565 .map(|i| (((i * 13 + salt * 7) % 89) as f32 / 89.0 - 0.5) * 0.8)
1566 .collect()
1567 };
1568 for (label, a_log, lb) in [
1569 ("per-head standard", vecf(nh, 40), None),
1570 ("per-dim lower-bound", vecf(dk, 41), Some(-5.0f32)),
1571 ] {
1572 let w = KdaWeights {
1573 q_proj: synth(nh * dk, hs, 1),
1574 k_proj: synth(nh * dk, hs, 2),
1575 v_proj: synth(nh * dv, hs, 3),
1576 conv_q: vecf(nh * dk * kk, 4),
1577 conv_k: vecf(nh * dk * kk, 5),
1578 conv_v: vecf(nh * dv * kk, 6),
1579 f_a: synth(rank, hs, 7),
1580 f_b: synth(nh * dk, rank, 8),
1581 dt_bias: vecf(nh * dk, 9),
1582 a_log: a_log.clone(),
1583 b_proj: synth(nh, hs, 10),
1584 gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
1585 o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
1586 o_proj: synth(hs, nh * dv, 13),
1587 gate_lower_bound: lb,
1588 };
1589 let cfg = KdaCfg {
1590 num_heads: nh,
1591 head_k_dim: dk,
1592 head_v_dim: dv,
1593 conv_kernel: kk,
1594 hidden_size: hs,
1595 rms_eps: 1e-6,
1596 };
1597 let xs: Vec<Vec<f32>> = (0..6)
1598 .map(|t| {
1599 (0..hs)
1600 .map(|i| ((t * hs + i) as f32 * 0.37).sin() * 0.5)
1601 .collect()
1602 })
1603 .collect();
1604
1605 let mut state = Vec::new();
1607 let got: Vec<Vec<f32>> = xs
1608 .iter()
1609 .map(|x| kda_forward(x, &w, &cfg, &mut state, None))
1610 .collect();
1611
1612 let mv = |t: &QTensor, x: &[f32]| -> Vec<f32> {
1614 let mut o = vec![0.0f32; t.rows()];
1615 t.matvec(x, &mut o, None);
1616 o
1617 };
1618 let mut hist: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = Vec::new(); let mut s_state = vec![0f64; nh * dk * dv];
1620 let mut want: Vec<Vec<f32>> = Vec::new();
1621 for x in &xs {
1622 let (xq, xk, xv) = (mv(&w.q_proj, x), mv(&w.k_proj, x), mv(&w.v_proj, x));
1623 hist.push((xq, xk, xv));
1624 let conv = |sel: fn(&(Vec<f32>, Vec<f32>, Vec<f32>)) -> &Vec<f32>,
1626 taps: &[f32],
1627 n: usize|
1628 -> Vec<f32> {
1629 (0..n)
1630 .map(|c| {
1631 let t = &taps[c * kk..(c + 1) * kk];
1632 let mut acc = 0f64;
1633 for j in 0..kk {
1634 let idx = hist.len() as i64 - (kk as i64 - j as i64);
1635 if idx >= 0 {
1636 acc += sel(&hist[idx as usize])[c] as f64 * t[j] as f64;
1637 }
1638 }
1639 silu(acc)
1640 })
1641 .map(|v| v as f32)
1642 .collect()
1643 };
1644 let cq = conv(|h| &h.0, &w.conv_q, nh * dk);
1645 let ck = conv(|h| &h.1, &w.conv_k, nh * dk);
1646 let cv = conv(|h| &h.2, &w.conv_v, nh * dv);
1647 let f = mv(&w.f_b, &mv(&w.f_a, x));
1648 let bb = mv(&w.b_proj, x);
1649 let gate_out = match &w.gate {
1650 KdaOutGate::LowRank(ga, gb) => mv(gb, &mv(ga, x)),
1651 KdaOutGate::Full(g) => mv(g, x),
1652 };
1653 let mut of = vec![0f32; nh * dv];
1654 for h in 0..nh {
1655 let q: Vec<f64> = {
1657 let sl = &cq[h * dk..(h + 1) * dk];
1658 let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
1659 let inv = 1.0 / ((n + 1e-6).sqrt() * (dk as f64).sqrt());
1660 sl.iter().map(|&v| v as f64 * inv).collect()
1661 };
1662 let k: Vec<f64> = {
1663 let sl = &ck[h * dk..(h + 1) * dk];
1664 let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
1665 let inv = 1.0 / (n + 1e-6).sqrt();
1666 sl.iter().map(|&v| v as f64 * inv).collect()
1667 };
1668 let v: Vec<f64> = cv[h * dv..(h + 1) * dv].iter().map(|&v| v as f64).collect();
1669 let g: Vec<f64> = (0..dk)
1671 .map(|d| {
1672 let a = if w.a_log.len() == nh {
1673 w.a_log[h] as f64
1674 } else {
1675 w.a_log[d] as f64
1676 };
1677 let raw = f[h * dk + d] as f64 + w.dt_bias[h * dk + d] as f64;
1678 match w.gate_lower_bound {
1679 Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
1680 None => -a.exp() * softplus(raw),
1681 }
1682 })
1683 .collect();
1684 let beta = sigmoid(bb[h] as f64);
1685 let s = &mut s_state[h * dk * dv..(h + 1) * dk * dv];
1686 for di in 0..dk {
1688 for dj in 0..dv {
1689 s[di * dv + dj] *= g[di].exp();
1690 }
1691 }
1692 let mut kv = vec![0f64; dv];
1694 for di in 0..dk {
1695 for dj in 0..dv {
1696 kv[dj] += k[di] * s[di * dv + dj];
1697 }
1698 }
1699 for di in 0..dk {
1700 for dj in 0..dv {
1701 s[di * dv + dj] += beta * k[di] * (v[dj] - kv[dj]);
1702 }
1703 }
1704 let mut o = vec![0f64; dv];
1705 for di in 0..dk {
1706 for dj in 0..dv {
1707 o[dj] += q[di] * s[di * dv + dj];
1708 }
1709 }
1710 let ss: f64 = o.iter().map(|&v| v * v).sum();
1712 let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
1713 for dj in 0..dv {
1714 of[h * dv + dj] = (o[dj]
1715 * inv
1716 * w.o_norm[dj] as f64
1717 * sigmoid(gate_out[h * dv + dj] as f64))
1718 as f32;
1719 }
1720 }
1721 want.push(mv(&w.o_proj, &of));
1722 }
1723
1724 for (t, (g, e)) in got.iter().zip(&want).enumerate() {
1725 for (i, (a, b)) in g.iter().zip(e.iter()).enumerate() {
1726 assert!((a - b).abs() < 2e-4, "{label}: t={t} i={i}: {a} vs {b}");
1727 }
1728 }
1729 }
1730
1731 let w = KdaWeights {
1733 q_proj: synth(nh * dk, hs, 1),
1734 k_proj: synth(nh * dk, hs, 2),
1735 v_proj: synth(nh * dv, hs, 3),
1736 conv_q: vecf(nh * dk * kk, 4),
1737 conv_k: vecf(nh * dk * kk, 5),
1738 conv_v: vecf(nh * dv * kk, 6),
1739 f_a: synth(rank, hs, 7),
1740 f_b: synth(nh * dk, rank, 8),
1741 dt_bias: vecf(nh * dk, 9),
1742 a_log: vecf(nh, 40),
1743 b_proj: synth(nh, hs, 10),
1744 gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
1745 o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
1746 o_proj: synth(hs, nh * dv, 13),
1747 gate_lower_bound: None,
1748 };
1749 let cfg = KdaCfg {
1750 num_heads: nh,
1751 head_k_dim: dk,
1752 head_v_dim: dv,
1753 conv_kernel: kk,
1754 hidden_size: hs,
1755 rms_eps: 1e-6,
1756 };
1757 let xs: Vec<f32> = (0..5 * hs).map(|i| (i as f32 * 0.29).cos() * 0.4).collect();
1758 let mut st1 = Vec::new();
1759 let seq: Vec<f32> = (0..5)
1760 .flat_map(|t| kda_forward(&xs[t * hs..(t + 1) * hs], &w, &cfg, &mut st1, None))
1761 .collect();
1762 let mut st2 = Vec::new();
1763 let bat = kda_forward_batch(&xs, 5, &w, &cfg, &mut st2, None);
1764 for (i, (a, b)) in seq.iter().zip(&bat).enumerate() {
1765 assert!((a - b).abs() < 1e-5, "batch i={i}: {a} vs {b}");
1766 }
1767 assert_eq!(st1, st2, "state must match after the chunk");
1768 }
1769
1770 use super::*;
1771
1772 fn tiny() -> (VmfPhaseWeights, VmfPhaseCfg) {
1773 let cfg = VmfPhaseCfg {
1774 num_heads: 2,
1775 nphase: 3,
1776 value_head_dim: 4,
1777 hidden_size: 8,
1778 phase_mass: 0.0,
1779 };
1780 let synth = |rows: usize, cols: usize, salt: usize| {
1781 QTensor::from_f32(
1782 (0..rows * cols)
1783 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
1784 .collect(),
1785 rows,
1786 cols,
1787 )
1788 };
1789 let w = VmfPhaseWeights {
1790 thq: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 1),
1791 thk: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 2),
1792 v_proj: synth(cfg.num_heads * cfg.value_head_dim, cfg.hidden_size, 3),
1793 out_proj: synth(cfg.hidden_size, cfg.num_heads * cfg.value_head_dim, 4),
1794 decay: (0..cfg.num_heads * 2 * cfg.nphase)
1795 .map(|i| 0.9 + 0.005 * (i % 10) as f64)
1796 .collect(),
1797 conv: None,
1798 k_gate: None,
1799 phase_delta: false,
1800 };
1801 (w, cfg)
1802 }
1803
1804 #[test]
1805 fn state_persists_and_changes_output() {
1806 let (w, cfg) = tiny();
1807 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
1808 let mut state = Vec::new();
1809 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1810 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1811 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
1813 assert_eq!(state.len(), cfg.state_len());
1814 }
1815
1816 #[test]
1817 fn phase_delta_matches_independent_token_oracle_and_ignores_phase_mass() {
1818 let cfg = VmfPhaseCfg {
1823 num_heads: 1,
1824 nphase: 2,
1825 value_head_dim: 1,
1826 hidden_size: 2,
1827 phase_mass: 17.0,
1829 };
1830 let w = VmfPhaseWeights {
1831 thq: QTensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], 2, 2),
1832 thk: QTensor::from_f32(vec![0.0, 1.0, 1.0, 0.0], 2, 2),
1833 v_proj: QTensor::from_f32(vec![1.0, 0.0], 1, 2),
1834 out_proj: QTensor::from_f32(vec![1.0, 0.0], 2, 1),
1835 decay: vec![0.9, 0.8, 0.7, 0.6],
1836 conv: None,
1837 k_gate: None, phase_delta: true,
1839 };
1840 let xs = [[0.3f32, 0.7f32], [-0.2, 0.4], [0.8, -0.6]];
1841 let c = 1.0f64 / 2.0f64.sqrt();
1842 let mut oracle_state = [0.0f64; 4];
1843 let mut state = Vec::new();
1844 for x in xs {
1845 let q = [
1846 c * (x[0] as f64).cos(),
1847 c * (x[1] as f64).cos(),
1848 c * (x[0] as f64).sin(),
1849 c * (x[1] as f64).sin(),
1850 ];
1851 let k = [
1852 c * (x[1] as f64).cos(),
1853 c * (x[0] as f64).cos(),
1854 c * (x[1] as f64).sin(),
1855 c * (x[0] as f64).sin(),
1856 ];
1857 let value = x[0] as f64;
1858 let mut read = 0.0;
1859 for f in 0..4 {
1860 oracle_state[f] *= w.decay[f];
1861 read += k[f] * oracle_state[f];
1862 }
1863 for f in 0..4 {
1864 oracle_state[f] += k[f] * (value - read);
1865 }
1866 let want: f32 = q
1867 .iter()
1868 .zip(oracle_state)
1869 .map(|(qf, sf)| (qf * sf) as f32)
1870 .sum();
1871 let got = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1872 assert!(
1873 (got[0] - want).abs() < 2e-6,
1874 "oracle {want} vs runtime {}",
1875 got[0]
1876 );
1877 assert!(got[1].abs() < 2e-7);
1878 }
1879 assert_eq!(state.len(), cfg.state_len());
1880 for (got, want) in state.iter().zip(oracle_state) {
1881 assert!((*got as f64 - want).abs() < 2e-6);
1882 }
1883 }
1884
1885 #[test]
1886 fn phase_delta_fused_readout_is_bit_exact_over_long_continuation() {
1887 for (nh, nph, dv) in [(1, 1, 1), (3, 5, 7), (8, 32, 128)] {
1890 let cfg = VmfPhaseCfg {
1891 num_heads: nh,
1892 nphase: nph,
1893 value_head_dim: dv,
1894 hidden_size: nh * dv,
1895 phase_mass: 19.0,
1896 };
1897 let p2 = 2 * nph;
1898 let scale = 1.0f64 / (nph as f64).sqrt();
1899 let mut state: Vec<f32> = (0..cfg.state_len() + 11)
1900 .map(|i| (i as f32 * 0.37).sin())
1901 .collect();
1902 let mut reference = state.clone();
1903 let tail = state[cfg.state_len()..].to_vec();
1904 let decay: Vec<f64> = (0..nh * p2).map(|i| [0.0, 0.8, 0.99, 1.0][i % 4]).collect();
1905 for step in 0..128 {
1906 let q: Vec<f32> = (0..nh * nph)
1907 .map(|i| ((i + step * 3) as f32 * 0.13).sin() * 4.0)
1908 .collect();
1909 let k: Vec<f32> = (0..nh * nph)
1910 .map(|i| ((i + step * 7) as f32 * 0.29).cos() * 4.0)
1911 .collect();
1912 let v: Vec<f32> = (0..nh * dv)
1913 .map(|i| ((i + step) as f32 * 0.43).sin())
1914 .collect();
1915 let gates: Vec<f32> = (0..nh).map(|h| [0.0, 0.3, 1.0][(step + h) % 3]).collect();
1916 let kap = (step % 4 != 0).then_some(gates.as_slice());
1917 let mut out = vec![0.125f32; nh * dv];
1918 let mut want = out.clone();
1919 for h in 0..nh {
1920 let feature = |theta: &[f32], f: usize| {
1921 let angle = theta[h * nph + f % nph] as f64;
1922 scale * if f < nph { angle.cos() } else { angle.sin() }
1923 };
1924 let beta = kap.map_or(1.0f64, |g| g[h] as f64);
1925 let mut read = vec![0.0f64; dv];
1926 for f in 0..p2 {
1927 for d in 0..dv {
1928 let at = (h * p2 + f) * dv + d;
1929 read[d] += feature(&k, f) * (decay[h * p2 + f] * reference[at] as f64);
1930 }
1931 }
1932 for f in 0..p2 {
1933 for d in 0..dv {
1934 let at = (h * p2 + f) * dv + d;
1935 reference[at] = (decay[h * p2 + f] * reference[at] as f64
1936 + beta * feature(&k, f) * (v[h * dv + d] as f64 - read[d]))
1937 as f32;
1938 }
1939 }
1940 for f in 0..p2 {
1941 for d in 0..dv {
1942 want[h * dv + d] +=
1943 (feature(&q, f) * reference[(h * p2 + f) * dv + d] as f64) as f32;
1944 }
1945 }
1946 }
1947 phase_delta_step_f32(&q, &k, &v, &decay, kap, &cfg, &mut state, &mut out);
1948 assert_eq!(out, want, "shape {nh}/{nph}/{dv}, step {step}");
1949 assert_eq!(state, reference, "state at step {step}");
1950 assert_eq!(&state[cfg.state_len()..], tail);
1951 }
1952 }
1953 }
1954
1955 #[test]
1956 fn phase_delta_pair_reset_and_legacy_paths_are_distinct() {
1957 let (mut delta, mut cfg) = tiny();
1958 delta.phase_delta = true;
1959 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
1962 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
1963 let mut seq_state = Vec::new();
1964 let d1 = vmf_phase_forward(&x1, &delta, &cfg, &mut seq_state, None);
1965 let d2 = vmf_phase_forward(&x2, &delta, &cfg, &mut seq_state, None);
1966
1967 let mut pair_state = Vec::new();
1968 let mut scratch = Vec::new();
1969 let (p1, p2) = vmf_phase_pair(&x1, &x2, &delta, &cfg, &mut pair_state, &mut scratch, None);
1970 assert_eq!(d1, p1);
1971 assert_eq!(d2, p2);
1972 std::mem::swap(&mut pair_state, &mut scratch);
1973 assert_eq!(seq_state, pair_state);
1974
1975 let mut reset = seq_state;
1978 reset.clear();
1979 let after_reset = vmf_phase_forward(&x1, &delta, &cfg, &mut reset, None);
1980 assert_eq!(after_reset, d1);
1981
1982 cfg.phase_mass = 0.0;
1983 let mut legacy = delta;
1984 legacy.phase_delta = false;
1985 let mut legacy_state = Vec::new();
1986 let old = vmf_phase_forward(&x1, &legacy, &cfg, &mut legacy_state, None);
1987 assert!(
1988 old.iter().zip(&d1).any(|(a, b)| (a - b).abs() > 1e-5),
1989 "legacy additive and normalized Phase-Delta paths must differ"
1990 );
1991 }
1992
1993 #[test]
1997 fn phase_mass_zero_is_noop_and_positive_shifts() {
1998 let (w, cfg0) = tiny();
1999 let mut cfg_m = cfg0;
2000 cfg_m.phase_mass = 1.0;
2001 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.4).sin()).collect();
2002
2003 let mut s0 = Vec::new();
2004 let base = vmf_phase_forward(&x, &w, &cfg0, &mut s0, None);
2005 let mut s0b = Vec::new();
2007 let base2 = vmf_phase_forward(&x, &w, &cfg0, &mut s0b, None);
2008 assert_eq!(base, base2, "mass=0 must be deterministic/no-op");
2009 let mut sm = Vec::new();
2011 let massed = vmf_phase_forward(&x, &w, &cfg_m, &mut sm, None);
2012 assert!(
2013 base.iter().zip(&massed).any(|(a, b)| (a - b).abs() > 1e-5),
2014 "mass>0 must change the output"
2015 );
2016 assert!(massed.iter().all(|v| v.is_finite()));
2017 }
2018
2019 #[test]
2024 fn kappa_gate_open_matches_none_and_closed_writes_nothing() {
2025 let (mut w, cfg) = tiny();
2026 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
2027
2028 let mut s_none = Vec::new();
2029 let base1 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
2030 let base2 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
2031
2032 w.k_gate = Some((
2034 QTensor::from_f32(
2035 vec![0.0; cfg.num_heads * cfg.hidden_size],
2036 cfg.num_heads,
2037 cfg.hidden_size,
2038 ),
2039 vec![20.0; cfg.num_heads],
2040 ));
2041 let mut s_open = Vec::new();
2042 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
2043 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
2044 for (a, b) in base1.iter().zip(&o1).chain(base2.iter().zip(&o2)) {
2045 assert!(
2046 (a - b).abs() < 1e-5,
2047 "open κ must match gateless: {a} vs {b}"
2048 );
2049 }
2050
2051 w.k_gate = Some((
2053 QTensor::from_f32(
2054 vec![0.0; cfg.num_heads * cfg.hidden_size],
2055 cfg.num_heads,
2056 cfg.hidden_size,
2057 ),
2058 vec![-20.0; cfg.num_heads],
2059 ));
2060 let mut s_closed = Vec::new();
2061 let oc = vmf_phase_forward(&x, &w, &cfg, &mut s_closed, None);
2062 assert!(
2063 s_closed.iter().all(|&v| v.abs() < 1e-7),
2064 "closed κ: state must stay empty"
2065 );
2066 assert!(
2067 oc.iter().all(|&v| v.abs() < 1e-6),
2068 "closed κ: empty-state readout"
2069 );
2070 }
2071
2072 #[test]
2073 fn pair_matches_two_singles_bitexact() {
2074 let (w, cfg) = tiny();
2075 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
2076 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
2077
2078 let mut s_ref = Vec::new();
2080 let r1 = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
2081 let r2 = vmf_phase_forward(&x2, &w, &cfg, &mut s_ref, None);
2082
2083 let mut s = Vec::new();
2085 let mut scratch = Vec::new();
2086 let (p1, p2) = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2087 assert_eq!(r1, p1, "lane 1 must be bit-identical");
2088 assert_eq!(r2, p2, "lane 2 must be bit-identical");
2089 std::mem::swap(&mut s, &mut scratch);
2091 assert_eq!(s, s_ref, "accepted state must equal sequential state");
2092 }
2093
2094 #[test]
2095 fn rejected_draft_leaves_state_at_lane1() {
2096 let (w, cfg) = tiny();
2097 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
2098 let x2 = vec![0.5f32; 8];
2099
2100 let mut s_ref = Vec::new();
2101 let _ = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
2102
2103 let mut s = Vec::new();
2104 let mut scratch = Vec::new();
2105 let _ = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2106 assert_eq!(s, s_ref);
2108 }
2109
2110 fn tiny_gdn() -> (GdnWeights, GdnCfg) {
2113 let cfg = GdnCfg {
2114 num_v_heads: 4,
2115 num_k_heads: 2,
2116 key_head_dim: 3,
2117 value_head_dim: 5,
2118 conv_kernel: 4,
2119 hidden_size: 8,
2120 rms_eps: 1e-6,
2121 output_gate_sigmoid: false,
2122 };
2123 let c_dim = cfg.conv_dim();
2124 let vd = cfg.num_v_heads * cfg.value_head_dim;
2125 let synth = |rows: usize, cols: usize, salt: usize| {
2126 QTensor::from_f32(
2127 (0..rows * cols)
2128 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
2129 .collect(),
2130 rows,
2131 cols,
2132 )
2133 };
2134 let vecf = |n: usize, salt: usize| -> Vec<f32> {
2135 (0..n)
2136 .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.6)
2137 .collect()
2138 };
2139 let w = GdnWeights {
2140 in_proj_qkv: synth(c_dim, cfg.hidden_size, 1),
2141 in_proj_z: synth(vd, cfg.hidden_size, 2),
2142 in_proj_a: synth(cfg.num_v_heads, cfg.hidden_size, 3),
2143 in_proj_b: synth(cfg.num_v_heads, cfg.hidden_size, 4),
2144 conv1d: vecf(c_dim * cfg.conv_kernel, 5),
2145 a_log: (0..cfg.num_v_heads).map(|i| 0.2 + 0.3 * i as f32).collect(),
2146 dt_bias: vecf(cfg.num_v_heads, 6),
2147 norm: vec![1.0; cfg.value_head_dim],
2148 out_proj: synth(cfg.hidden_size, vd, 7),
2149 };
2150 (w, cfg)
2151 }
2152
2153 #[test]
2154 fn gdn_state_persists_and_changes_output() {
2155 let (w, cfg) = tiny_gdn();
2156 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
2157 let mut state = Vec::new();
2158 let o1 = gdn_forward(&x, &w, &cfg, &mut state, None);
2159 let o2 = gdn_forward(&x, &w, &cfg, &mut state, None);
2160 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
2161 assert_eq!(state.len(), cfg.state_len());
2162 }
2163
2164 #[test]
2165 fn gdn_pair_matches_two_singles_bitexact() {
2166 let (w, cfg) = tiny_gdn();
2167 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
2168 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
2169
2170 let mut s_ref = Vec::new();
2171 let r1 = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
2172 let r2 = gdn_forward(&x2, &w, &cfg, &mut s_ref, None);
2173
2174 let mut s = Vec::new();
2175 let mut scratch = Vec::new();
2176 let (p1, p2) = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2177 assert_eq!(r1, p1, "lane 1 must be bit-identical");
2178 assert_eq!(r2, p2, "lane 2 must be bit-identical");
2179 std::mem::swap(&mut s, &mut scratch);
2180 assert_eq!(s, s_ref, "accepted state must equal sequential state");
2181 }
2182
2183 #[test]
2184 fn gdn_rejected_draft_leaves_state_at_lane1() {
2185 let (w, cfg) = tiny_gdn();
2186 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
2187 let x2 = vec![0.5f32; 8];
2188
2189 let mut s_ref = Vec::new();
2190 let _ = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
2191
2192 let mut s = Vec::new();
2193 let mut scratch = Vec::new();
2194 let _ = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2195 assert_eq!(s, s_ref);
2196 }
2197
2198 #[test]
2202 fn gdn_conv_ring_matches_explicit_causal_conv() {
2203 let (w, cfg) = tiny_gdn();
2204 let seq: Vec<Vec<f32>> = (0..6)
2205 .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.17).sin()).collect())
2206 .collect();
2207
2208 let mut s_inc = Vec::new();
2211 for (t, x) in seq.iter().enumerate() {
2212 let inc = gdn_forward(x, &w, &cfg, &mut s_inc, None);
2213 let mut s_replay = Vec::new();
2214 let mut replay = Vec::new();
2215 for xr in &seq[..=t] {
2216 replay = gdn_forward(xr, &w, &cfg, &mut s_replay, None);
2217 }
2218 assert_eq!(inc, replay, "position {t}: ring must equal replay");
2219 }
2220 }
2221
2222 fn tiny_short_conv() -> (ShortConvWeights, ShortConvCfg) {
2223 let cfg = ShortConvCfg {
2224 hidden_size: 8,
2225 kernel: 3,
2226 };
2227 let synth = |rows: usize, cols: usize, salt: usize| {
2228 QTensor::from_f32(
2229 (0..rows * cols)
2230 .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.5)
2231 .collect(),
2232 rows,
2233 cols,
2234 )
2235 };
2236 let w = ShortConvWeights {
2237 in_proj: synth(3 * cfg.hidden_size, cfg.hidden_size, 1),
2238 conv: (0..cfg.hidden_size * cfg.kernel)
2239 .map(|i| ((i * 7 % 13) as f32 / 13.0 - 0.5) * 0.8)
2240 .collect(),
2241 out_proj: synth(cfg.hidden_size, cfg.hidden_size, 2),
2242 };
2243 (w, cfg)
2244 }
2245
2246 #[test]
2249 fn short_conv_ring_matches_explicit_causal_conv() {
2250 let (w, cfg) = tiny_short_conv();
2251 let seq: Vec<Vec<f32>> = (0..6)
2252 .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.19).cos()).collect())
2253 .collect();
2254 let mut s_inc = Vec::new();
2255 for (t, x) in seq.iter().enumerate() {
2256 let inc = short_conv_forward(x, &w, &cfg, &mut s_inc, None);
2257 let mut s_replay = Vec::new();
2258 let mut replay = Vec::new();
2259 for xr in &seq[..=t] {
2260 replay = short_conv_forward(xr, &w, &cfg, &mut s_replay, None);
2261 }
2262 assert_eq!(inc, replay, "position {t}: ring must equal replay");
2263 assert_eq!(s_inc.len(), cfg.state_len());
2264 }
2265 }
2266
2267 #[test]
2270 fn short_conv_batch_matches_sequential() {
2271 let (w, cfg) = tiny_short_conv();
2272 let b = 5;
2273 let xs: Vec<f32> = (0..b * cfg.hidden_size)
2274 .map(|i| (i as f32 * 0.13).sin() * 0.6)
2275 .collect();
2276
2277 let mut s_seq = Vec::new();
2278 let mut seq_out = vec![0.0f32; b * cfg.hidden_size];
2279 for bi in 0..b {
2280 let o = short_conv_forward(
2281 &xs[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size],
2282 &w,
2283 &cfg,
2284 &mut s_seq,
2285 None,
2286 );
2287 seq_out[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size].copy_from_slice(&o);
2288 }
2289
2290 let mut s_batch = Vec::new();
2291 let batch_out = short_conv_forward_batch(&xs, b, &w, &cfg, &mut s_batch, None);
2292 assert_eq!(
2293 seq_out, batch_out,
2294 "batch conv must match sequential decode"
2295 );
2296 assert_eq!(s_seq, s_batch, "ring state must match after the chunk");
2297 }
2298}