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(
77 thq: &[f32],
78 thk: &[f32],
79 v: &[f32],
80 decay: &[f64],
81 kap: Option<&[f32]>,
82 cfg: &VmfPhaseCfg,
83 state: &mut [f32],
84 out: &mut [f32],
85) {
86 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
87 let mscale = 1.0f64 / (1.0 + cfg.phase_mass as f64);
89 let p2 = 2 * nph;
90 for h in 0..nh {
91 let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
92 let thk_h = &thk[h * nph..(h + 1) * nph];
93 let thq_h = &thq[h * nph..(h + 1) * nph];
94 let vt = &v[h * dv..(h + 1) * dv];
95 let ot = &mut out[h * dv..(h + 1) * dv];
96 let dec = &decay[h * p2..(h + 1) * p2];
97 let kh = kap.map_or(1.0f64, |k| k[h] as f64);
99 for f in 0..p2 {
100 let (fk, fq) = if f < nph {
102 ((thk_h[f] as f64 * mscale).cos(), (thq_h[f] as f64 * mscale).cos())
103 } else {
104 (
105 (thk_h[f - nph] as f64 * mscale).sin(),
106 (thq_h[f - nph] as f64 * mscale).sin(),
107 )
108 };
109 let fkw = fk * kh;
110 let row = &mut s[f * dv..(f + 1) * dv];
111 let dcf = dec[f];
112 for d in 0..dv {
113 let cell = dcf * row[d] as f64 + fkw * vt[d] as f64;
115 row[d] = cell as f32;
116 ot[d] += (fq * cell) as f32; }
118 }
119 }
120}
121
122fn kappa_of(x: &[f32], w: &VmfPhaseWeights, nh: usize, pool: Option<&Pool>) -> Option<Vec<f32>> {
125 let (kw, kb) = w.k_gate.as_ref()?;
126 let mut k = vec![0.0f32; nh];
127 kw.matvec(x, &mut k, pool);
128 for (v, b) in k.iter_mut().zip(kb) {
129 *v = 1.0 / (1.0 + (-(*v + b)).exp());
130 }
131 Some(k)
132}
133
134pub fn vmf_phase_forward(
136 x: &[f32],
137 w: &VmfPhaseWeights,
138 cfg: &VmfPhaseCfg,
139 state: &mut Vec<f32>,
140 pool: Option<&Pool>,
141) -> Vec<f32> {
142 if state.len() != cfg.state_len() {
143 *state = vec![0f32; cfg.state_len()];
144 }
145 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
146
147 let mut thq = vec![0.0f32; nh * nph];
148 w.thq.matvec(x, &mut thq, pool);
149 let mut thk = vec![0.0f32; nh * nph];
150 w.thk.matvec(x, &mut thk, pool);
151 let mut v = vec![0.0f32; nh * dv];
152 w.v_proj.matvec(x, &mut v, pool);
153
154 let kap = kappa_of(x, w, nh, pool);
155 let mut o = vec![0.0f32; nh * dv];
156 phase_step(&thq, &thk, &v, &w.decay, kap.as_deref(), cfg, state, &mut o);
157
158 let mut out = vec![0.0f32; cfg.hidden_size];
159 w.out_proj.matvec(&o, &mut out, pool);
160 out
161}
162
163#[allow(clippy::too_many_arguments)]
168pub fn vmf_phase_pair(
169 x1: &[f32],
170 x2: &[f32],
171 w: &VmfPhaseWeights,
172 cfg: &VmfPhaseCfg,
173 state: &mut Vec<f32>,
174 scratch: &mut Vec<f32>,
175 pool: Option<&Pool>,
176) -> (Vec<f32>, Vec<f32>) {
177 if state.len() != cfg.state_len() {
178 *state = vec![0f32; cfg.state_len()];
179 }
180 let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
181
182 let mut thq1 = vec![0.0f32; nh * nph];
183 let mut thq2 = vec![0.0f32; nh * nph];
184 w.thq.matvec2(x1, x2, &mut thq1, &mut thq2, pool);
185 let mut thk1 = vec![0.0f32; nh * nph];
186 let mut thk2 = vec![0.0f32; nh * nph];
187 w.thk.matvec2(x1, x2, &mut thk1, &mut thk2, pool);
188 let mut v1 = vec![0.0f32; nh * dv];
189 let mut v2 = vec![0.0f32; nh * dv];
190 w.v_proj.matvec2(x1, x2, &mut v1, &mut v2, pool);
191
192 let kap1 = kappa_of(x1, w, nh, pool);
194 let mut o1 = vec![0.0f32; nh * dv];
195 phase_step(&thq1, &thk1, &v1, &w.decay, kap1.as_deref(), cfg, state, &mut o1);
196
197 let kap2 = kappa_of(x2, w, nh, pool);
199 scratch.clear();
200 scratch.extend_from_slice(state);
201 let mut o2 = vec![0.0f32; nh * dv];
202 phase_step(&thq2, &thk2, &v2, &w.decay, kap2.as_deref(), cfg, scratch, &mut o2);
203
204 let mut out1 = vec![0.0f32; cfg.hidden_size];
205 let mut out2 = vec![0.0f32; cfg.hidden_size];
206 w.out_proj.matvec2(&o1, &o2, &mut out1, &mut out2, pool);
207 (out1, out2)
208}
209
210pub struct GdnWeights {
215 pub in_proj_qkv: QTensor,
217 pub in_proj_z: QTensor,
219 pub in_proj_a: QTensor,
221 pub in_proj_b: QTensor,
223 pub conv1d: Vec<f32>,
225 pub a_log: Vec<f32>,
227 pub dt_bias: Vec<f32>,
229 pub norm: Vec<f32>,
231 pub out_proj: QTensor,
233}
234
235#[derive(Clone, Copy)]
236pub struct GdnCfg {
237 pub num_v_heads: usize,
238 pub num_k_heads: usize,
239 pub key_head_dim: usize,
240 pub value_head_dim: usize,
241 pub conv_kernel: usize,
242 pub hidden_size: usize,
243 pub rms_eps: f64,
244}
245
246impl GdnCfg {
247 pub fn conv_dim(&self) -> usize {
248 2 * self.num_k_heads * self.key_head_dim + self.num_v_heads * self.value_head_dim
249 }
250
251 pub fn state_len(&self) -> usize {
254 (self.conv_kernel - 1) * self.conv_dim()
255 + self.num_v_heads * self.key_head_dim * self.value_head_dim
256 }
257}
258
259fn softplus(x: f64) -> f64 {
260 if x > 20.0 {
261 x
262 } else {
263 x.exp().ln_1p()
264 }
265}
266
267fn sigmoid(x: f64) -> f64 {
268 1.0 / (1.0 + (-x).exp())
269}
270
271fn silu(x: f64) -> f64 {
272 x / (1.0 + (-x).exp())
273}
274
275#[derive(Clone, Copy)]
278struct SendMutF32(*mut f32);
279unsafe impl Send for SendMutF32 {}
280unsafe impl Sync for SendMutF32 {}
281
282#[allow(clippy::too_many_arguments)]
295fn gdn_step(
296 qkv: &[f32],
297 z: &[f32],
298 a: &[f32],
299 b: &[f32],
300 w: &GdnWeights,
301 cfg: &GdnCfg,
302 state: &mut [f32],
303 of: &mut [f32],
304 pool: Option<&Pool>,
305) {
306 let (nv, nk, dk, dv, kk) = (
307 cfg.num_v_heads,
308 cfg.num_k_heads,
309 cfg.key_head_dim,
310 cfg.value_head_dim,
311 cfg.conv_kernel,
312 );
313 let c_dim = cfg.conv_dim();
314 let (kd, rep) = (nk * dk, nv / nk);
315 let (ring, s_all) = state.split_at_mut((kk - 1) * c_dim);
316
317 let mut cq = vec![0f32; c_dim];
321 for c in 0..c_dim {
322 let taps = &w.conv1d[c * kk..(c + 1) * kk];
323 let mut acc = qkv[c] as f64 * taps[kk - 1] as f64;
324 for j in 0..kk - 1 {
325 acc += ring[j * c_dim + c] as f64 * taps[j] as f64;
326 }
327 cq[c] = silu(acc) as f32;
328 }
329 if kk > 1 {
331 ring.copy_within(c_dim.., 0);
332 let tail = (kk - 2) * c_dim;
333 ring[tail..tail + c_dim].copy_from_slice(&qkv[..c_dim]);
334 }
335
336 let cq = &cq;
337 let s_ptr = SendMutF32(s_all.as_mut_ptr());
338 let of_ptr = SendMutF32(of.as_mut_ptr());
339 let head_range = |h0: usize, h1: usize| {
340 let (s_ptr, of_ptr) = (s_ptr, of_ptr);
343 let mut kv = crate::attention::take_buf(dv);
345 let mut delta = crate::attention::take_buf(dv);
346 let mut o = crate::attention::take_buf(dv);
347 let mut kf = crate::attention::take_buf(dk);
348 let mut qf = crate::attention::take_buf(dk);
349 for h in h0..h1 {
350 let ko = h / rep; let (qs, ks) = (ko * dk, kd + ko * dk);
352 let (mut nq, mut nkn) = (0f64, 0f64);
354 for d in 0..dk {
355 nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
356 nkn += (cq[ks + d] as f64) * (cq[ks + d] as f64);
357 }
358 let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
359 let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
360 for d in 0..dk {
361 qf[d] = cq[qs + d] * invq;
362 kf[d] = cq[ks + d] * invk;
363 }
364
365 let g = (-(w.a_log[h] as f64).exp()
366 * softplus(a[h] as f64 + w.dt_bias[h] as f64))
367 .exp() as f32;
368 let beta = sigmoid(b[h] as f64) as f32;
369
370 let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
372 let oh =
373 unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
374 let vt = &cq[2 * kd + h * dv..2 * kd + (h + 1) * dv];
375
376 kv[..dv].fill(0.0);
380 for di in 0..dk {
381 let kfd = kf[di];
382 let row = &s[di * dv..(di + 1) * dv];
383 for dj in 0..dv {
384 kv[dj] += row[dj] * kfd; }
386 }
387 for dj in 0..dv {
388 delta[dj] = (vt[dj] - g * kv[dj]) * beta;
389 }
390 o[..dv].fill(0.0);
391 for di in 0..dk {
392 let kfd = kf[di];
393 let qfd = qf[di];
394 let row = &mut s[di * dv..(di + 1) * dv];
395 for dj in 0..dv {
396 let cell = g * row[dj] + kfd * delta[dj];
397 row[dj] = cell;
398 o[dj] += qfd * cell; }
400 }
401 let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
403 let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
404 for dj in 0..dv {
405 oh[dj] =
406 ((o[dj] as f64 * inv) * w.norm[dj] as f64 * silu(z[h * dv + dj] as f64)) as f32;
407 }
408 }
409 crate::attention::recycle_buf(&mut kv);
410 crate::attention::recycle_buf(&mut delta);
411 crate::attention::recycle_buf(&mut o);
412 crate::attention::recycle_buf(&mut kf);
413 crate::attention::recycle_buf(&mut qf);
414 };
415 match pool {
416 Some(pool) if nv >= 4 => pool.run(&|widx, n| {
417 let chunk = nv.div_ceil(n);
418 let h0 = (widx * chunk).min(nv);
419 let h1 = (h0 + chunk).min(nv);
420 if h0 < h1 {
421 head_range(h0, h1);
422 }
423 }),
424 _ => head_range(0, nv),
425 }
426}
427
428pub fn gdn_forward(
430 x: &[f32],
431 w: &GdnWeights,
432 cfg: &GdnCfg,
433 state: &mut Vec<f32>,
434 pool: Option<&Pool>,
435) -> Vec<f32> {
436 if state.len() != cfg.state_len() {
437 *state = vec![0f32; cfg.state_len()];
438 }
439 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
440
441 let mut qkv = vec![0.0f32; c_dim];
442 let mut z = vec![0.0f32; vd];
443 let mut a = vec![0.0f32; cfg.num_v_heads];
444 let mut b = vec![0.0f32; cfg.num_v_heads];
445 let cpu_projs = |qkv: &mut Vec<f32>, z: &mut Vec<f32>, a: &mut Vec<f32>, b: &mut Vec<f32>| {
449 QTensor::matvec_many(
450 [&w.in_proj_qkv, &w.in_proj_z, &w.in_proj_a, &w.in_proj_b],
451 x,
452 [
453 qkv.as_mut_slice(),
454 z.as_mut_slice(),
455 a.as_mut_slice(),
456 b.as_mut_slice(),
457 ],
458 pool,
459 );
460 };
461 let mut done = false;
462 if crate::gpu::enabled_here() && gdn_projs_eligible(w) {
463 match crate::gpu::probe_arm(crate::gpu::OpClass::Batch) {
464 crate::gpu::ProbeArm::Gpu => {
465 let t0 = std::time::Instant::now();
466 if gdn_projs_gpu(w, x, &mut qkv, &mut z) {
467 crate::gpu::probe_record(crate::gpu::OpClass::Batch, true, t0.elapsed());
468 w.in_proj_a.matvec(x, &mut a, pool);
469 w.in_proj_b.matvec(x, &mut b, pool);
470 done = true;
471 }
472 }
473 crate::gpu::ProbeArm::CpuTimed => {
474 let t0 = std::time::Instant::now();
475 crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
476 crate::gpu::probe_record(crate::gpu::OpClass::Batch, false, t0.elapsed());
477 done = true;
478 }
479 crate::gpu::ProbeArm::Cpu => {
480 crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
481 done = true;
482 }
483 }
484 }
485 if !done {
486 cpu_projs(&mut qkv, &mut z, &mut a, &mut b);
487 }
488
489 let mut of = vec![0.0f32; vd];
490 gdn_step(&qkv, &z, &a, &b, w, cfg, state, &mut of, pool);
491
492 let mut out = vec![0.0f32; cfg.hidden_size];
493 w.out_proj.matvec(&of, &mut out, pool);
494 out
495}
496
497pub fn gdn_forward_batch(
502 xs: &[f32],
503 b: usize,
504 w: &GdnWeights,
505 cfg: &GdnCfg,
506 state: &mut Vec<f32>,
507 pool: Option<&Pool>,
508) -> Vec<f32> {
509 if state.len() != cfg.state_len() {
510 *state = vec![0f32; cfg.state_len()];
511 }
512 let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
513 let nv = cfg.num_v_heads;
514
515 let mut qkv = vec![0.0f32; b * c_dim];
516 w.in_proj_qkv.matmat(xs, b, &mut qkv, pool);
517 let mut z = vec![0.0f32; b * vd];
518 w.in_proj_z.matmat(xs, b, &mut z, pool);
519 let mut a = vec![0.0f32; b * nv];
520 w.in_proj_a.matmat(xs, b, &mut a, pool);
521 let mut bb = vec![0.0f32; b * nv];
522 w.in_proj_b.matmat(xs, b, &mut bb, pool);
523
524 let mut of = vec![0.0f32; b * vd];
525 for bi in 0..b {
526 gdn_step(
527 &qkv[bi * c_dim..(bi + 1) * c_dim],
528 &z[bi * vd..(bi + 1) * vd],
529 &a[bi * nv..(bi + 1) * nv],
530 &bb[bi * nv..(bi + 1) * nv],
531 w,
532 cfg,
533 state,
534 &mut of[bi * vd..(bi + 1) * vd],
535 pool,
536 );
537 }
538 let mut out = vec![0.0f32; b * cfg.hidden_size];
539 w.out_proj.matmat(&of, b, &mut out, pool);
540 out
541}
542
543fn gdn_projs_eligible(w: &GdnWeights) -> bool {
547 w.in_proj_qkv.is_q1()
548 || std::env::var("CMF_GPU_GDN").map(|v| v == "1").unwrap_or(false)
549}
550
551fn gdn_projs_gpu(w: &GdnWeights, x: &[f32], qkv: &mut [f32], z: &mut [f32]) -> bool {
553 use crate::gpu::matvec_batch;
554 use crate::qtensor::QTensor;
555 if !crate::gpu::enabled_here() {
556 return false;
557 }
558 fn part<'a>(
559 t: &'a QTensor,
560 x: &[f32],
561 ) -> Option<(std::sync::Arc<cortiq_core::CmfModel>, crate::gpu::BatchJob<'a>)> {
562 use crate::gpu::BatchJob;
563 use crate::qtensor::prescale;
564 use cortiq_core::TensorDtype;
565 match t {
566 QTensor::Mapped {
567 model,
568 idx,
569 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
570 rows,
571 cols,
572 row_scale,
573 col_field,
574 ..
575 } => Some((
576 model.clone(),
577 BatchJob {
578 idx: *idx,
579 rows: *rows,
580 cols: *cols,
581 row_scale,
582 xs: prescale(x, col_field, *dt).into_owned(),
583 q1: false,
584 },
585 )),
586 QTensor::Mapped {
587 model,
588 idx,
589 dtype: TensorDtype::Q1,
590 rows,
591 cols,
592 ..
593 } => Some((
594 model.clone(),
595 BatchJob {
596 idx: *idx,
597 rows: *rows,
598 cols: *cols,
599 row_scale: &[],
600 xs: x.to_vec(),
601 q1: true,
602 },
603 )),
604 _ => None,
605 }
606 }
607 let Some((model, jq)) = part(&w.in_proj_qkv, x) else { return false };
608 let Some((_, jz)) = part(&w.in_proj_z, x) else { return false };
609 matvec_batch(&model, &[jq, jz], &mut [qkv, z])
610}
611
612#[allow(clippy::too_many_arguments)]
615pub fn gdn_pair(
616 x1: &[f32],
617 x2: &[f32],
618 w: &GdnWeights,
619 cfg: &GdnCfg,
620 state: &mut Vec<f32>,
621 scratch: &mut Vec<f32>,
622 pool: Option<&Pool>,
623) -> (Vec<f32>, Vec<f32>) {
624 if state.len() != cfg.state_len() {
625 *state = vec![0f32; cfg.state_len()];
626 }
627 let (c_dim, vd, nv) = (
628 cfg.conv_dim(),
629 cfg.num_v_heads * cfg.value_head_dim,
630 cfg.num_v_heads,
631 );
632
633 let mut qkv1 = vec![0.0f32; c_dim];
634 let mut qkv2 = vec![0.0f32; c_dim];
635 w.in_proj_qkv.matvec2(x1, x2, &mut qkv1, &mut qkv2, pool);
636 let mut z1 = vec![0.0f32; vd];
637 let mut z2 = vec![0.0f32; vd];
638 w.in_proj_z.matvec2(x1, x2, &mut z1, &mut z2, pool);
639 let mut a1 = vec![0.0f32; nv];
640 let mut a2 = vec![0.0f32; nv];
641 w.in_proj_a.matvec2(x1, x2, &mut a1, &mut a2, pool);
642 let mut b1 = vec![0.0f32; nv];
643 let mut b2 = vec![0.0f32; nv];
644 w.in_proj_b.matvec2(x1, x2, &mut b1, &mut b2, pool);
645
646 let mut of1 = vec![0.0f32; vd];
647 gdn_step(&qkv1, &z1, &a1, &b1, w, cfg, state, &mut of1, pool);
648
649 scratch.clear();
650 scratch.extend_from_slice(state);
651 let mut of2 = vec![0.0f32; vd];
652 gdn_step(&qkv2, &z2, &a2, &b2, w, cfg, scratch, &mut of2, pool);
653
654 let mut out1 = vec![0.0f32; cfg.hidden_size];
655 let mut out2 = vec![0.0f32; cfg.hidden_size];
656 w.out_proj.matvec2(&of1, &of2, &mut out1, &mut out2, pool);
657 (out1, out2)
658}
659
660pub struct ShortConvWeights {
667 pub in_proj: QTensor,
669 pub conv: Vec<f32>,
673 pub out_proj: QTensor,
675}
676
677#[derive(Clone, Copy)]
678pub struct ShortConvCfg {
679 pub hidden_size: usize,
680 pub kernel: usize,
682}
683
684impl ShortConvCfg {
685 pub fn state_len(&self) -> usize {
687 (self.kernel - 1) * self.hidden_size
688 }
689}
690
691fn short_conv_step(
701 bcx: &[f32],
702 conv: &[f32],
703 cfg: &ShortConvCfg,
704 ring_state: &mut [f32],
705 y: &mut [f32],
706) {
707 let (h, k) = (cfg.hidden_size, cfg.kernel);
708 let ring = k - 1;
709 let (bg, cg, xg) = (&bcx[0..h], &bcx[h..2 * h], &bcx[2 * h..3 * h]);
710 for c in 0..h {
711 let bx = bg[c] * xg[c];
712 let wc = &conv[c * k..(c + 1) * k];
713 let mut acc = wc[k - 1] * bx;
715 let rc = &mut ring_state[c * ring..c * ring + ring];
716 for s in 0..ring {
717 acc += wc[k - 2 - s] * rc[s];
718 }
719 y[c] = cg[c] * acc;
720 for s in (1..ring).rev() {
722 rc[s] = rc[s - 1];
723 }
724 if ring > 0 {
725 rc[0] = bx;
726 }
727 }
728}
729
730pub fn short_conv_forward(
732 x: &[f32],
733 w: &ShortConvWeights,
734 cfg: &ShortConvCfg,
735 state: &mut Vec<f32>,
736 pool: Option<&Pool>,
737) -> Vec<f32> {
738 if state.len() != cfg.state_len() {
739 *state = vec![0f32; cfg.state_len()];
740 }
741 let h = cfg.hidden_size;
742 let mut bcx = vec![0.0f32; 3 * h];
743 w.in_proj.matvec(x, &mut bcx, pool);
744 let mut y = vec![0.0f32; h];
745 short_conv_step(&bcx, &w.conv, cfg, state, &mut y);
746 let mut out = vec![0.0f32; h];
747 w.out_proj.matvec(&y, &mut out, pool);
748 out
749}
750
751pub fn short_conv_forward_batch(
756 xs: &[f32],
757 b: usize,
758 w: &ShortConvWeights,
759 cfg: &ShortConvCfg,
760 state: &mut Vec<f32>,
761 pool: Option<&Pool>,
762) -> Vec<f32> {
763 if state.len() != cfg.state_len() {
764 *state = vec![0f32; cfg.state_len()];
765 }
766 let h = cfg.hidden_size;
767 let mut bcx = vec![0.0f32; b * 3 * h];
768 w.in_proj.matmat(xs, b, &mut bcx, pool);
769 let mut y = vec![0.0f32; b * h];
770 for bi in 0..b {
771 short_conv_step(
772 &bcx[bi * 3 * h..(bi + 1) * 3 * h],
773 &w.conv,
774 cfg,
775 state,
776 &mut y[bi * h..(bi + 1) * h],
777 );
778 }
779 let mut out = vec![0.0f32; b * h];
780 w.out_proj.matmat(&y, b, &mut out, pool);
781 out
782}
783
784#[allow(clippy::too_many_arguments)]
789pub fn short_conv_pair(
790 x1: &[f32],
791 x2: &[f32],
792 w: &ShortConvWeights,
793 cfg: &ShortConvCfg,
794 state: &mut Vec<f32>,
795 scratch: &mut Vec<f32>,
796 pool: Option<&Pool>,
797) -> (Vec<f32>, Vec<f32>) {
798 if state.len() != cfg.state_len() {
799 *state = vec![0f32; cfg.state_len()];
800 }
801 let h = cfg.hidden_size;
802 let mut bcx1 = vec![0.0f32; 3 * h];
803 let mut bcx2 = vec![0.0f32; 3 * h];
804 w.in_proj.matvec2(x1, x2, &mut bcx1, &mut bcx2, pool);
805
806 let mut y1 = vec![0.0f32; h];
807 short_conv_step(&bcx1, &w.conv, cfg, state, &mut y1);
808 scratch.clear();
809 scratch.extend_from_slice(state);
810 let mut y2 = vec![0.0f32; h];
811 short_conv_step(&bcx2, &w.conv, cfg, scratch, &mut y2);
812
813 let mut out1 = vec![0.0f32; h];
814 let mut out2 = vec![0.0f32; h];
815 w.out_proj.matvec2(&y1, &y2, &mut out1, &mut out2, pool);
816 (out1, out2)
817}
818
819#[cfg(test)]
820mod tests {
821 use super::*;
822
823 fn tiny() -> (VmfPhaseWeights, VmfPhaseCfg) {
824 let cfg = VmfPhaseCfg {
825 num_heads: 2,
826 nphase: 3,
827 value_head_dim: 4,
828 hidden_size: 8,
829 phase_mass: 0.0,
830 };
831 let synth = |rows: usize, cols: usize, salt: usize| {
832 QTensor::from_f32(
833 (0..rows * cols)
834 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
835 .collect(),
836 rows,
837 cols,
838 )
839 };
840 let w = VmfPhaseWeights {
841 thq: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 1),
842 thk: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 2),
843 v_proj: synth(cfg.num_heads * cfg.value_head_dim, cfg.hidden_size, 3),
844 out_proj: synth(cfg.hidden_size, cfg.num_heads * cfg.value_head_dim, 4),
845 decay: (0..cfg.num_heads * 2 * cfg.nphase)
846 .map(|i| 0.9 + 0.005 * (i % 10) as f64)
847 .collect(),
848 k_gate: None,
849 };
850 (w, cfg)
851 }
852
853 #[test]
854 fn state_persists_and_changes_output() {
855 let (w, cfg) = tiny();
856 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
857 let mut state = Vec::new();
858 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
859 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
860 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
862 assert_eq!(state.len(), cfg.state_len());
863 }
864
865 #[test]
869 fn phase_mass_zero_is_noop_and_positive_shifts() {
870 let (w, cfg0) = tiny();
871 let mut cfg_m = cfg0.clone();
872 cfg_m.phase_mass = 1.0;
873 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.4).sin()).collect();
874
875 let mut s0 = Vec::new();
876 let base = vmf_phase_forward(&x, &w, &cfg0, &mut s0, None);
877 let mut s0b = Vec::new();
879 let base2 = vmf_phase_forward(&x, &w, &cfg0, &mut s0b, None);
880 assert_eq!(base, base2, "mass=0 must be deterministic/no-op");
881 let mut sm = Vec::new();
883 let massed = vmf_phase_forward(&x, &w, &cfg_m, &mut sm, None);
884 assert!(
885 base.iter().zip(&massed).any(|(a, b)| (a - b).abs() > 1e-5),
886 "mass>0 must change the output"
887 );
888 assert!(massed.iter().all(|v| v.is_finite()));
889 }
890
891 #[test]
896 fn kappa_gate_open_matches_none_and_closed_writes_nothing() {
897 let (mut w, cfg) = tiny();
898 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
899
900 let mut s_none = Vec::new();
901 let base1 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
902 let base2 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
903
904 w.k_gate = Some((
906 QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
907 vec![20.0; cfg.num_heads],
908 ));
909 let mut s_open = Vec::new();
910 let o1 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
911 let o2 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
912 for (a, b) in base1.iter().zip(&o1).chain(base2.iter().zip(&o2)) {
913 assert!((a - b).abs() < 1e-5, "open κ must match gateless: {a} vs {b}");
914 }
915
916 w.k_gate = Some((
918 QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
919 vec![-20.0; cfg.num_heads],
920 ));
921 let mut s_closed = Vec::new();
922 let oc = vmf_phase_forward(&x, &w, &cfg, &mut s_closed, None);
923 assert!(s_closed.iter().all(|&v| v.abs() < 1e-7), "closed κ: state must stay empty");
924 assert!(oc.iter().all(|&v| v.abs() < 1e-6), "closed κ: empty-condensate readout");
925 }
926
927 #[test]
928 fn pair_matches_two_singles_bitexact() {
929 let (w, cfg) = tiny();
930 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
931 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
932
933 let mut s_ref = Vec::new();
935 let r1 = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
936 let r2 = vmf_phase_forward(&x2, &w, &cfg, &mut s_ref, None);
937
938 let mut s = Vec::new();
940 let mut scratch = Vec::new();
941 let (p1, p2) = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
942 assert_eq!(r1, p1, "lane 1 must be bit-identical");
943 assert_eq!(r2, p2, "lane 2 must be bit-identical");
944 std::mem::swap(&mut s, &mut scratch);
946 assert_eq!(s, s_ref, "accepted state must equal sequential state");
947 }
948
949 #[test]
950 fn rejected_draft_leaves_state_at_lane1() {
951 let (w, cfg) = tiny();
952 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
953 let x2 = vec![0.5f32; 8];
954
955 let mut s_ref = Vec::new();
956 let _ = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
957
958 let mut s = Vec::new();
959 let mut scratch = Vec::new();
960 let _ = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
961 assert_eq!(s, s_ref);
963 }
964
965 fn tiny_gdn() -> (GdnWeights, GdnCfg) {
968 let cfg = GdnCfg {
969 num_v_heads: 4,
970 num_k_heads: 2,
971 key_head_dim: 3,
972 value_head_dim: 5,
973 conv_kernel: 4,
974 hidden_size: 8,
975 rms_eps: 1e-6,
976 };
977 let c_dim = cfg.conv_dim();
978 let vd = cfg.num_v_heads * cfg.value_head_dim;
979 let synth = |rows: usize, cols: usize, salt: usize| {
980 QTensor::from_f32(
981 (0..rows * cols)
982 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
983 .collect(),
984 rows,
985 cols,
986 )
987 };
988 let vecf = |n: usize, salt: usize| -> Vec<f32> {
989 (0..n)
990 .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.6)
991 .collect()
992 };
993 let w = GdnWeights {
994 in_proj_qkv: synth(c_dim, cfg.hidden_size, 1),
995 in_proj_z: synth(vd, cfg.hidden_size, 2),
996 in_proj_a: synth(cfg.num_v_heads, cfg.hidden_size, 3),
997 in_proj_b: synth(cfg.num_v_heads, cfg.hidden_size, 4),
998 conv1d: vecf(c_dim * cfg.conv_kernel, 5),
999 a_log: (0..cfg.num_v_heads).map(|i| 0.2 + 0.3 * i as f32).collect(),
1000 dt_bias: vecf(cfg.num_v_heads, 6),
1001 norm: vec![1.0; cfg.value_head_dim],
1002 out_proj: synth(cfg.hidden_size, vd, 7),
1003 };
1004 (w, cfg)
1005 }
1006
1007 #[test]
1008 fn gdn_state_persists_and_changes_output() {
1009 let (w, cfg) = tiny_gdn();
1010 let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
1011 let mut state = Vec::new();
1012 let o1 = gdn_forward(&x, &w, &cfg, &mut state, None);
1013 let o2 = gdn_forward(&x, &w, &cfg, &mut state, None);
1014 assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
1015 assert_eq!(state.len(), cfg.state_len());
1016 }
1017
1018 #[test]
1019 fn gdn_pair_matches_two_singles_bitexact() {
1020 let (w, cfg) = tiny_gdn();
1021 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
1022 let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
1023
1024 let mut s_ref = Vec::new();
1025 let r1 = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
1026 let r2 = gdn_forward(&x2, &w, &cfg, &mut s_ref, None);
1027
1028 let mut s = Vec::new();
1029 let mut scratch = Vec::new();
1030 let (p1, p2) = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
1031 assert_eq!(r1, p1, "lane 1 must be bit-identical");
1032 assert_eq!(r2, p2, "lane 2 must be bit-identical");
1033 std::mem::swap(&mut s, &mut scratch);
1034 assert_eq!(s, s_ref, "accepted state must equal sequential state");
1035 }
1036
1037 #[test]
1038 fn gdn_rejected_draft_leaves_state_at_lane1() {
1039 let (w, cfg) = tiny_gdn();
1040 let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
1041 let x2 = vec![0.5f32; 8];
1042
1043 let mut s_ref = Vec::new();
1044 let _ = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
1045
1046 let mut s = Vec::new();
1047 let mut scratch = Vec::new();
1048 let _ = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
1049 assert_eq!(s, s_ref);
1050 }
1051
1052 #[test]
1056 fn gdn_conv_ring_matches_explicit_causal_conv() {
1057 let (w, cfg) = tiny_gdn();
1058 let seq: Vec<Vec<f32>> = (0..6)
1059 .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.17).sin()).collect())
1060 .collect();
1061
1062 let mut s_inc = Vec::new();
1065 for (t, x) in seq.iter().enumerate() {
1066 let inc = gdn_forward(x, &w, &cfg, &mut s_inc, None);
1067 let mut s_replay = Vec::new();
1068 let mut replay = Vec::new();
1069 for xr in &seq[..=t] {
1070 replay = gdn_forward(xr, &w, &cfg, &mut s_replay, None);
1071 }
1072 assert_eq!(inc, replay, "position {t}: ring must equal replay");
1073 }
1074 }
1075
1076 fn tiny_short_conv() -> (ShortConvWeights, ShortConvCfg) {
1077 let cfg = ShortConvCfg { hidden_size: 8, kernel: 3 };
1078 let synth = |rows: usize, cols: usize, salt: usize| {
1079 QTensor::from_f32(
1080 (0..rows * cols)
1081 .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.5)
1082 .collect(),
1083 rows,
1084 cols,
1085 )
1086 };
1087 let w = ShortConvWeights {
1088 in_proj: synth(3 * cfg.hidden_size, cfg.hidden_size, 1),
1089 conv: (0..cfg.hidden_size * cfg.kernel)
1090 .map(|i| ((i * 7 % 13) as f32 / 13.0 - 0.5) * 0.8)
1091 .collect(),
1092 out_proj: synth(cfg.hidden_size, cfg.hidden_size, 2),
1093 };
1094 (w, cfg)
1095 }
1096
1097 #[test]
1100 fn short_conv_ring_matches_explicit_causal_conv() {
1101 let (w, cfg) = tiny_short_conv();
1102 let seq: Vec<Vec<f32>> = (0..6)
1103 .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.19).cos()).collect())
1104 .collect();
1105 let mut s_inc = Vec::new();
1106 for (t, x) in seq.iter().enumerate() {
1107 let inc = short_conv_forward(x, &w, &cfg, &mut s_inc, None);
1108 let mut s_replay = Vec::new();
1109 let mut replay = Vec::new();
1110 for xr in &seq[..=t] {
1111 replay = short_conv_forward(xr, &w, &cfg, &mut s_replay, None);
1112 }
1113 assert_eq!(inc, replay, "position {t}: ring must equal replay");
1114 assert_eq!(s_inc.len(), cfg.state_len());
1115 }
1116 }
1117
1118 #[test]
1121 fn short_conv_batch_matches_sequential() {
1122 let (w, cfg) = tiny_short_conv();
1123 let b = 5;
1124 let xs: Vec<f32> =
1125 (0..b * cfg.hidden_size).map(|i| (i as f32 * 0.13).sin() * 0.6).collect();
1126
1127 let mut s_seq = Vec::new();
1128 let mut seq_out = vec![0.0f32; b * cfg.hidden_size];
1129 for bi in 0..b {
1130 let o = short_conv_forward(
1131 &xs[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size],
1132 &w,
1133 &cfg,
1134 &mut s_seq,
1135 None,
1136 );
1137 seq_out[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size].copy_from_slice(&o);
1138 }
1139
1140 let mut s_batch = Vec::new();
1141 let batch_out = short_conv_forward_batch(&xs, b, &w, &cfg, &mut s_batch, None);
1142 assert_eq!(seq_out, batch_out, "batch conv must match sequential decode");
1143 assert_eq!(s_seq, s_batch, "ring state must match after the chunk");
1144 }
1145}