#[inline]
fn idx(row: usize, col: usize, cols: usize) -> usize {
row * cols + col
}
pub fn l2_normalize_rows(x: &[f32], rows: usize, cols: usize, eps: f32) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
let row = &x[r * cols..(r + 1) * cols];
let sum_sq: f32 = row.iter().map(|v| v * v).sum();
if !sum_sq.is_finite() {
continue;
}
let denom = (sum_sq + eps).sqrt();
for c in 0..cols {
out[idx(r, c, cols)] = row[c] / denom;
}
}
out
}
pub fn solve_unit_lower(a: &[f32], b: &[f32], c: usize, d: usize) -> Vec<f32> {
let mut x = b.to_vec();
for i in 1..c {
for col in 0..d {
let mut acc = 0.0f32;
for (j, xj) in x.chunks_exact(d).enumerate().take(i) {
acc += a[idx(i, j, c)] * xj[col];
}
x[idx(i, col, d)] -= acc;
}
}
x
}
#[allow(clippy::too_many_arguments)]
pub fn sequential_gdn(
q: &[f32],
k: &[f32],
v: &[f32],
beta: &[f32],
alpha: &[f32],
h0: &[f32],
t: usize,
dk: usize,
dv: usize,
apply_scale: bool,
) -> (Vec<f32>, Vec<f32>) {
let scale = if apply_scale {
1.0 / (dk as f32).sqrt()
} else {
1.0
};
let mut h = h0.to_vec(); let mut out = vec![0.0f32; t * dv];
for i in 0..t {
let qi = &q[i * dk..(i + 1) * dk];
let ki = &k[i * dk..(i + 1) * dk];
let vi = &v[i * dv..(i + 1) * dv];
let a = alpha[i];
let b = beta[i];
let mut h_decayed = vec![0.0f32; dv * dk];
for (hd, hv) in h_decayed.iter_mut().zip(h.iter()) {
*hd = a * hv;
}
let mut kv = vec![0.0f32; dv];
for row in 0..dv {
let hrow = &h_decayed[row * dk..(row + 1) * dk];
kv[row] = hrow.iter().zip(ki).map(|(hv, kv_)| hv * kv_).sum();
}
let mut r = vec![0.0f32; dv];
for row in 0..dv {
r[row] = b * (vi[row] - kv[row]);
}
for row in 0..dv {
for col in 0..dk {
h_decayed[idx(row, col, dk)] += r[row] * ki[col];
}
}
h = h_decayed;
let out_row = &mut out[i * dv..(i + 1) * dv];
for (row, out_val) in out_row.iter_mut().enumerate() {
let hrow = &h[row * dk..(row + 1) * dk];
let dot: f32 = hrow.iter().zip(qi).map(|(hv, qv)| hv * qv).sum();
*out_val = dot * scale;
}
}
(out, h)
}
#[allow(clippy::too_many_arguments)]
pub fn chunkwise_gdn(
q: &[f32],
k: &[f32],
v: &[f32],
beta: &[f32],
alpha: &[f32],
h0: &[f32],
t_total: usize,
dk: usize,
dv: usize,
chunk_size: usize,
apply_scale: bool,
) -> (Vec<f32>, Vec<f32>) {
assert!(chunk_size > 0, "chunk_size must be positive");
let scale = if apply_scale {
1.0 / (dk as f32).sqrt()
} else {
1.0
};
let mut h = h0.to_vec(); let mut out = vec![0.0f32; t_total * dv];
let mut start = 0usize;
while start < t_total {
let end = (start + chunk_size).min(t_total);
let c = end - start;
let q_c = &q[start * dk..end * dk];
let k_c = &k[start * dk..end * dk];
let v_c = &v[start * dv..end * dv];
let b_c = &beta[start..end];
let a_c = &alpha[start..end];
let mut gamma_log = vec![0.0f32; c];
let mut running_log = 0.0f32;
for (i, gl) in gamma_log.iter_mut().enumerate() {
running_log += a_c[i].max(f32::MIN_POSITIVE).ln();
*gl = running_log;
}
let gamma_log_end = gamma_log[c - 1];
let mut g = vec![0.0f32; c * c];
for i in 0..c {
let ki = &k_c[i * dk..(i + 1) * dk];
for j in 0..c {
let kj = &k_c[j * dk..(j + 1) * dk];
g[idx(i, j, c)] = ki.iter().zip(kj).map(|(a, b)| a * b).sum();
}
}
let mut a_plain = vec![0.0f32; c * c];
for i in 0..c {
a_plain[idx(i, i, c)] = 1.0;
for j in 0..i {
a_plain[idx(i, j, c)] = b_c[i] * g[idx(i, j, c)];
}
}
let mut rhs_k = vec![0.0f32; c * dk];
for i in 0..c {
for col in 0..dk {
rhs_k[idx(i, col, dk)] = b_c[i] * k_c[idx(i, col, dk)];
}
}
let mut rhs_v = vec![0.0f32; c * dv];
for i in 0..c {
for col in 0..dv {
rhs_v[idx(i, col, dv)] = b_c[i] * v_c[idx(i, col, dv)];
}
}
let w = solve_unit_lower(&a_plain, &rhs_k, c, dk);
let mut a_gamma = vec![0.0f32; c * c];
for i in 0..c {
a_gamma[idx(i, i, c)] = 1.0;
for j in 0..i {
a_gamma[idx(i, j, c)] =
b_c[i] * g[idx(i, j, c)] * (gamma_log[i] - gamma_log[j]).exp();
}
}
let u = solve_unit_lower(&a_gamma, &rhs_v, c, dv);
let mut w_h0 = vec![0.0f32; c * dv];
for i in 0..c {
let wi = &w[i * dk..(i + 1) * dk];
for row in 0..dv {
let hrow = &h[row * dk..(row + 1) * dk];
w_h0[idx(i, row, dv)] = wi.iter().zip(hrow).map(|(a, b)| a * b).sum();
}
}
let mut r = vec![0.0f32; c * dv];
for i in 0..c {
let gamma_i = gamma_log[i].exp();
for col in 0..dv {
r[idx(i, col, dv)] = u[idx(i, col, dv)] - gamma_i * w_h0[idx(i, col, dv)];
}
}
let mut q_h0 = vec![0.0f32; c * dv];
for i in 0..c {
let qi = &q_c[i * dk..(i + 1) * dk];
for row in 0..dv {
let hrow = &h[row * dk..(row + 1) * dk];
q_h0[idx(i, row, dv)] = qi.iter().zip(hrow).map(|(a, b)| a * b).sum();
}
}
let mut qk = vec![0.0f32; c * c];
for i in 0..c {
let qi = &q_c[i * dk..(i + 1) * dk];
for j in 0..c {
let kj = &k_c[j * dk..(j + 1) * dk];
qk[idx(i, j, c)] = qi.iter().zip(kj).map(|(a, b)| a * b).sum();
}
}
let mut lqk = vec![0.0f32; c * c];
for i in 0..c {
for j in 0..=i {
lqk[idx(i, j, c)] = qk[idx(i, j, c)] * (gamma_log[i] - gamma_log[j]).exp();
}
}
for i in 0..c {
let mut lqk_r_row = vec![0.0f32; dv];
for j in 0..=i {
let lij = lqk[idx(i, j, c)];
for col in 0..dv {
lqk_r_row[col] += lij * r[idx(j, col, dv)];
}
}
let gamma_i = gamma_log[i].exp();
for col in 0..dv {
let val = gamma_i * q_h0[idx(i, col, dv)] + lqk_r_row[col];
out[(start + i) * dv + col] = val * scale;
}
}
let mut k_right = vec![0.0f32; c * dk];
for i in 0..c {
let scale_i = (gamma_log_end - gamma_log[i]).exp();
for col in 0..dk {
k_right[idx(i, col, dk)] = scale_i * k_c[idx(i, col, dk)];
}
}
let gamma_end = gamma_log_end.exp();
let mut h_next = vec![0.0f32; dv * dk];
for (hn, hv) in h_next.iter_mut().zip(h.iter()) {
*hn = gamma_end * hv;
}
for i in 0..c {
for row in 0..dv {
let rv = r[idx(i, row, dv)];
for col in 0..dk {
h_next[idx(row, col, dk)] += rv * k_right[idx(i, col, dk)];
}
}
}
h = h_next;
start = end;
}
(out, h)
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
use std::path::PathBuf;
#[derive(Deserialize)]
struct Fixture {
length: usize,
dk: usize,
dv: usize,
chunk_size: usize,
q: Vec<Vec<f32>>,
k: Vec<Vec<f32>>,
v: Vec<Vec<f32>>,
beta: Vec<f32>,
alpha: Vec<f32>,
h0: Vec<Vec<f32>>,
out_seq: Vec<Vec<f32>>,
h_seq: Vec<Vec<f32>>,
out_chk: Vec<Vec<f32>>,
h_chk: Vec<Vec<f32>>,
}
fn flatten(rows: &[Vec<f32>]) -> Vec<f32> {
rows.iter().flat_map(|r| r.iter().copied()).collect()
}
fn fixtures_dir() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("gdn_chunk")
}
fn load_fixture(name: &str) -> Fixture {
let path = fixtures_dir().join(name);
let data = std::fs::read_to_string(&path)
.unwrap_or_else(|e| panic!("failed to read {}: {e}", path.display()));
serde_json::from_str(&data)
.unwrap_or_else(|e| panic!("bad JSON in {}: {e}", path.display()))
}
fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "length mismatch in max_abs_diff");
let mut max = 0.0f32;
for (x, y) in a.iter().zip(b.iter()) {
let d = (x - y).abs();
if !d.is_finite() {
return d;
}
if d > max {
max = d;
}
}
max
}
fn count_nonfinite(v: &[f32]) -> usize {
v.iter().filter(|x| !x.is_finite()).count()
}
const TOL: f32 = 1.0e-5;
fn run_case(fixture_name: &str) {
let fx = load_fixture(fixture_name);
let q = flatten(&fx.q);
let k = flatten(&fx.k);
let v = flatten(&fx.v);
let h0 = flatten(&fx.h0);
let expected_out_seq = flatten(&fx.out_seq);
let expected_h_seq = flatten(&fx.h_seq);
let expected_out_chk = flatten(&fx.out_chk);
let expected_h_chk = flatten(&fx.h_chk);
let (rust_out_seq, rust_h_seq) = sequential_gdn(
&q, &k, &v, &fx.beta, &fx.alpha, &h0, fx.length, fx.dk, fx.dv, true,
);
let (rust_out_chk, rust_h_chk) = chunkwise_gdn(
&q,
&k,
&v,
&fx.beta,
&fx.alpha,
&h0,
fx.length,
fx.dk,
fx.dv,
fx.chunk_size,
true,
);
for (label, buf) in [
("rust_out_seq", &rust_out_seq),
("rust_h_seq", &rust_h_seq),
("rust_out_chk", &rust_out_chk),
("rust_h_chk", &rust_h_chk),
] {
let n = count_nonfinite(buf);
assert_eq!(
n, 0,
"{fixture_name}: {label} contains {n} non-finite (NaN/Inf) values — \
a decay-underflow regression the max-fold gate would silently pass"
);
}
let d_seq_out = max_abs_diff(&rust_out_seq, &expected_out_seq);
let d_seq_h = max_abs_diff(&rust_h_seq, &expected_h_seq);
let d_chk_out = max_abs_diff(&rust_out_chk, &expected_out_chk);
let d_chk_h = max_abs_diff(&rust_h_chk, &expected_h_chk);
let d_internal_out = max_abs_diff(&rust_out_chk, &rust_out_seq);
let d_internal_h = max_abs_diff(&rust_h_chk, &rust_h_seq);
println!(
"gdn_chunk_ref fixture={fixture_name} \
rust_seq_vs_numpy_seq(out={d_seq_out:.3e}, h={d_seq_h:.3e}) \
rust_chk_vs_numpy_chk(out={d_chk_out:.3e}, h={d_chk_h:.3e}) \
rust_chk_vs_rust_seq(out={d_internal_out:.3e}, h={d_internal_h:.3e})"
);
assert!(
d_seq_out <= TOL,
"{fixture_name}: rust sequential_gdn out diverged from NumPy: max_abs={d_seq_out}"
);
assert!(
d_seq_h <= TOL,
"{fixture_name}: rust sequential_gdn h diverged from NumPy: max_abs={d_seq_h}"
);
assert!(
d_chk_out <= TOL,
"{fixture_name}: rust chunkwise_gdn out diverged from NumPy: max_abs={d_chk_out}"
);
assert!(
d_chk_h <= TOL,
"{fixture_name}: rust chunkwise_gdn h diverged from NumPy: max_abs={d_chk_h}"
);
assert!(
d_internal_out <= TOL,
"{fixture_name}: rust chunkwise_gdn out diverged from rust sequential_gdn (the \
equivalence gate the Metal kernels rely on): max_abs={d_internal_out}"
);
assert!(
d_internal_h <= TOL,
"{fixture_name}: rust chunkwise_gdn h diverged from rust sequential_gdn (the \
equivalence gate the Metal kernels rely on): max_abs={d_internal_h}"
);
}
#[test]
fn gdn_chunk_ref_parity_seed7_chunk64() {
run_case("case_seed7_len191_chunk64.json");
}
#[test]
fn gdn_chunk_ref_parity_seed11_chunk128() {
run_case("case_seed11_len384_chunk128.json");
}
#[test]
fn gdn_chunk_ref_parity_seed23_uneven_tail() {
run_case("case_seed23_len130_chunk64.json");
}
#[test]
fn gdn_chunk_ref_parity_seed31_strong_decay() {
run_case("case_seed31_len160_chunk64_strongdecay.json");
}
#[test]
fn max_abs_diff_is_nan_and_inf_honest() {
let clean = [1.0f32, 2.0, 3.0];
let nan_side = [1.0f32, f32::NAN, 3.0];
let d = max_abs_diff(&clean, &nan_side);
assert!(
!d.is_finite(),
"max_abs_diff silently dropped a NaN operand (got {d}); the gate is blind"
);
let inf_side = [1.0f32, f32::INFINITY, 3.0];
let di = max_abs_diff(&clean, &inf_side);
assert!(
!di.is_finite(),
"max_abs_diff silently dropped an Inf operand (got {di}); the gate is blind"
);
}
#[test]
fn seed31_fixture_triggers_linear_gamma_underflow() {
let fx = load_fixture("case_seed31_len160_chunk64_strongdecay.json");
let c = fx.chunk_size;
let mut gamma = 1.0f32;
let mut underflowed = false;
for i in 0..c {
gamma *= fx.alpha[i];
if gamma == 0.0 {
underflowed = true;
break;
}
}
assert!(
underflowed,
"seed31 strong-decay fixture no longer underflows the linear cumprod within a \
chunk; it no longer guards the gamma-underflow regression this PR fixes"
);
let gamma_inv = 1.0f32 / gamma;
assert!(
!gamma_inv.is_finite(),
"expected 1/0 = inf, got {gamma_inv}"
);
assert!(
(0.0f32 * gamma_inv).is_nan(),
"expected the pre-fix decay factor 0*inf to be NaN"
);
}
#[test]
fn l2_normalize_rows_fail_closed_table() {
let cols = 3usize;
let cases: &[(&str, &[f32], bool)] = &[
("nan_lane", &[f32::NAN, 1.0, 2.0], true),
("pos_inf_lane", &[f32::INFINITY, 1.0, 2.0], true),
("neg_inf_lane", &[f32::NEG_INFINITY, 1.0, 2.0], true),
("all_zero", &[0.0, 0.0, 0.0], true),
];
for (label, row, expect_all_zero) in cases {
let out = l2_normalize_rows(row, 1, cols, 1e-6);
assert!(
out.iter().all(|x| x.is_finite()),
"case {label}: l2_normalize_rows output must be fully finite, got {out:?}"
);
if *expect_all_zero {
assert!(
out.iter().all(|x| *x == 0.0),
"case {label}: expected the whole row zeroed, got {out:?}"
);
}
}
let ordinary = l2_normalize_rows(&[3.0f32, 4.0], 1, 2, 0.0);
assert!((ordinary[0] - 0.6).abs() < 1e-6);
assert!((ordinary[1] - 0.8).abs() < 1e-6);
let multi = l2_normalize_rows(&[f32::NAN, 0.0, 0.0, 3.0, 4.0, 0.0], 2, 3, 0.0);
assert!(
multi[0..3].iter().all(|x| *x == 0.0),
"row 0 (poisoned) must be all-zero, got {:?}",
&multi[0..3]
);
assert!(
(multi[3] - 0.6).abs() < 1e-6 && (multi[4] - 0.8).abs() < 1e-6,
"row 1 (clean) must be unaffected by row 0's poisoning, got {:?}",
&multi[3..6]
);
}
}