use crate::pool::Pool;
use crate::qtensor::QTensor;
pub struct VmfPhaseWeights {
pub thq: QTensor,
pub thk: QTensor,
pub v_proj: QTensor,
pub out_proj: QTensor,
pub decay: Vec<f64>,
pub k_gate: Option<(QTensor, Vec<f32>)>,
}
#[derive(Clone, Copy)]
pub struct VmfPhaseCfg {
pub num_heads: usize,
pub nphase: usize,
pub value_head_dim: usize,
pub hidden_size: usize,
pub phase_mass: f32,
}
impl VmfPhaseCfg {
pub fn state_len(&self) -> usize {
self.num_heads * 2 * self.nphase * self.value_head_dim
}
}
fn phase_step(
thq: &[f32],
thk: &[f32],
v: &[f32],
decay: &[f64],
kap: Option<&[f32]>,
cfg: &VmfPhaseCfg,
state: &mut [f32],
out: &mut [f32],
) {
let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
let mscale = 1.0f64 / (1.0 + cfg.phase_mass as f64);
let p2 = 2 * nph;
for h in 0..nh {
let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
let thk_h = &thk[h * nph..(h + 1) * nph];
let thq_h = &thq[h * nph..(h + 1) * nph];
let vt = &v[h * dv..(h + 1) * dv];
let ot = &mut out[h * dv..(h + 1) * dv];
let dec = &decay[h * p2..(h + 1) * p2];
let kh = kap.map_or(1.0f64, |k| k[h] as f64);
for f in 0..p2 {
let (fk, fq) = if f < nph {
((thk_h[f] as f64 * mscale).cos(), (thq_h[f] as f64 * mscale).cos())
} else {
(
(thk_h[f - nph] as f64 * mscale).sin(),
(thq_h[f - nph] as f64 * mscale).sin(),
)
};
let fkw = fk * kh;
let row = &mut s[f * dv..(f + 1) * dv];
let dcf = dec[f];
for d in 0..dv {
let cell = dcf * row[d] as f64 + fkw * vt[d] as f64;
row[d] = cell as f32;
ot[d] += (fq * cell) as f32; }
}
}
}
fn kappa_of(x: &[f32], w: &VmfPhaseWeights, nh: usize, pool: Option<&Pool>) -> Option<Vec<f32>> {
let (kw, kb) = w.k_gate.as_ref()?;
let mut k = vec![0.0f32; nh];
kw.matvec(x, &mut k, pool);
for (v, b) in k.iter_mut().zip(kb) {
*v = 1.0 / (1.0 + (-(*v + b)).exp());
}
Some(k)
}
pub fn vmf_phase_forward(
x: &[f32],
w: &VmfPhaseWeights,
cfg: &VmfPhaseCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
let mut thq = vec![0.0f32; nh * nph];
w.thq.matvec(x, &mut thq, pool);
let mut thk = vec![0.0f32; nh * nph];
w.thk.matvec(x, &mut thk, pool);
let mut v = vec![0.0f32; nh * dv];
w.v_proj.matvec(x, &mut v, pool);
let kap = kappa_of(x, w, nh, pool);
let mut o = vec![0.0f32; nh * dv];
phase_step(&thq, &thk, &v, &w.decay, kap.as_deref(), cfg, state, &mut o);
let mut out = vec![0.0f32; cfg.hidden_size];
w.out_proj.matvec(&o, &mut out, pool);
out
}
#[allow(clippy::too_many_arguments)]
pub fn vmf_phase_pair(
x1: &[f32],
x2: &[f32],
w: &VmfPhaseWeights,
cfg: &VmfPhaseCfg,
state: &mut Vec<f32>,
scratch: &mut Vec<f32>,
pool: Option<&Pool>,
) -> (Vec<f32>, Vec<f32>) {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
let mut thq1 = vec![0.0f32; nh * nph];
let mut thq2 = vec![0.0f32; nh * nph];
w.thq.matvec2(x1, x2, &mut thq1, &mut thq2, pool);
let mut thk1 = vec![0.0f32; nh * nph];
let mut thk2 = vec![0.0f32; nh * nph];
w.thk.matvec2(x1, x2, &mut thk1, &mut thk2, pool);
let mut v1 = vec![0.0f32; nh * dv];
let mut v2 = vec![0.0f32; nh * dv];
w.v_proj.matvec2(x1, x2, &mut v1, &mut v2, pool);
let kap1 = kappa_of(x1, w, nh, pool);
let mut o1 = vec![0.0f32; nh * dv];
phase_step(&thq1, &thk1, &v1, &w.decay, kap1.as_deref(), cfg, state, &mut o1);
let kap2 = kappa_of(x2, w, nh, pool);
scratch.clear();
scratch.extend_from_slice(state);
let mut o2 = vec![0.0f32; nh * dv];
phase_step(&thq2, &thk2, &v2, &w.decay, kap2.as_deref(), cfg, scratch, &mut o2);
let mut out1 = vec![0.0f32; cfg.hidden_size];
let mut out2 = vec![0.0f32; cfg.hidden_size];
w.out_proj.matvec2(&o1, &o2, &mut out1, &mut out2, pool);
(out1, out2)
}
pub struct GdnWeights {
pub in_proj_qkv: QTensor,
pub in_proj_z: QTensor,
pub in_proj_a: QTensor,
pub in_proj_b: QTensor,
pub conv1d: Vec<f32>,
pub a_log: Vec<f32>,
pub dt_bias: Vec<f32>,
pub norm: Vec<f32>,
pub out_proj: QTensor,
}
#[derive(Clone, Copy)]
pub struct GdnCfg {
pub num_v_heads: usize,
pub num_k_heads: usize,
pub key_head_dim: usize,
pub value_head_dim: usize,
pub conv_kernel: usize,
pub hidden_size: usize,
pub rms_eps: f64,
}
impl GdnCfg {
pub fn conv_dim(&self) -> usize {
2 * self.num_k_heads * self.key_head_dim + self.num_v_heads * self.value_head_dim
}
pub fn state_len(&self) -> usize {
(self.conv_kernel - 1) * self.conv_dim()
+ self.num_v_heads * self.key_head_dim * self.value_head_dim
}
}
fn softplus(x: f64) -> f64 {
if x > 20.0 {
x
} else {
x.exp().ln_1p()
}
}
fn sigmoid(x: f64) -> f64 {
1.0 / (1.0 + (-x).exp())
}
fn silu(x: f64) -> f64 {
x / (1.0 + (-x).exp())
}
#[derive(Clone, Copy)]
struct SendMutF32(*mut f32);
unsafe impl Send for SendMutF32 {}
unsafe impl Sync for SendMutF32 {}
#[allow(clippy::too_many_arguments)]
fn gdn_step(
qkv: &[f32],
z: &[f32],
a: &[f32],
b: &[f32],
w: &GdnWeights,
cfg: &GdnCfg,
state: &mut [f32],
of: &mut [f32],
pool: Option<&Pool>,
) {
let (nv, nk, dk, dv, kk) = (
cfg.num_v_heads,
cfg.num_k_heads,
cfg.key_head_dim,
cfg.value_head_dim,
cfg.conv_kernel,
);
let c_dim = cfg.conv_dim();
let (kd, rep) = (nk * dk, nv / nk);
let (ring, s_all) = state.split_at_mut((kk - 1) * c_dim);
let mut cq = vec![0f32; c_dim];
for c in 0..c_dim {
let taps = &w.conv1d[c * kk..(c + 1) * kk];
let mut acc = qkv[c] as f64 * taps[kk - 1] as f64;
for j in 0..kk - 1 {
acc += ring[j * c_dim + c] as f64 * taps[j] as f64;
}
cq[c] = silu(acc) as f32;
}
if kk > 1 {
ring.copy_within(c_dim.., 0);
let tail = (kk - 2) * c_dim;
ring[tail..tail + c_dim].copy_from_slice(&qkv[..c_dim]);
}
let cq = &cq;
let s_ptr = SendMutF32(s_all.as_mut_ptr());
let of_ptr = SendMutF32(of.as_mut_ptr());
let head_range = |h0: usize, h1: usize| {
let (s_ptr, of_ptr) = (s_ptr, of_ptr);
let mut kv = crate::attention::take_buf(dv);
let mut delta = crate::attention::take_buf(dv);
let mut o = crate::attention::take_buf(dv);
let mut kf = crate::attention::take_buf(dk);
let mut qf = crate::attention::take_buf(dk);
for h in h0..h1 {
let ko = h / rep; let (qs, ks) = (ko * dk, kd + ko * dk);
let (mut nq, mut nkn) = (0f64, 0f64);
for d in 0..dk {
nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
nkn += (cq[ks + d] as f64) * (cq[ks + d] as f64);
}
let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
for d in 0..dk {
qf[d] = cq[qs + d] * invq;
kf[d] = cq[ks + d] * invk;
}
let g = (-(w.a_log[h] as f64).exp()
* softplus(a[h] as f64 + w.dt_bias[h] as f64))
.exp() as f32;
let beta = sigmoid(b[h] as f64) as f32;
let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
let oh =
unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
let vt = &cq[2 * kd + h * dv..2 * kd + (h + 1) * dv];
kv[..dv].fill(0.0);
for di in 0..dk {
let kfd = kf[di];
let row = &s[di * dv..(di + 1) * dv];
for dj in 0..dv {
kv[dj] += row[dj] * kfd; }
}
for dj in 0..dv {
delta[dj] = (vt[dj] - g * kv[dj]) * beta;
}
o[..dv].fill(0.0);
for di in 0..dk {
let kfd = kf[di];
let qfd = qf[di];
let row = &mut s[di * dv..(di + 1) * dv];
for dj in 0..dv {
let cell = g * row[dj] + kfd * delta[dj];
row[dj] = cell;
o[dj] += qfd * cell; }
}
let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
for dj in 0..dv {
oh[dj] =
((o[dj] as f64 * inv) * w.norm[dj] as f64 * silu(z[h * dv + dj] as f64)) as f32;
}
}
crate::attention::recycle_buf(&mut kv);
crate::attention::recycle_buf(&mut delta);
crate::attention::recycle_buf(&mut o);
crate::attention::recycle_buf(&mut kf);
crate::attention::recycle_buf(&mut qf);
};
match pool {
Some(pool) if nv >= 4 => pool.run(&|widx, n| {
let chunk = nv.div_ceil(n);
let h0 = (widx * chunk).min(nv);
let h1 = (h0 + chunk).min(nv);
if h0 < h1 {
head_range(h0, h1);
}
}),
_ => head_range(0, nv),
}
}
pub fn gdn_forward(
x: &[f32],
w: &GdnWeights,
cfg: &GdnCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
let mut qkv = vec![0.0f32; c_dim];
let mut z = vec![0.0f32; vd];
let mut a = vec![0.0f32; cfg.num_v_heads];
let mut b = vec![0.0f32; cfg.num_v_heads];
let cpu_projs = |qkv: &mut Vec<f32>, z: &mut Vec<f32>, a: &mut Vec<f32>, b: &mut Vec<f32>| {
QTensor::matvec_many(
[&w.in_proj_qkv, &w.in_proj_z, &w.in_proj_a, &w.in_proj_b],
x,
[
qkv.as_mut_slice(),
z.as_mut_slice(),
a.as_mut_slice(),
b.as_mut_slice(),
],
pool,
);
};
let mut done = false;
if crate::gpu::enabled_here() && gdn_projs_eligible(w) {
match crate::gpu::probe_arm(crate::gpu::OpClass::Batch) {
crate::gpu::ProbeArm::Gpu => {
let t0 = std::time::Instant::now();
if gdn_projs_gpu(w, x, &mut qkv, &mut z) {
crate::gpu::probe_record(crate::gpu::OpClass::Batch, true, t0.elapsed());
w.in_proj_a.matvec(x, &mut a, pool);
w.in_proj_b.matvec(x, &mut b, pool);
done = true;
}
}
crate::gpu::ProbeArm::CpuTimed => {
let t0 = std::time::Instant::now();
crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
crate::gpu::probe_record(crate::gpu::OpClass::Batch, false, t0.elapsed());
done = true;
}
crate::gpu::ProbeArm::Cpu => {
crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
done = true;
}
}
}
if !done {
cpu_projs(&mut qkv, &mut z, &mut a, &mut b);
}
let mut of = vec![0.0f32; vd];
gdn_step(&qkv, &z, &a, &b, w, cfg, state, &mut of, pool);
let mut out = vec![0.0f32; cfg.hidden_size];
w.out_proj.matvec(&of, &mut out, pool);
out
}
pub fn gdn_forward_batch(
xs: &[f32],
b: usize,
w: &GdnWeights,
cfg: &GdnCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
let nv = cfg.num_v_heads;
let mut qkv = vec![0.0f32; b * c_dim];
w.in_proj_qkv.matmat(xs, b, &mut qkv, pool);
let mut z = vec![0.0f32; b * vd];
w.in_proj_z.matmat(xs, b, &mut z, pool);
let mut a = vec![0.0f32; b * nv];
w.in_proj_a.matmat(xs, b, &mut a, pool);
let mut bb = vec![0.0f32; b * nv];
w.in_proj_b.matmat(xs, b, &mut bb, pool);
let mut of = vec![0.0f32; b * vd];
for bi in 0..b {
gdn_step(
&qkv[bi * c_dim..(bi + 1) * c_dim],
&z[bi * vd..(bi + 1) * vd],
&a[bi * nv..(bi + 1) * nv],
&bb[bi * nv..(bi + 1) * nv],
w,
cfg,
state,
&mut of[bi * vd..(bi + 1) * vd],
pool,
);
}
let mut out = vec![0.0f32; b * cfg.hidden_size];
w.out_proj.matmat(&of, b, &mut out, pool);
out
}
fn gdn_projs_eligible(w: &GdnWeights) -> bool {
w.in_proj_qkv.is_q1()
|| std::env::var("CMF_GPU_GDN").map(|v| v == "1").unwrap_or(false)
}
fn gdn_projs_gpu(w: &GdnWeights, x: &[f32], qkv: &mut [f32], z: &mut [f32]) -> bool {
use crate::gpu::matvec_batch;
use crate::qtensor::QTensor;
if !crate::gpu::enabled_here() {
return false;
}
fn part<'a>(
t: &'a QTensor,
x: &[f32],
) -> Option<(std::sync::Arc<cortiq_core::CmfModel>, crate::gpu::BatchJob<'a>)> {
use crate::gpu::BatchJob;
use crate::qtensor::prescale;
use cortiq_core::TensorDtype;
match t {
QTensor::Mapped {
model,
idx,
dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
rows,
cols,
row_scale,
col_field,
..
} => Some((
model.clone(),
BatchJob {
idx: *idx,
rows: *rows,
cols: *cols,
row_scale,
xs: prescale(x, col_field, *dt).into_owned(),
q1: false,
},
)),
QTensor::Mapped {
model,
idx,
dtype: TensorDtype::Q1,
rows,
cols,
..
} => Some((
model.clone(),
BatchJob {
idx: *idx,
rows: *rows,
cols: *cols,
row_scale: &[],
xs: x.to_vec(),
q1: true,
},
)),
_ => None,
}
}
let Some((model, jq)) = part(&w.in_proj_qkv, x) else { return false };
let Some((_, jz)) = part(&w.in_proj_z, x) else { return false };
matvec_batch(&model, &[jq, jz], &mut [qkv, z])
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_pair(
x1: &[f32],
x2: &[f32],
w: &GdnWeights,
cfg: &GdnCfg,
state: &mut Vec<f32>,
scratch: &mut Vec<f32>,
pool: Option<&Pool>,
) -> (Vec<f32>, Vec<f32>) {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (c_dim, vd, nv) = (
cfg.conv_dim(),
cfg.num_v_heads * cfg.value_head_dim,
cfg.num_v_heads,
);
let mut qkv1 = vec![0.0f32; c_dim];
let mut qkv2 = vec![0.0f32; c_dim];
w.in_proj_qkv.matvec2(x1, x2, &mut qkv1, &mut qkv2, pool);
let mut z1 = vec![0.0f32; vd];
let mut z2 = vec![0.0f32; vd];
w.in_proj_z.matvec2(x1, x2, &mut z1, &mut z2, pool);
let mut a1 = vec![0.0f32; nv];
let mut a2 = vec![0.0f32; nv];
w.in_proj_a.matvec2(x1, x2, &mut a1, &mut a2, pool);
let mut b1 = vec![0.0f32; nv];
let mut b2 = vec![0.0f32; nv];
w.in_proj_b.matvec2(x1, x2, &mut b1, &mut b2, pool);
let mut of1 = vec![0.0f32; vd];
gdn_step(&qkv1, &z1, &a1, &b1, w, cfg, state, &mut of1, pool);
scratch.clear();
scratch.extend_from_slice(state);
let mut of2 = vec![0.0f32; vd];
gdn_step(&qkv2, &z2, &a2, &b2, w, cfg, scratch, &mut of2, pool);
let mut out1 = vec![0.0f32; cfg.hidden_size];
let mut out2 = vec![0.0f32; cfg.hidden_size];
w.out_proj.matvec2(&of1, &of2, &mut out1, &mut out2, pool);
(out1, out2)
}
#[cfg(test)]
mod tests {
use super::*;
fn tiny() -> (VmfPhaseWeights, VmfPhaseCfg) {
let cfg = VmfPhaseCfg {
num_heads: 2,
nphase: 3,
value_head_dim: 4,
hidden_size: 8,
phase_mass: 0.0,
};
let synth = |rows: usize, cols: usize, salt: usize| {
QTensor::from_f32(
(0..rows * cols)
.map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
.collect(),
rows,
cols,
)
};
let w = VmfPhaseWeights {
thq: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 1),
thk: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 2),
v_proj: synth(cfg.num_heads * cfg.value_head_dim, cfg.hidden_size, 3),
out_proj: synth(cfg.hidden_size, cfg.num_heads * cfg.value_head_dim, 4),
decay: (0..cfg.num_heads * 2 * cfg.nphase)
.map(|i| 0.9 + 0.005 * (i % 10) as f64)
.collect(),
k_gate: None,
};
(w, cfg)
}
#[test]
fn state_persists_and_changes_output() {
let (w, cfg) = tiny();
let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
let mut state = Vec::new();
let o1 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
let o2 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
assert_eq!(state.len(), cfg.state_len());
}
#[test]
fn phase_mass_zero_is_noop_and_positive_shifts() {
let (w, cfg0) = tiny();
let mut cfg_m = cfg0.clone();
cfg_m.phase_mass = 1.0;
let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.4).sin()).collect();
let mut s0 = Vec::new();
let base = vmf_phase_forward(&x, &w, &cfg0, &mut s0, None);
let mut s0b = Vec::new();
let base2 = vmf_phase_forward(&x, &w, &cfg0, &mut s0b, None);
assert_eq!(base, base2, "mass=0 must be deterministic/no-op");
let mut sm = Vec::new();
let massed = vmf_phase_forward(&x, &w, &cfg_m, &mut sm, None);
assert!(
base.iter().zip(&massed).any(|(a, b)| (a - b).abs() > 1e-5),
"mass>0 must change the output"
);
assert!(massed.iter().all(|v| v.is_finite()));
}
#[test]
fn kappa_gate_open_matches_none_and_closed_writes_nothing() {
let (mut w, cfg) = tiny();
let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
let mut s_none = Vec::new();
let base1 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
let base2 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
w.k_gate = Some((
QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
vec![20.0; cfg.num_heads],
));
let mut s_open = Vec::new();
let o1 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
let o2 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
for (a, b) in base1.iter().zip(&o1).chain(base2.iter().zip(&o2)) {
assert!((a - b).abs() < 1e-5, "open κ must match gateless: {a} vs {b}");
}
w.k_gate = Some((
QTensor::from_f32(vec![0.0; cfg.num_heads * cfg.hidden_size], cfg.num_heads, cfg.hidden_size),
vec![-20.0; cfg.num_heads],
));
let mut s_closed = Vec::new();
let oc = vmf_phase_forward(&x, &w, &cfg, &mut s_closed, None);
assert!(s_closed.iter().all(|&v| v.abs() < 1e-7), "closed κ: state must stay empty");
assert!(oc.iter().all(|&v| v.abs() < 1e-6), "closed κ: empty-condensate readout");
}
#[test]
fn pair_matches_two_singles_bitexact() {
let (w, cfg) = tiny();
let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
let mut s_ref = Vec::new();
let r1 = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
let r2 = vmf_phase_forward(&x2, &w, &cfg, &mut s_ref, None);
let mut s = Vec::new();
let mut scratch = Vec::new();
let (p1, p2) = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
assert_eq!(r1, p1, "lane 1 must be bit-identical");
assert_eq!(r2, p2, "lane 2 must be bit-identical");
std::mem::swap(&mut s, &mut scratch);
assert_eq!(s, s_ref, "accepted state must equal sequential state");
}
#[test]
fn rejected_draft_leaves_state_at_lane1() {
let (w, cfg) = tiny();
let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
let x2 = vec![0.5f32; 8];
let mut s_ref = Vec::new();
let _ = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
let mut s = Vec::new();
let mut scratch = Vec::new();
let _ = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
assert_eq!(s, s_ref);
}
fn tiny_gdn() -> (GdnWeights, GdnCfg) {
let cfg = GdnCfg {
num_v_heads: 4,
num_k_heads: 2,
key_head_dim: 3,
value_head_dim: 5,
conv_kernel: 4,
hidden_size: 8,
rms_eps: 1e-6,
};
let c_dim = cfg.conv_dim();
let vd = cfg.num_v_heads * cfg.value_head_dim;
let synth = |rows: usize, cols: usize, salt: usize| {
QTensor::from_f32(
(0..rows * cols)
.map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
.collect(),
rows,
cols,
)
};
let vecf = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.6)
.collect()
};
let w = GdnWeights {
in_proj_qkv: synth(c_dim, cfg.hidden_size, 1),
in_proj_z: synth(vd, cfg.hidden_size, 2),
in_proj_a: synth(cfg.num_v_heads, cfg.hidden_size, 3),
in_proj_b: synth(cfg.num_v_heads, cfg.hidden_size, 4),
conv1d: vecf(c_dim * cfg.conv_kernel, 5),
a_log: (0..cfg.num_v_heads).map(|i| 0.2 + 0.3 * i as f32).collect(),
dt_bias: vecf(cfg.num_v_heads, 6),
norm: vec![1.0; cfg.value_head_dim],
out_proj: synth(cfg.hidden_size, vd, 7),
};
(w, cfg)
}
#[test]
fn gdn_state_persists_and_changes_output() {
let (w, cfg) = tiny_gdn();
let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
let mut state = Vec::new();
let o1 = gdn_forward(&x, &w, &cfg, &mut state, None);
let o2 = gdn_forward(&x, &w, &cfg, &mut state, None);
assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
assert_eq!(state.len(), cfg.state_len());
}
#[test]
fn gdn_pair_matches_two_singles_bitexact() {
let (w, cfg) = tiny_gdn();
let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
let mut s_ref = Vec::new();
let r1 = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
let r2 = gdn_forward(&x2, &w, &cfg, &mut s_ref, None);
let mut s = Vec::new();
let mut scratch = Vec::new();
let (p1, p2) = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
assert_eq!(r1, p1, "lane 1 must be bit-identical");
assert_eq!(r2, p2, "lane 2 must be bit-identical");
std::mem::swap(&mut s, &mut scratch);
assert_eq!(s, s_ref, "accepted state must equal sequential state");
}
#[test]
fn gdn_rejected_draft_leaves_state_at_lane1() {
let (w, cfg) = tiny_gdn();
let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
let x2 = vec![0.5f32; 8];
let mut s_ref = Vec::new();
let _ = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
let mut s = Vec::new();
let mut scratch = Vec::new();
let _ = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
assert_eq!(s, s_ref);
}
#[test]
fn gdn_conv_ring_matches_explicit_causal_conv() {
let (w, cfg) = tiny_gdn();
let seq: Vec<Vec<f32>> = (0..6)
.map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.17).sin()).collect())
.collect();
let mut s_inc = Vec::new();
for (t, x) in seq.iter().enumerate() {
let inc = gdn_forward(x, &w, &cfg, &mut s_inc, None);
let mut s_replay = Vec::new();
let mut replay = Vec::new();
for xr in &seq[..=t] {
replay = gdn_forward(xr, &w, &cfg, &mut s_replay, None);
}
assert_eq!(inc, replay, "position {t}: ring must equal replay");
}
}
}