use cortiq_engine::fcd_ops as ops;
use cortiq_engine::nystrom::{NystromState, O1Rect};
fn synth(n: usize, salt: u64) -> Vec<f64> {
(0..n)
.map(|i| {
let x = (i as u64)
.wrapping_mul(6364136223846793005)
.wrapping_add(salt.wrapping_mul(1442695040888963407) ^ 0x9E3779B97F4A7C15);
let x = (x ^ (x >> 31)).wrapping_mul(0xBF58476D1CE4E5B9);
((x >> 11) as f64 / (1u64 << 53) as f64) - 0.5
})
.collect()
}
fn compare(t: usize, d: usize, dv: usize, p: usize, m: usize, w: usize, sink: usize) -> (f64, f64) {
let cfg = ops::NysCfg { m, w, sink, prefill: Some(p) };
let mut q = synth(t * d, 20);
let mut k = synth(t * d, 21);
let v = synth(t * dv, 22);
for x in q.iter_mut().chain(k.iter_mut()) {
*x *= 2.0;
}
let mut want = vec![0f64; t * dv];
ops::nystrom_head_fwd(&q, &k, &v, t, d, dv, &cfg, &mut want);
let f32s = |x: &[f64]| x.iter().map(|&a| a as f32).collect::<Vec<f32>>();
let (q32, k32, v32) = (f32s(&q), f32s(&k), f32s(&v));
let mut st = NystromState::new(m, w, sink).with_rect(O1Rect::Aggregate);
st.prefill(&q32[..p * d], &k32[..p * d], &v32[..p * dv], p, d, dv);
let mut got = vec![0.0f32; dv];
let (mut max_diff, mut max_out) = (0.0f64, 0.0f64);
for i in p..t {
st.step(
&q32[i * d..(i + 1) * d],
&k32[i * d..(i + 1) * d],
&v32[i * dv..(i + 1) * dv],
&mut got,
);
for c in 0..dv {
let expect = want[i * dv + c];
max_diff = max_diff.max((expect - got[c] as f64).abs());
max_out = max_out.max(expect.abs());
}
}
(max_diff, max_out)
}
#[test]
fn trainer_forward_matches_runtime_kernel() {
for &(t, p, m, w, sink) in &[
(160usize, 80usize, 8usize, 32usize, 4usize),
(160, 80, 16, 32, 0), (200, 100, 8, 48, 4), (120, 60, 4, 16, 2), ] {
let (d, dv) = (6usize, 5usize);
assert!(p > w + sink + 8, "fixture must seal a real skeleton");
let (diff, out) = compare(t, d, dv, p, m, w, sink);
println!(
"t={t} p={p} m={m} w={w} sink={sink}: max|trainer-runtime| = {diff:.3e} \
(max|out| {out:.3e})"
);
assert!(
diff < 2e-5,
"train/serve skew: t={t} p={p} m={m} w={w} sink={sink} max diff {diff:.3e}"
);
}
}
#[test]
fn aggregate_guard_differs_from_per_key_clamp() {
let (t, d, dv, p, m, w, sink) = (160usize, 6usize, 5usize, 80usize, 8usize, 32usize, 4usize);
let cfg = ops::NysCfg { m, w, sink, prefill: Some(p) };
let mut q = synth(t * d, 20);
let mut k = synth(t * d, 21);
let v = synth(t * dv, 22);
for x in q.iter_mut().chain(k.iter_mut()) {
*x *= 2.0;
}
let mut out = vec![0f64; t * dv];
ops::nystrom_head_fwd(&q, &k, &v, t, d, dv, &cfg, &mut out);
let f32s = |x: &[f64]| x.iter().map(|&a| a as f32).collect::<Vec<f32>>();
let (q32, k32, v32) = (f32s(&q), f32s(&k), f32s(&v));
let mut st = NystromState::new(m, w, sink).with_rect(O1Rect::Fm);
st.prefill(&q32[..p * d], &k32[..p * d], &v32[..p * dv], p, d, dv);
let mut got = vec![0.0f32; dv];
let mut max_diff = 0.0f64;
for i in p..t {
st.step(
&q32[i * d..(i + 1) * d],
&k32[i * d..(i + 1) * d],
&v32[i * dv..(i + 1) * dv],
&mut got,
);
for c in 0..dv {
max_diff = max_diff.max((out[i * dv + c] - got[c] as f64).abs());
}
}
println!("aggregate vs per-key(Fm) rectifier: max diff {max_diff:.3e}");
assert!(
max_diff > 1e-3,
"fixture no longer exercises negative far mass ({max_diff:.3e}): the parity \
test above would pass under a per-key clamp too"
);
}
#[test]
fn short_prompt_falls_back_to_exact_both_sides() {
let (t, d, dv, p, m, w, sink) = (160usize, 6usize, 5usize, 20usize, 8usize, 32usize, 4usize);
assert!(p <= w + sink + 8, "fixture must be in the degenerate regime");
let cfg = ops::NysCfg { m, w, sink, prefill: Some(p) };
let q = synth(t * d, 30);
let k = synth(t * d, 31);
let v = synth(t * dv, 32);
let mut got = vec![0f64; t * dv];
ops::nystrom_head_fwd(&q, &k, &v, t, d, dv, &cfg, &mut got);
let scale = 1.0 / (d as f64).sqrt();
for ti in 0..t {
let mut lg = vec![0f64; ti + 1];
let mut c = f64::NEG_INFINITY;
for (j, l) in lg.iter_mut().enumerate() {
*l = (0..d).map(|x| q[ti * d + x] * k[j * d + x]).sum::<f64>() * scale;
c = c.max(*l);
}
let mut den = 0f64;
let mut acc = vec![0f64; dv];
for (j, &l) in lg.iter().enumerate() {
let e = (l - c).exp();
den += e;
for x in 0..dv {
acc[x] += e * v[j * dv + x];
}
}
for x in 0..dv {
let want = acc[x] / den;
assert!(
(want - got[ti * dv + x]).abs() < 1e-12,
"row {ti}: short-prompt window must be exact attention"
);
}
}
}