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)
}
pub struct ShortConvWeights {
pub in_proj: QTensor,
pub conv: Vec<f32>,
pub out_proj: QTensor,
}
#[derive(Clone, Copy)]
pub struct ShortConvCfg {
pub hidden_size: usize,
pub kernel: usize,
}
impl ShortConvCfg {
pub fn state_len(&self) -> usize {
(self.kernel - 1) * self.hidden_size
}
}
fn short_conv_step(
bcx: &[f32],
conv: &[f32],
cfg: &ShortConvCfg,
ring_state: &mut [f32],
y: &mut [f32],
) {
let (h, k) = (cfg.hidden_size, cfg.kernel);
let ring = k - 1;
let (bg, cg, xg) = (&bcx[0..h], &bcx[h..2 * h], &bcx[2 * h..3 * h]);
for c in 0..h {
let bx = bg[c] * xg[c];
let wc = &conv[c * k..(c + 1) * k];
let mut acc = wc[k - 1] * bx;
let rc = &mut ring_state[c * ring..c * ring + ring];
for s in 0..ring {
acc += wc[k - 2 - s] * rc[s];
}
y[c] = cg[c] * acc;
for s in (1..ring).rev() {
rc[s] = rc[s - 1];
}
if ring > 0 {
rc[0] = bx;
}
}
}
pub fn short_conv_forward(
x: &[f32],
w: &ShortConvWeights,
cfg: &ShortConvCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let h = cfg.hidden_size;
let mut bcx = vec![0.0f32; 3 * h];
w.in_proj.matvec(x, &mut bcx, pool);
let mut y = vec![0.0f32; h];
short_conv_step(&bcx, &w.conv, cfg, state, &mut y);
let mut out = vec![0.0f32; h];
w.out_proj.matvec(&y, &mut out, pool);
out
}
pub fn short_conv_forward_batch(
xs: &[f32],
b: usize,
w: &ShortConvWeights,
cfg: &ShortConvCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let h = cfg.hidden_size;
let mut bcx = vec![0.0f32; b * 3 * h];
w.in_proj.matmat(xs, b, &mut bcx, pool);
let mut y = vec![0.0f32; b * h];
for bi in 0..b {
short_conv_step(
&bcx[bi * 3 * h..(bi + 1) * 3 * h],
&w.conv,
cfg,
state,
&mut y[bi * h..(bi + 1) * h],
);
}
let mut out = vec![0.0f32; b * h];
w.out_proj.matmat(&y, b, &mut out, pool);
out
}
#[allow(clippy::too_many_arguments)]
pub fn short_conv_pair(
x1: &[f32],
x2: &[f32],
w: &ShortConvWeights,
cfg: &ShortConvCfg,
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 h = cfg.hidden_size;
let mut bcx1 = vec![0.0f32; 3 * h];
let mut bcx2 = vec![0.0f32; 3 * h];
w.in_proj.matvec2(x1, x2, &mut bcx1, &mut bcx2, pool);
let mut y1 = vec![0.0f32; h];
short_conv_step(&bcx1, &w.conv, cfg, state, &mut y1);
scratch.clear();
scratch.extend_from_slice(state);
let mut y2 = vec![0.0f32; h];
short_conv_step(&bcx2, &w.conv, cfg, scratch, &mut y2);
let mut out1 = vec![0.0f32; h];
let mut out2 = vec![0.0f32; h];
w.out_proj.matvec2(&y1, &y2, &mut out1, &mut out2, pool);
(out1, out2)
}
pub struct KdaWeights {
pub q_proj: QTensor,
pub k_proj: QTensor,
pub v_proj: QTensor,
pub conv_q: Vec<f32>,
pub conv_k: Vec<f32>,
pub conv_v: Vec<f32>,
pub f_a: QTensor,
pub f_b: QTensor,
pub dt_bias: Vec<f32>,
pub a_log: Vec<f32>,
pub b_proj: QTensor,
pub gate: KdaOutGate,
pub o_norm: Vec<f32>,
pub o_proj: QTensor,
pub gate_lower_bound: Option<f32>,
}
pub enum KdaOutGate {
Full(QTensor),
LowRank(QTensor, QTensor),
}
#[derive(Clone, Copy)]
pub struct KdaCfg {
pub num_heads: usize,
pub head_k_dim: usize,
pub head_v_dim: usize,
pub conv_kernel: usize,
pub hidden_size: usize,
pub rms_eps: f64,
}
impl KdaCfg {
pub fn state_len(&self) -> usize {
let (nh, dk, dv, kk) = (
self.num_heads,
self.head_k_dim,
self.head_v_dim,
self.conv_kernel,
);
(kk - 1) * (2 * nh * dk + nh * dv) + nh * dk * dv
}
}
fn kda_conv(raw: &[f32], taps: &[f32], ring: &mut [f32], kk: usize, out: &mut [f32]) {
let c_dim = raw.len();
for c in 0..c_dim {
let t = &taps[c * kk..(c + 1) * kk];
let mut acc = raw[c] as f64 * t[kk - 1] as f64;
for j in 0..kk - 1 {
acc += ring[j * c_dim + c] as f64 * t[j] as f64;
}
out[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(raw);
}
}
#[inline]
fn kda_log_decay(w: &KdaWeights, cfg: &KdaCfg, h: usize, d: usize, f: f32) -> f64 {
let (nh, dk) = (cfg.num_heads, cfg.head_k_dim);
let a = if w.a_log.len() == nh {
w.a_log[h] as f64
} else if w.a_log.len() == dk {
w.a_log[d] as f64
} else {
w.a_log[h * dk + d] as f64
};
let raw = f as f64 + w.dt_bias[h * dk + d] as f64;
match w.gate_lower_bound {
Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
None => -a.exp() * softplus(raw),
}
}
#[allow(clippy::too_many_arguments)]
fn kda_step(
xq: &[f32],
xk: &[f32],
xv: &[f32],
f: &[f32],
b: &[f32],
gate_out: &[f32],
w: &KdaWeights,
cfg: &KdaCfg,
state: &mut [f32],
of: &mut [f32],
pool: Option<&Pool>,
) {
let (nh, dk, dv, kk) = (
cfg.num_heads,
cfg.head_k_dim,
cfg.head_v_dim,
cfg.conv_kernel,
);
let (kd, vd) = (nh * dk, nh * dv);
let ring_q_len = (kk - 1) * kd;
let ring_v_len = (kk - 1) * vd;
let (ring_q, rest) = state.split_at_mut(ring_q_len);
let (ring_k, rest) = rest.split_at_mut(ring_q_len);
let (ring_v, s_all) = rest.split_at_mut(ring_v_len);
let mut cq = vec![0f32; kd];
let mut ck = vec![0f32; kd];
let mut cv = vec![0f32; vd];
kda_conv(xq, &w.conv_q, ring_q, kk, &mut cq);
kda_conv(xk, &w.conv_k, ring_k, kk, &mut ck);
kda_conv(xv, &w.conv_v, ring_v, kk, &mut cv);
let (cq, ck, cv) = (&cq, &ck, &cv);
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);
let mut gd = crate::attention::take_buf(dk);
for h in h0..h1 {
let qs = h * 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 += (ck[qs + d] as f64) * (ck[qs + 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] = ck[qs + d] * invk;
gd[d] = kda_log_decay(w, cfg, h, d, f[qs + d]).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 = &cv[h * dv..(h + 1) * dv];
kv[..dv].fill(0.0);
for di in 0..dk {
let kg = kf[di] * gd[di];
let row = &s[di * dv..(di + 1) * dv];
for dj in 0..dv {
kv[dj] += row[dj] * kg;
}
}
for dj in 0..dv {
delta[dj] = (vt[dj] - kv[dj]) * beta;
}
o[..dv].fill(0.0);
for di in 0..dk {
let (kfd, qfd, gdd) = (kf[di], qf[di], gd[di]);
let row = &mut s[di * dv..(di + 1) * dv];
for dj in 0..dv {
let cell = gdd * 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.o_norm[dj] as f64
* sigmoid(gate_out[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);
crate::attention::recycle_buf(&mut gd);
};
match pool {
Some(pool) if nh >= 4 => pool.run(&|widx, n| {
let chunk = nh.div_ceil(n);
let h0 = (widx * chunk).min(nh);
let h1 = (h0 + chunk).min(nh);
if h0 < h1 {
head_range(h0, h1);
}
}),
_ => head_range(0, nh),
}
}
fn kda_gate_out(w: &KdaWeights, x: &[f32], vd: usize, pool: Option<&Pool>) -> Vec<f32> {
let mut g = vec![0.0f32; vd];
match &w.gate {
KdaOutGate::Full(gp) => gp.matvec(x, &mut g, pool),
KdaOutGate::LowRank(ga, gb) => {
let mut low = vec![0.0f32; ga.rows()];
ga.matvec(x, &mut low, pool);
gb.matvec(&low, &mut g, pool);
}
}
g
}
pub fn kda_forward(
x: &[f32],
w: &KdaWeights,
cfg: &KdaCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (nh, dk, dv) = (cfg.num_heads, cfg.head_k_dim, cfg.head_v_dim);
let (kd, vd) = (nh * dk, nh * dv);
let mut xq = vec![0.0f32; kd];
let mut xk = vec![0.0f32; kd];
let mut xv = vec![0.0f32; vd];
let mut fl = vec![0.0f32; w.f_a.rows()];
let mut b = vec![0.0f32; nh];
QTensor::matvec_many(
[&w.q_proj, &w.k_proj, &w.v_proj, &w.f_a],
x,
[
xq.as_mut_slice(),
xk.as_mut_slice(),
xv.as_mut_slice(),
fl.as_mut_slice(),
],
pool,
);
w.b_proj.matvec(x, &mut b, pool);
let mut f = vec![0.0f32; kd];
w.f_b.matvec(&fl, &mut f, pool);
let gate_out = kda_gate_out(w, x, vd, pool);
let mut of = vec![0.0f32; vd];
kda_step(&xq, &xk, &xv, &f, &b, &gate_out, w, cfg, state, &mut of, pool);
let mut out = vec![0.0f32; cfg.hidden_size];
w.o_proj.matvec(&of, &mut out, pool);
out
}
pub fn kda_forward_batch(
xs: &[f32],
bsz: usize,
w: &KdaWeights,
cfg: &KdaCfg,
state: &mut Vec<f32>,
pool: Option<&Pool>,
) -> Vec<f32> {
if state.len() != cfg.state_len() {
*state = vec![0f32; cfg.state_len()];
}
let (nh, dk, dv, hs) = (
cfg.num_heads,
cfg.head_k_dim,
cfg.head_v_dim,
cfg.hidden_size,
);
let (kd, vd) = (nh * dk, nh * dv);
let mut xq = vec![0.0f32; bsz * kd];
w.q_proj.matmat(xs, bsz, &mut xq, pool);
let mut xk = vec![0.0f32; bsz * kd];
w.k_proj.matmat(xs, bsz, &mut xk, pool);
let mut xv = vec![0.0f32; bsz * vd];
w.v_proj.matmat(xs, bsz, &mut xv, pool);
let rank = w.f_a.rows();
let mut fl = vec![0.0f32; bsz * rank];
w.f_a.matmat(xs, bsz, &mut fl, pool);
let mut f = vec![0.0f32; bsz * kd];
w.f_b.matmat(&fl, bsz, &mut f, pool);
let mut b = vec![0.0f32; bsz * nh];
w.b_proj.matmat(xs, bsz, &mut b, pool);
let mut gate_out = vec![0.0f32; bsz * vd];
match &w.gate {
KdaOutGate::Full(gp) => gp.matmat(xs, bsz, &mut gate_out, pool),
KdaOutGate::LowRank(ga, gb) => {
let mut low = vec![0.0f32; bsz * ga.rows()];
ga.matmat(xs, bsz, &mut low, pool);
gb.matmat(&low, bsz, &mut gate_out, pool);
}
}
let mut of = vec![0.0f32; bsz * vd];
for bi in 0..bsz {
let mut oh = vec![0.0f32; vd];
kda_step(
&xq[bi * kd..(bi + 1) * kd],
&xk[bi * kd..(bi + 1) * kd],
&xv[bi * vd..(bi + 1) * vd],
&f[bi * kd..(bi + 1) * kd],
&b[bi * nh..(bi + 1) * nh],
&gate_out[bi * vd..(bi + 1) * vd],
w,
cfg,
state,
&mut oh,
pool,
);
of[bi * vd..(bi + 1) * vd].copy_from_slice(&oh);
}
let mut out = vec![0.0f32; bsz * hs];
w.o_proj.matmat(&of, bsz, &mut out, pool);
out
}
#[cfg(test)]
mod tests {
#[test]
fn kda_forward_matches_naive_reference() {
let (nh, dk, dv, kk, hs, rank) = (2usize, 4usize, 4usize, 3usize, 6usize, 3usize);
let synth = |rows: usize, cols: usize, salt: usize| -> QTensor {
QTensor::from_f32(
(0..rows * cols)
.map(|i| (((i * 31 + salt * 17) % 101) as f32 / 101.0 - 0.5) * 0.6)
.collect(),
rows,
cols,
)
};
let vecf = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 13 + salt * 7) % 89) as f32 / 89.0 - 0.5) * 0.8)
.collect()
};
for (label, a_log, lb) in [
("per-head standard", vecf(nh, 40), None),
("per-dim lower-bound", vecf(dk, 41), Some(-5.0f32)),
] {
let w = KdaWeights {
q_proj: synth(nh * dk, hs, 1),
k_proj: synth(nh * dk, hs, 2),
v_proj: synth(nh * dv, hs, 3),
conv_q: vecf(nh * dk * kk, 4),
conv_k: vecf(nh * dk * kk, 5),
conv_v: vecf(nh * dv * kk, 6),
f_a: synth(rank, hs, 7),
f_b: synth(nh * dk, rank, 8),
dt_bias: vecf(nh * dk, 9),
a_log: a_log.clone(),
b_proj: synth(nh, hs, 10),
gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
o_proj: synth(hs, nh * dv, 13),
gate_lower_bound: lb,
};
let cfg = KdaCfg {
num_heads: nh,
head_k_dim: dk,
head_v_dim: dv,
conv_kernel: kk,
hidden_size: hs,
rms_eps: 1e-6,
};
let xs: Vec<Vec<f32>> = (0..6)
.map(|t| (0..hs).map(|i| ((t * hs + i) as f32 * 0.37).sin() * 0.5).collect())
.collect();
let mut state = Vec::new();
let got: Vec<Vec<f32>> = xs
.iter()
.map(|x| kda_forward(x, &w, &cfg, &mut state, None))
.collect();
let mv = |t: &QTensor, x: &[f32]| -> Vec<f32> {
let mut o = vec![0.0f32; t.rows()];
t.matvec(x, &mut o, None);
o
};
let mut hist: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = Vec::new(); let mut s_state = vec![0f64; nh * dk * dv];
let mut want: Vec<Vec<f32>> = Vec::new();
for x in &xs {
let (xq, xk, xv) = (mv(&w.q_proj, x), mv(&w.k_proj, x), mv(&w.v_proj, x));
hist.push((xq, xk, xv));
let conv = |sel: fn(&(Vec<f32>, Vec<f32>, Vec<f32>)) -> &Vec<f32>,
taps: &[f32],
n: usize|
-> Vec<f32> {
(0..n)
.map(|c| {
let t = &taps[c * kk..(c + 1) * kk];
let mut acc = 0f64;
for j in 0..kk {
let idx = hist.len() as i64 - (kk as i64 - j as i64);
if idx >= 0 {
acc += sel(&hist[idx as usize])[c] as f64 * t[j] as f64;
}
}
silu(acc)
})
.map(|v| v as f32)
.collect()
};
let cq = conv(|h| &h.0, &w.conv_q, nh * dk);
let ck = conv(|h| &h.1, &w.conv_k, nh * dk);
let cv = conv(|h| &h.2, &w.conv_v, nh * dv);
let f = mv(&w.f_b, &mv(&w.f_a, x));
let bb = mv(&w.b_proj, x);
let gate_out = match &w.gate {
KdaOutGate::LowRank(ga, gb) => mv(gb, &mv(ga, x)),
KdaOutGate::Full(g) => mv(g, x),
};
let mut of = vec![0f32; nh * dv];
for h in 0..nh {
let q: Vec<f64> = {
let sl = &cq[h * dk..(h + 1) * dk];
let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
let inv = 1.0 / ((n + 1e-6).sqrt() * (dk as f64).sqrt());
sl.iter().map(|&v| v as f64 * inv).collect()
};
let k: Vec<f64> = {
let sl = &ck[h * dk..(h + 1) * dk];
let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
let inv = 1.0 / (n + 1e-6).sqrt();
sl.iter().map(|&v| v as f64 * inv).collect()
};
let v: Vec<f64> = cv[h * dv..(h + 1) * dv].iter().map(|&v| v as f64).collect();
let g: Vec<f64> = (0..dk)
.map(|d| {
let a = if w.a_log.len() == nh {
w.a_log[h] as f64
} else {
w.a_log[d] as f64
};
let raw = f[h * dk + d] as f64 + w.dt_bias[h * dk + d] as f64;
match w.gate_lower_bound {
Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
None => -a.exp() * softplus(raw),
}
})
.collect();
let beta = sigmoid(bb[h] as f64);
let s = &mut s_state[h * dk * dv..(h + 1) * dk * dv];
for di in 0..dk {
for dj in 0..dv {
s[di * dv + dj] *= g[di].exp();
}
}
let mut kv = vec![0f64; dv];
for di in 0..dk {
for dj in 0..dv {
kv[dj] += k[di] * s[di * dv + dj];
}
}
for di in 0..dk {
for dj in 0..dv {
s[di * dv + dj] += beta * k[di] * (v[dj] - kv[dj]);
}
}
let mut o = vec![0f64; dv];
for di in 0..dk {
for dj in 0..dv {
o[dj] += q[di] * s[di * dv + dj];
}
}
let ss: f64 = o.iter().map(|&v| v * v).sum();
let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
for dj in 0..dv {
of[h * dv + dj] = (o[dj] * inv
* w.o_norm[dj] as f64
* sigmoid(gate_out[h * dv + dj] as f64))
as f32;
}
}
want.push(mv(&w.o_proj, &of));
}
for (t, (g, e)) in got.iter().zip(&want).enumerate() {
for (i, (a, b)) in g.iter().zip(e.iter()).enumerate() {
assert!(
(a - b).abs() < 2e-4,
"{label}: t={t} i={i}: {a} vs {b}"
);
}
}
}
let w = KdaWeights {
q_proj: synth(nh * dk, hs, 1),
k_proj: synth(nh * dk, hs, 2),
v_proj: synth(nh * dv, hs, 3),
conv_q: vecf(nh * dk * kk, 4),
conv_k: vecf(nh * dk * kk, 5),
conv_v: vecf(nh * dv * kk, 6),
f_a: synth(rank, hs, 7),
f_b: synth(nh * dk, rank, 8),
dt_bias: vecf(nh * dk, 9),
a_log: vecf(nh, 40),
b_proj: synth(nh, hs, 10),
gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
o_proj: synth(hs, nh * dv, 13),
gate_lower_bound: None,
};
let cfg = KdaCfg {
num_heads: nh,
head_k_dim: dk,
head_v_dim: dv,
conv_kernel: kk,
hidden_size: hs,
rms_eps: 1e-6,
};
let xs: Vec<f32> = (0..5 * hs).map(|i| (i as f32 * 0.29).cos() * 0.4).collect();
let mut st1 = Vec::new();
let seq: Vec<f32> = (0..5)
.flat_map(|t| kda_forward(&xs[t * hs..(t + 1) * hs], &w, &cfg, &mut st1, None))
.collect();
let mut st2 = Vec::new();
let bat = kda_forward_batch(&xs, 5, &w, &cfg, &mut st2, None);
for (i, (a, b)) in seq.iter().zip(&bat).enumerate() {
assert!((a - b).abs() < 1e-5, "batch i={i}: {a} vs {b}");
}
assert_eq!(st1, st2, "state must match after the chunk");
}
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");
}
}
fn tiny_short_conv() -> (ShortConvWeights, ShortConvCfg) {
let cfg = ShortConvCfg {
hidden_size: 8,
kernel: 3,
};
let synth = |rows: usize, cols: usize, salt: usize| {
QTensor::from_f32(
(0..rows * cols)
.map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.5)
.collect(),
rows,
cols,
)
};
let w = ShortConvWeights {
in_proj: synth(3 * cfg.hidden_size, cfg.hidden_size, 1),
conv: (0..cfg.hidden_size * cfg.kernel)
.map(|i| ((i * 7 % 13) as f32 / 13.0 - 0.5) * 0.8)
.collect(),
out_proj: synth(cfg.hidden_size, cfg.hidden_size, 2),
};
(w, cfg)
}
#[test]
fn short_conv_ring_matches_explicit_causal_conv() {
let (w, cfg) = tiny_short_conv();
let seq: Vec<Vec<f32>> = (0..6)
.map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.19).cos()).collect())
.collect();
let mut s_inc = Vec::new();
for (t, x) in seq.iter().enumerate() {
let inc = short_conv_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 = short_conv_forward(xr, &w, &cfg, &mut s_replay, None);
}
assert_eq!(inc, replay, "position {t}: ring must equal replay");
assert_eq!(s_inc.len(), cfg.state_len());
}
}
#[test]
fn short_conv_batch_matches_sequential() {
let (w, cfg) = tiny_short_conv();
let b = 5;
let xs: Vec<f32> = (0..b * cfg.hidden_size)
.map(|i| (i as f32 * 0.13).sin() * 0.6)
.collect();
let mut s_seq = Vec::new();
let mut seq_out = vec![0.0f32; b * cfg.hidden_size];
for bi in 0..b {
let o = short_conv_forward(
&xs[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size],
&w,
&cfg,
&mut s_seq,
None,
);
seq_out[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size].copy_from_slice(&o);
}
let mut s_batch = Vec::new();
let batch_out = short_conv_forward_batch(&xs, b, &w, &cfg, &mut s_batch, None);
assert_eq!(
seq_out, batch_out,
"batch conv must match sequential decode"
);
assert_eq!(s_seq, s_batch, "ring state must match after the chunk");
}
}