use super::codes::{SparseCode, solve_row_codes};
use super::scoring::{ScoreRouteStats, TileScorer};
use super::{SparseDictConfig, SparseDictFit};
use ndarray::{Array2, ArrayView2, Axis};
use rayon::prelude::*;
use std::collections::HashMap;
pub(super) fn route_and_code_all(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
scorer: &TileScorer,
s: usize,
code_ridge: f32,
minibatch: usize,
score_mode: gam_gpu::GpuMode,
mut score_route_stats: Option<&mut ScoreRouteStats>,
) -> Result<Vec<SparseCode>, String> {
let n = x.nrows();
let batch = minibatch.max(1);
let mut codes: Vec<SparseCode> = Vec::with_capacity(n);
let mut start = 0usize;
while start < n {
let end = (start + batch).min(n);
let block = x.slice(ndarray::s![start..end, ..]);
let routed = scorer.route_minibatch_with_mode(block, decoder, score_mode)?;
if let Some(stats) = score_route_stats.as_deref_mut() {
stats.record_result(&routed);
}
let active_lists = routed.selections;
let mut block_codes: Vec<SparseCode> = block
.axis_iter(Axis(0))
.into_par_iter()
.zip(active_lists.into_par_iter())
.map(|(row, active)| solve_row_codes(row, decoder, &active, s, code_ridge))
.collect();
codes.append(&mut block_codes);
start = end;
}
Ok(codes)
}
pub(super) fn run(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
) -> Result<SparseDictFit, String> {
validate(x, config)?;
let n = x.nrows();
let p = x.ncols();
let k = config.n_atoms;
let s = config.active.min(k).max(1);
let mut decoder = seed_decoder(x, k);
unit_norm_rows(&mut decoder);
let scorer = TileScorer::new(s, config.score_tile);
let mut score_route_stats = ScoreRouteStats::default();
let mut pending_eq = DecoderNormalEq::zeros(k, p);
let mut prev_ev = f64::NEG_INFINITY;
let mut converged = false;
let mut epochs_run = 0usize;
let mut decoder_solve_stats = DecoderSolveStats::default();
let mut codes = route_and_code_all(
x,
decoder.view(),
&scorer,
s,
config.code_ridge,
config.minibatch,
config.score_mode,
Some(&mut score_route_stats),
)?;
for epoch in 0..config.max_epochs {
epochs_run = epoch + 1;
pending_eq.accumulate(x, &codes);
let sigma = residual_scale(x, &codes, decoder.view());
let (stats, gate) = solve_decoder_with_routability_gate(
&mut decoder,
&pending_eq,
config.decoder_ridge as f64,
sigma,
);
decoder_solve_stats = stats;
pending_eq.clear_refreshed_atoms(&gate);
unit_norm_rows(&mut decoder);
let revived = revive_dead_atoms(x, &codes, &mut decoder);
if revived > 0 {
unit_norm_rows(&mut decoder);
}
codes = route_and_code_all(
x,
decoder.view(),
&scorer,
s,
config.code_ridge,
config.minibatch,
config.score_mode,
Some(&mut score_route_stats),
)?;
let ev = explained_variance(x, &codes, decoder.view());
let improve = ev - prev_ev;
if revived == 0 && improve.abs() <= config.tolerance && epoch > 0 {
converged = true;
break;
}
prev_ev = ev;
}
let final_ev = explained_variance(x, &codes, decoder.view());
let (indices, code_mat) = pack_codes(&codes, n, s);
Ok(SparseDictFit {
decoder,
indices,
codes: code_mat,
explained_variance: final_ev,
epochs: epochs_run,
converged,
active: s,
score_route_stats,
decoder_solve_stats,
})
}
fn validate(x: ArrayView2<'_, f32>, config: &SparseDictConfig) -> Result<(), String> {
if x.nrows() == 0 || x.ncols() == 0 {
return Err("fit_sparse_dictionary requires a non-empty N×P matrix".to_string());
}
if !x.iter().all(|v| v.is_finite()) {
return Err("fit_sparse_dictionary input must be finite".to_string());
}
if config.n_atoms == 0 {
return Err("fit_sparse_dictionary requires K >= 1".to_string());
}
if config.active == 0 {
return Err("fit_sparse_dictionary requires active (top_s) >= 1".to_string());
}
if config.max_epochs == 0 {
return Err("fit_sparse_dictionary requires max_epochs >= 1".to_string());
}
if !(config.code_ridge.is_finite() && config.code_ridge >= 0.0) {
return Err("fit_sparse_dictionary code_ridge must be finite and non-negative".to_string());
}
if !(config.decoder_ridge.is_finite() && config.decoder_ridge >= 0.0) {
return Err(
"fit_sparse_dictionary decoder_ridge must be finite and non-negative".to_string(),
);
}
if !config.tolerance.is_finite() {
return Err("fit_sparse_dictionary tolerance must be finite".to_string());
}
Ok(())
}
pub(super) fn seed_decoder(x: ArrayView2<'_, f32>, k: usize) -> Array2<f32> {
let n = x.nrows();
let p = x.ncols();
let mut decoder = Array2::<f32>::zeros((k, p));
let mut first = 0usize;
let mut best = f32::NEG_INFINITY;
for i in 0..n {
let r = x.row(i);
let nrm: f32 = r.iter().map(|v| v * v).sum();
if nrm > best {
best = nrm;
first = i;
}
}
decoder.row_mut(0).assign(&x.row(first));
let mut min_dist2 = vec![f32::INFINITY; n];
for atom in 1..k {
let prev = decoder.row(atom - 1);
for i in 0..n {
let mut d2 = 0.0f32;
let xi = x.row(i);
for c in 0..p {
let d = xi[c] - prev[c];
d2 += d * d;
}
if d2 < min_dist2[i] {
min_dist2[i] = d2;
}
}
let chosen = if atom < n {
let mut bi = 0usize;
let mut bv = f32::NEG_INFINITY;
for i in 0..n {
if min_dist2[i] > bv {
bv = min_dist2[i];
bi = i;
}
}
bi
} else {
atom % n
};
decoder.row_mut(atom).assign(&x.row(chosen));
}
decoder
}
pub(super) struct DecoderNormalEq {
pub(super) diag: Vec<f64>,
pub(super) b: Array2<f64>,
pub(super) off: HashMap<(u32, u32), f64>,
pub(super) firings: Vec<usize>,
pub(super) amplitude_sum: Vec<f64>,
}
impl DecoderNormalEq {
pub(super) fn zeros(k: usize, p: usize) -> Self {
Self {
diag: vec![0.0f64; k],
b: Array2::<f64>::zeros((k, p)),
off: HashMap::new(),
firings: vec![0; k],
amplitude_sum: vec![0.0; k],
}
}
pub(super) fn accumulate(&mut self, x: ArrayView2<'_, f32>, codes: &[SparseCode]) {
let p = self.b.ncols();
for (row_idx, code) in codes.iter().enumerate() {
let xi = x.row(row_idx);
let xi_slice = xi.as_slice();
for a in 0..code.indices.len() {
let ca = code.codes[a] as f64;
if ca == 0.0 {
continue;
}
let ka = code.indices[a];
self.firings[ka as usize] += 1;
self.amplitude_sum[ka as usize] += ca.abs();
self.diag[ka as usize] += ca * ca;
let brow = ka as usize;
let mut brow_view = self.b.row_mut(brow);
match (brow_view.as_slice_mut(), xi_slice) {
(Some(bs), Some(xs)) => {
for (bref, &xv) in bs.iter_mut().zip(xs.iter()) {
*bref += ca * xv as f64;
}
}
_ => {
for c in 0..p {
brow_view[c] += ca * xi[c] as f64;
}
}
}
for bsel in (a + 1)..code.indices.len() {
let cb = code.codes[bsel] as f64;
if cb == 0.0 {
continue;
}
let kb = code.indices[bsel];
if ka == kb {
self.diag[ka as usize] += 2.0 * ca * cb;
continue;
}
let key = if ka < kb { (ka, kb) } else { (kb, ka) };
*self.off.entry(key).or_insert(0.0) += ca * cb;
}
}
}
}
pub(super) fn clear_refreshed_atoms(&mut self, gate: &[RoutabilityGateDecision]) {
for decision in gate.iter() {
if !decision.refresh {
continue;
}
let atom = decision.atom;
self.diag[atom] = 0.0;
self.firings[atom] = 0;
self.amplitude_sum[atom] = 0.0;
self.b.row_mut(atom).fill(0.0);
}
self.off
.retain(|&(a, b), _| !gate[a as usize].refresh && !gate[b as usize].refresh);
}
}
pub(super) const DEAD_DENOM: f64 = 1.0e-12;
const MAX_DIRECT_BLOCK: usize = 8;
#[derive(Clone, Copy, Debug)]
pub struct DecoderSolveStats {
pub mean_cofiring_degree: f64,
pub giant_component_fraction: f64,
pub component_count: usize,
pub max_component_size: usize,
pub cg_columns: usize,
pub cg_iterations: usize,
pub cg_kappa_hat: Option<f64>,
pub cg_relative_residual: f64,
pub cg_residual_stop: f64,
}
impl Default for DecoderSolveStats {
fn default() -> Self {
Self {
mean_cofiring_degree: 0.0,
giant_component_fraction: 0.0,
component_count: 0,
max_component_size: 0,
cg_columns: 0,
cg_iterations: 0,
cg_kappa_hat: None,
cg_relative_residual: 0.0,
cg_residual_stop: 0.0,
}
}
}
impl DecoderSolveStats {
fn record_cg(&mut self, result: &CgSolveResult) {
self.cg_columns += 1;
self.cg_iterations += result.iterations;
self.cg_relative_residual = self.cg_relative_residual.max(result.relative_residual);
if let Some(kappa) = result.kappa_hat {
self.cg_kappa_hat = Some(self.cg_kappa_hat.map_or(kappa, |old| old.max(kappa)));
}
}
}
#[derive(Clone, Copy, Debug)]
pub(super) struct RoutabilityGateDecision {
pub(super) atom: usize,
pub(super) refresh: bool,
pub(super) firings: usize,
pub(super) mean_amplitude: f64,
pub(super) z_alpha: f64,
pub(super) margin: f64,
pub(super) threshold: f64,
pub(super) standard_error: f64,
}
fn routability_z_alpha(firings: usize) -> f64 {
(firings.max(2) as f64).ln().sqrt()
}
pub(super) fn routability_gate_decisions(
eq: &DecoderNormalEq,
residual_scale: f64,
) -> Vec<RoutabilityGateDecision> {
(0..eq.diag.len())
.map(|atom| {
let firings = eq.firings[atom];
if firings == 0 || eq.diag[atom] <= DEAD_DENOM {
return RoutabilityGateDecision {
atom,
refresh: false,
firings,
mean_amplitude: 0.0,
z_alpha: routability_z_alpha(firings),
margin: 0.0,
threshold: f64::INFINITY,
standard_error: f64::INFINITY,
};
}
let n = firings as f64;
let mean_amplitude = eq.amplitude_sum[atom] / n;
let z_alpha = routability_z_alpha(firings);
let charge_floor = if residual_scale > 0.0 {
residual_scale * z_alpha / n.sqrt()
} else {
0.0
};
let margin = if mean_amplitude > 0.0 {
1.0 - charge_floor / mean_amplitude
} else {
f64::NEG_INFINITY
};
let standard_error = if residual_scale > 0.0 && mean_amplitude > 0.0 {
residual_scale / (mean_amplitude * n.sqrt())
} else if mean_amplitude > 0.0 {
0.0
} else {
f64::INFINITY
};
let threshold = if margin > 0.0 && mean_amplitude > 0.0 {
let denom = mean_amplitude * margin;
(z_alpha * residual_scale / denom).powi(2)
} else {
f64::INFINITY
};
RoutabilityGateDecision {
atom,
refresh: n >= threshold,
firings,
mean_amplitude,
z_alpha,
margin,
threshold,
standard_error,
}
})
.collect()
}
pub(super) fn solve_decoder_with_routability_gate(
decoder: &mut Array2<f32>,
eq: &DecoderNormalEq,
ridge: f64,
residual_scale: f64,
) -> (DecoderSolveStats, Vec<RoutabilityGateDecision>) {
let gate = routability_gate_decisions(eq, residual_scale);
let mut candidate = decoder.clone();
let stats = solve_decoder(&mut candidate, eq, ridge);
for decision in gate.iter() {
if !decision.refresh {
log::debug!(
"[SAE routability] atom {} deferred: firings={} mean_amplitude={:.4} \
z_alpha={:.4} margin={:.4} standard_error={:.4} threshold={:.4}",
decision.atom,
decision.firings,
decision.mean_amplitude,
decision.z_alpha,
decision.margin,
decision.standard_error,
decision.threshold,
);
continue;
}
let src = candidate.row(decision.atom);
let mut dst = decoder.row_mut(decision.atom);
dst.assign(&src);
}
(stats, gate)
}
fn revive_dead_atoms(
x: ArrayView2<'_, f32>,
codes: &[SparseCode],
decoder: &mut Array2<f32>,
) -> usize {
let n = x.nrows();
let p = x.ncols();
let k = decoder.nrows();
let mut alive = vec![false; k];
for code in codes.iter() {
for (j, &idx) in code.indices.iter().enumerate() {
if code.codes[j] != 0.0 {
alive[idx as usize] = true;
}
}
}
let dead: Vec<usize> = (0..k).filter(|&a| !alive[a]).collect();
if dead.is_empty() {
return 0;
}
let mut resid = Array2::<f32>::zeros((n, p));
let mut resid_norm2 = vec![0.0f64; n];
for i in 0..n {
let xi = x.row(i);
let mut ri = resid.row_mut(i);
for c in 0..p {
ri[c] = xi[c];
}
let code = &codes[i];
for j in 0..code.indices.len() {
let cj = code.codes[j];
if cj == 0.0 {
continue;
}
let drow = decoder.row(code.indices[j] as usize);
for c in 0..p {
ri[c] -= cj * drow[c];
}
}
let mut acc = 0.0f64;
for c in 0..p {
acc += ri[c] as f64 * ri[c] as f64;
}
resid_norm2[i] = acc;
}
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| {
resid_norm2[b]
.partial_cmp(&resid_norm2[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.cmp(&b))
});
let mut revived = 0usize;
for (t, &atom) in dead.iter().enumerate() {
if t >= n {
break; }
let row = order[t];
if resid_norm2[row] <= (DEAD_DENOM as f64) {
break; }
let src = resid.row(row);
let mut dst = decoder.row_mut(atom);
for c in 0..p {
dst[c] = src[c];
}
revived += 1;
}
revived
}
pub(super) fn solve_decoder(
decoder: &mut Array2<f32>,
eq: &DecoderNormalEq,
ridge: f64,
) -> DecoderSolveStats {
let k = eq.diag.len();
let p = eq.b.ncols();
let mut neigh: Vec<Vec<(u32, f64)>> = vec![Vec::new(); k];
for (&(a, b), &val) in eq.off.iter() {
neigh[a as usize].push((b, val));
neigh[b as usize].push((a, val));
}
for list in neigh.iter_mut() {
list.sort_by_key(|&(nb, _)| nb);
}
let mut stats = DecoderSolveStats {
mean_cofiring_degree: if k == 0 {
0.0
} else {
2.0 * eq.off.len() as f64 / k as f64
},
cg_residual_stop: ridge.max(DEAD_DENOM),
..DecoderSolveStats::default()
};
let mut visited = vec![false; k];
for start in 0..k {
if visited[start] {
continue;
}
if neigh[start].is_empty() {
visited[start] = true;
stats.component_count += 1;
stats.max_component_size = stats.max_component_size.max(1);
let denom = eq.diag[start] + ridge;
if denom <= DEAD_DENOM {
continue;
}
for c in 0..p {
decoder[[start, c]] = (eq.b[[start, c]] / denom) as f32;
}
continue;
}
let mut comp = vec![start];
visited[start] = true;
let mut head = 0usize;
while head < comp.len() {
let node = comp[head];
head += 1;
for &(nb, _) in &neigh[node] {
let nb = nb as usize;
if !visited[nb] {
visited[nb] = true;
comp.push(nb);
}
}
}
comp.sort_unstable();
stats.component_count += 1;
stats.max_component_size = stats.max_component_size.max(comp.len());
solve_component(decoder, eq, ridge, &comp, &neigh, p, &mut stats);
}
if k > 0 {
stats.giant_component_fraction = stats.max_component_size as f64 / k as f64;
}
stats
}
fn solve_component(
decoder: &mut Array2<f32>,
eq: &DecoderNormalEq,
ridge: f64,
comp: &[usize],
neigh: &[Vec<(u32, f64)>],
p: usize,
stats: &mut DecoderSolveStats,
) {
let m = comp.len();
let mut local: HashMap<usize, usize> = HashMap::with_capacity(m);
for (i, &a) in comp.iter().enumerate() {
local.insert(a, i);
}
if m <= MAX_DIRECT_BLOCK {
let mut mat = Array2::<f64>::zeros((m, m));
let mut rhs = Array2::<f64>::zeros((m, p));
for (i, &a) in comp.iter().enumerate() {
mat[[i, i]] = eq.diag[a] + ridge;
for &(nb, val) in &neigh[a] {
if let Some(&j) = local.get(&(nb as usize)) {
mat[[i, j]] = val;
}
}
for c in 0..p {
rhs[[i, c]] = eq.b[[a, c]];
}
}
let sol = cholesky_solve_block(&mat, &rhs);
for (i, &a) in comp.iter().enumerate() {
for c in 0..p {
decoder[[a, c]] = sol[[i, c]] as f32;
}
}
return;
}
let matvec = |xloc: &[f64]| -> Vec<f64> {
let mut y = vec![0.0f64; m];
for (i, &a) in comp.iter().enumerate() {
let mut acc = (eq.diag[a] + ridge) * xloc[i];
for &(nb, val) in &neigh[a] {
if let Some(&j) = local.get(&(nb as usize)) {
acc += val * xloc[j];
}
}
y[i] = acc;
}
y
};
let charge_floor = ridge.max(DEAD_DENOM);
let cap = m.saturating_mul(2).saturating_add(16);
for c in 0..p {
let mut bvec = vec![0.0f64; m];
let mut bnorm2 = 0.0f64;
for (i, &a) in comp.iter().enumerate() {
bvec[i] = eq.b[[a, c]];
bnorm2 += bvec[i] * bvec[i];
}
let bnorm = bnorm2.sqrt();
let mut xvec = vec![0.0f64; m];
if bnorm <= DEAD_DENOM {
for &a in comp {
decoder[[a, c]] = 0.0;
}
continue;
}
let result = cg_solve(&matvec, &bvec, charge_floor, cap);
xvec = result.x.clone();
stats.record_cg(&result);
for (i, &a) in comp.iter().enumerate() {
decoder[[a, c]] = xvec[i] as f32;
}
}
}
struct CgSolveResult {
x: Vec<f64>,
iterations: usize,
relative_residual: f64,
kappa_hat: Option<f64>,
}
fn cg_solve<F>(matvec: &F, b: &[f64], charge_floor: f64, cap: usize) -> CgSolveResult
where
F: Fn(&[f64]) -> Vec<f64>,
{
let n = b.len();
let bnorm = b.iter().map(|v| v * v).sum::<f64>().sqrt();
if bnorm <= DEAD_DENOM {
return CgSolveResult {
x: vec![0.0; n],
iterations: 0,
relative_residual: 0.0,
kappa_hat: None,
};
}
let mut x = vec![0.0f64; n];
let mut r = b.to_vec();
let mut pdir = r.clone();
let mut rs_old: f64 = r.iter().map(|v| v * v).sum();
let mut alphas = Vec::new();
let mut betas = Vec::new();
let mut relative_residual = 1.0;
for iter in 0..cap {
let ap = matvec(&pdir);
let mut pap = 0.0f64;
for i in 0..n {
pap += pdir[i] * ap[i];
}
if pap <= 0.0 || !pap.is_finite() {
return CgSolveResult {
x,
iterations: iter,
relative_residual,
kappa_hat: kappa_from_cg_tridiagonal(&alphas, &betas),
};
}
let alpha = rs_old / pap;
alphas.push(alpha);
for i in 0..n {
x[i] += alpha * pdir[i];
r[i] -= alpha * ap[i];
}
let rs_new: f64 = r.iter().map(|v| v * v).sum();
relative_residual = rs_new.sqrt() / bnorm;
if relative_residual <= charge_floor {
return CgSolveResult {
x,
iterations: iter + 1,
relative_residual,
kappa_hat: kappa_from_cg_tridiagonal(&alphas, &betas),
};
}
let beta = rs_new / rs_old;
if !beta.is_finite() {
return CgSolveResult {
x,
iterations: iter + 1,
relative_residual,
kappa_hat: kappa_from_cg_tridiagonal(&alphas, &betas),
};
}
betas.push(beta);
for i in 0..n {
pdir[i] = r[i] + beta * pdir[i];
}
rs_old = rs_new;
}
CgSolveResult {
x,
iterations: cap,
relative_residual,
kappa_hat: kappa_from_cg_tridiagonal(&alphas, &betas),
}
}
fn kappa_from_cg_tridiagonal(alphas: &[f64], betas: &[f64]) -> Option<f64> {
use faer::Side;
use gam_linalg::faer_ndarray::FaerEigh;
let n = alphas.len();
if n == 0 {
return None;
}
let mut tri = Array2::<f64>::zeros((n, n));
for i in 0..n {
let mut diag = 1.0 / alphas[i];
if i > 0 {
diag += betas[i - 1] / alphas[i - 1];
let off = betas[i - 1].sqrt() / alphas[i - 1];
tri[[i - 1, i]] = off;
tri[[i, i - 1]] = off;
}
tri[[i, i]] = diag;
}
let Ok((evals, _evecs)) = tri.eigh(Side::Lower) else {
return None;
};
let mut min_eval = f64::INFINITY;
let mut max_eval = 0.0f64;
for &eval in evals.iter() {
if eval.is_finite() && eval > 0.0 {
min_eval = min_eval.min(eval);
max_eval = max_eval.max(eval);
}
}
if min_eval.is_finite() && max_eval >= min_eval {
Some(max_eval / min_eval)
} else {
None
}
}
fn cholesky_solve_block(mat: &Array2<f64>, rhs: &Array2<f64>) -> Array2<f64> {
use faer::Side;
use gam_linalg::faer_ndarray::FaerCholesky;
let m = mat.nrows();
let mut a = mat.clone();
let mut bump = 0.0f64;
for _attempt in 0..6 {
if let Ok(factor) = a.cholesky(Side::Lower) {
return factor.solve_mat(rhs);
}
bump = if bump == 0.0 { 1.0e-8 } else { bump * 16.0 };
a = mat.clone();
for i in 0..m {
a[[i, i]] += bump;
}
}
let p = rhs.ncols();
let mut out = Array2::<f64>::zeros((m, p));
for i in 0..m {
let d = mat[[i, i]].max(DEAD_DENOM);
for c in 0..p {
out[[i, c]] = rhs[[i, c]] / d;
}
}
out
}
pub(super) fn unit_norm_rows(decoder: &mut Array2<f32>) {
for mut row in decoder.outer_iter_mut() {
let nrm: f32 = row.iter().map(|v| v * v).sum::<f32>().sqrt();
if nrm > 1.0e-12 {
row.mapv_inplace(|v| v / nrm);
let mut sign = 1.0f32;
for &v in row.iter() {
if v.abs() > 1.0e-9 {
sign = v.signum();
break;
}
}
if sign < 0.0 {
row.mapv_inplace(|v| -v);
}
}
}
}
fn explained_variance(
x: ArrayView2<'_, f32>,
codes: &[SparseCode],
decoder: ArrayView2<'_, f32>,
) -> f64 {
let n = x.nrows();
let p = x.ncols();
let mut means = vec![0.0f64; p];
for i in 0..n {
let xi = x.row(i);
for c in 0..p {
means[c] += xi[c] as f64;
}
}
for c in 0..p {
means[c] /= n as f64;
}
let mut rss = 0.0f64;
let mut tss = 0.0f64;
let mut recon = vec![0.0f64; p];
for i in 0..n {
for c in 0..p {
recon[c] = 0.0;
}
let code = &codes[i];
for j in 0..code.indices.len() {
let cj = code.codes[j] as f64;
if cj == 0.0 {
continue;
}
let drow = decoder.row(code.indices[j] as usize);
for c in 0..p {
recon[c] += cj * drow[c] as f64;
}
}
let xi = x.row(i);
for c in 0..p {
let r = xi[c] as f64 - recon[c];
rss += r * r;
let t = xi[c] as f64 - means[c];
tss += t * t;
}
}
if tss <= 1.0e-24 {
if rss <= 1.0e-24 { 1.0 } else { 0.0 }
} else {
1.0 - rss / tss
}
}
fn residual_scale(
x: ArrayView2<'_, f32>,
codes: &[SparseCode],
decoder: ArrayView2<'_, f32>,
) -> f64 {
let n = x.nrows();
let p = x.ncols();
let mut rss = 0.0f64;
let mut recon = vec![0.0f64; p];
for i in 0..n {
for c in 0..p {
recon[c] = 0.0;
}
let code = &codes[i];
for j in 0..code.indices.len() {
let cj = code.codes[j] as f64;
if cj == 0.0 {
continue;
}
let drow = decoder.row(code.indices[j] as usize);
for c in 0..p {
recon[c] += cj * drow[c] as f64;
}
}
let xi = x.row(i);
for c in 0..p {
let r = xi[c] as f64 - recon[c];
rss += r * r;
}
}
(rss / (n * p) as f64).sqrt()
}
fn pack_codes(codes: &[SparseCode], n: usize, s: usize) -> (Array2<u32>, Array2<f32>) {
let mut indices = Array2::<u32>::zeros((n, s));
let mut code_mat = Array2::<f32>::zeros((n, s));
for (i, code) in codes.iter().enumerate() {
for j in 0..s {
indices[[i, j]] = code.indices[j];
code_mat[[i, j]] = code.codes[j];
}
}
(indices, code_mat)
}
#[cfg(test)]
mod exact_solve_tests {
use super::{
DecoderNormalEq, cg_solve, explained_variance, route_and_code_all, solve_decoder,
solve_decoder_with_routability_gate,
};
use crate::sparse_dict::codes::SparseCode;
use crate::sparse_dict::scoring::TileScorer;
use crate::sparse_dict::{SparseDictConfig, fit_sparse_dictionary};
use ndarray::{Array2, ArrayView2};
use std::collections::HashMap;
fn assemble_normal_eq(
x: ArrayView2<'_, f32>,
codes: &[SparseCode],
k: usize,
p: usize,
) -> DecoderNormalEq {
let mut diag = vec![0.0f64; k];
let mut b = Array2::<f64>::zeros((k, p));
let mut off: HashMap<(u32, u32), f64> = HashMap::new();
let mut firings = vec![0usize; k];
let mut amplitude_sum = vec![0.0f64; k];
for (row_idx, code) in codes.iter().enumerate() {
let xi = x.row(row_idx);
let xi_slice = xi.as_slice();
for a in 0..code.indices.len() {
let ca = code.codes[a] as f64;
if ca == 0.0 {
continue;
}
let ka = code.indices[a];
firings[ka as usize] += 1;
amplitude_sum[ka as usize] += ca.abs();
diag[ka as usize] += ca * ca;
let brow = ka as usize;
let mut brow_view = b.row_mut(brow);
match (brow_view.as_slice_mut(), xi_slice) {
(Some(bs), Some(xs)) => {
for (bref, &xv) in bs.iter_mut().zip(xs.iter()) {
*bref += ca * xv as f64;
}
}
_ => {
for c in 0..p {
brow_view[c] += ca * xi[c] as f64;
}
}
}
for bsel in (a + 1)..code.indices.len() {
let cb = code.codes[bsel] as f64;
if cb == 0.0 {
continue;
}
let kb = code.indices[bsel];
if ka == kb {
diag[ka as usize] += 2.0 * ca * cb;
continue;
}
let key = if ka < kb { (ka, kb) } else { (kb, ka) };
*off.entry(key).or_insert(0.0) += ca * cb;
}
}
}
DecoderNormalEq {
diag,
b,
off,
firings,
amplitude_sum,
}
}
impl DecoderNormalEq {
fn matvec_col(&self, ridge: f64, x: &[f64]) -> Vec<f64> {
let k = self.diag.len();
let mut y = vec![0.0f64; k];
for i in 0..k {
y[i] = (self.diag[i] + ridge) * x[i];
}
for (&(a, b), &val) in self.off.iter() {
y[a as usize] += val * x[b as usize];
y[b as usize] += val * x[a as usize];
}
y
}
}
fn overlapping_problem() -> (Array2<f32>, Vec<SparseCode>, usize, usize) {
let k = 5usize;
let p = 4usize;
let supports: [[u32; 3]; 5] = [[0, 1, 2], [1, 2, 3], [2, 3, 4], [3, 4, 0], [4, 0, 1]];
let codevals: [[f32; 3]; 5] = [
[1.0, 0.5, -0.3],
[0.7, -0.2, 0.4],
[-0.6, 0.9, 0.1],
[0.3, -0.5, 0.8],
[0.2, 0.6, -0.4],
];
let codes: Vec<SparseCode> = supports
.iter()
.zip(codevals.iter())
.map(|(idx, cv)| SparseCode {
indices: idx.to_vec(),
codes: cv.to_vec(),
})
.collect();
let n = codes.len();
let mut x = Array2::<f32>::zeros((n, p));
for i in 0..n {
for c in 0..p {
x[[i, c]] = (((i * 7 + c * 3 + 1) % 13) as f32 - 6.0) / 4.0;
}
}
(x, codes, k, p)
}
fn accumulate_constant_rows(
eq: &mut DecoderNormalEq,
atom: u32,
rows: usize,
code: f32,
row: [f32; 2],
) {
let mut x = Array2::<f32>::zeros((rows, 2));
for i in 0..rows {
x[[i, 0]] = row[0];
x[[i, 1]] = row[1];
}
let codes: Vec<SparseCode> = (0..rows)
.map(|_| SparseCode {
indices: vec![atom],
codes: vec![code],
})
.collect();
eq.accumulate(x.view(), &codes);
}
fn normal_eq_residual(eq: &DecoderNormalEq, decoder: &Array2<f32>, ridge: f64) -> f64 {
let k = eq.diag.len();
let p = eq.b.ncols();
let mut rss = 0.0f64;
let mut bss = 0.0f64;
for c in 0..p {
let dcol: Vec<f64> = (0..k).map(|i| decoder[[i, c]] as f64).collect();
let y = eq.matvec_col(ridge, &dcol);
for i in 0..k {
let r = y[i] - eq.b[[i, c]];
rss += r * r;
bss += eq.b[[i, c]] * eq.b[[i, c]];
}
}
if bss <= 0.0 { 0.0 } else { (rss / bss).sqrt() }
}
#[test]
fn routability_gate_refreshes_well_fired_and_defers_starved_atom() {
let mut eq = DecoderNormalEq::zeros(2, 2);
accumulate_constant_rows(&mut eq, 0, 64, 1.0, [2.0, 0.0]);
accumulate_constant_rows(&mut eq, 1, 1, 1.0, [0.0, 3.0]);
let mut decoder = Array2::<f32>::zeros((2, 2));
decoder[[0, 1]] = 1.0;
decoder[[1, 0]] = 1.0;
let (_stats, gate) = solve_decoder_with_routability_gate(&mut decoder, &eq, 0.0, 1.0);
assert!(gate[0].refresh, "well-fired atom must refresh");
assert!(
gate[0].standard_error <= gate[0].margin,
"well-fired atom should clear the SE-to-margin gate"
);
assert!(!gate[1].refresh, "starved atom must defer");
assert!(
gate[1].standard_error > gate[1].margin,
"starved atom's refresh SE should exceed its charge-floor margin"
);
assert!(
decoder[[0, 0]] > 1.9 && decoder[[0, 1]].abs() < 1.0e-6,
"admitted atom should take its MOD row"
);
assert!(
decoder[[1, 0]] > 0.9 && decoder[[1, 1]].abs() < 1.0e-6,
"deferred atom should keep its previous row"
);
}
#[test]
fn deferred_atom_accumulates_until_routability_threshold_crosses() {
let mut eq = DecoderNormalEq::zeros(1, 2);
let mut decoder = Array2::<f32>::zeros((1, 2));
decoder[[0, 1]] = 1.0;
accumulate_constant_rows(&mut eq, 0, 1, 1.0, [3.0, 0.0]);
let (_stats_first, first_gate) =
solve_decoder_with_routability_gate(&mut decoder, &eq, 0.0, 1.0);
eq.clear_refreshed_atoms(&first_gate);
assert!(!first_gate[0].refresh, "single firing should defer");
assert_eq!(
eq.firings[0], 1,
"deferred atom's firing evidence must remain accumulated"
);
assert!(
decoder[[0, 1]] > 0.9,
"deferred atom must keep its old decoder direction"
);
accumulate_constant_rows(&mut eq, 0, 63, 1.0, [3.0, 0.0]);
let (_stats_second, second_gate) =
solve_decoder_with_routability_gate(&mut decoder, &eq, 0.0, 1.0);
eq.clear_refreshed_atoms(&second_gate);
assert!(
second_gate[0].refresh,
"accumulated firings should cross the routability threshold"
);
assert_eq!(
eq.firings[0], 0,
"refreshed atom's consumed evidence should be cleared"
);
assert!(
decoder[[0, 0]] > 2.9 && decoder[[0, 1]].abs() < 1.0e-6,
"eventually admitted atom should install its MOD row"
);
}
fn connected_tridiagonal_eq(k: usize, p: usize) -> DecoderNormalEq {
let mut diag = vec![0.0f64; k];
for (i, d) in diag.iter_mut().enumerate() {
*d = 1.8 + 0.03 * i as f64;
}
let mut off = std::collections::HashMap::new();
for i in 0..(k - 1) {
off.insert((i as u32, (i + 1) as u32), -0.25);
}
let mut b = Array2::<f64>::zeros((k, p));
for i in 0..k {
for c in 0..p {
b[[i, c]] = ((i * 5 + c * 7 + 3) % 17) as f64 / 11.0 - 0.6;
}
}
DecoderNormalEq {
diag,
b,
off,
firings: vec![4; k],
amplitude_sum: vec![4.0; k],
}
}
#[test]
fn exact_solver_drives_normal_eq_residual_below_tolerance() {
let (x, codes, k, p) = overlapping_problem();
let ridge = 1.0e-6f64;
let eq = assemble_normal_eq(x.view(), &codes, k, p);
assert!(
!eq.off.is_empty(),
"test problem must have off-diagonal coupling (overlapping supports)"
);
let mut decoder = Array2::<f32>::zeros((k, p));
solve_decoder(&mut decoder, &eq, ridge);
let rel = normal_eq_residual(&eq, &decoder, ridge);
assert!(
rel < 1.0e-6,
"coupled decoder solve must drive ‖(A+ρI)D−B‖/‖B‖ to the f32 floor \
(< 1e-6), got {rel}"
);
}
#[test]
fn block_solve_matches_independent_dense_solve() {
use faer::Side;
use gam_linalg::faer_ndarray::FaerCholesky;
let (x, codes, k, p) = overlapping_problem();
let ridge = 1.0e-6f64;
let eq = assemble_normal_eq(x.view(), &codes, k, p);
let mut decoder = Array2::<f32>::zeros((k, p));
solve_decoder(&mut decoder, &eq, ridge);
let mut mat = Array2::<f64>::zeros((k, k));
for i in 0..k {
mat[[i, i]] = eq.diag[i] + ridge;
}
for (&(a, b), &val) in eq.off.iter() {
mat[[a as usize, b as usize]] = val;
mat[[b as usize, a as usize]] = val;
}
let factor = mat.cholesky(Side::Lower).expect("dense SPD system");
let dense = factor.solve_mat(&eq.b);
for i in 0..k {
for c in 0..p {
let got = decoder[[i, c]] as f64;
let want = dense[[i, c]];
assert!(
(got - want).abs() <= 1.0e-5 + 1.0e-5 * want.abs(),
"block solve [{i},{c}] = {got} disagrees with dense solve {want}"
);
}
}
}
#[test]
fn matrix_free_cg_matches_dense_solve_to_charge_floor() {
use faer::Side;
use gam_linalg::faer_ndarray::FaerCholesky;
let k = 12usize;
let p = 3usize;
let ridge = 1.0e-5f64;
let eq = connected_tridiagonal_eq(k, p);
let mut decoder = Array2::<f32>::zeros((k, p));
let stats = solve_decoder(&mut decoder, &eq, ridge);
assert_eq!(stats.component_count, 1);
assert_eq!(stats.max_component_size, k);
assert_eq!(stats.cg_columns, p);
assert!(
stats.cg_relative_residual <= ridge,
"CG residual {} must stop below charge floor {ridge}",
stats.cg_relative_residual
);
let mut mat = Array2::<f64>::zeros((k, k));
for i in 0..k {
mat[[i, i]] = eq.diag[i] + ridge;
}
for (&(a, b), &val) in eq.off.iter() {
mat[[a as usize, b as usize]] = val;
mat[[b as usize, a as usize]] = val;
}
let dense = mat
.cholesky(Side::Lower)
.expect("dense SPD system")
.solve_mat(&eq.b);
let mut diff2 = 0.0f64;
let mut dense2 = 0.0f64;
for i in 0..k {
for c in 0..p {
let diff = decoder[[i, c]] as f64 - dense[[i, c]];
diff2 += diff * diff;
dense2 += dense[[i, c]] * dense[[i, c]];
}
}
let rel = (diff2 / dense2).sqrt();
assert!(
rel <= 5.0 * ridge,
"CG decoder must match dense solve to the charge floor, rel={rel}, floor={ridge}"
);
assert!(
stats.cg_kappa_hat.is_some(),
"CG path must report a Lanczos condition estimate"
);
}
#[test]
fn cg_lanczos_kappa_matches_true_condition_number() {
let eigenvalues = [1.0f64, 1.7, 2.9, 4.6, 8.0, 13.0];
let b = vec![1.0f64; eigenvalues.len()];
let matvec = |x: &[f64]| -> Vec<f64> {
eigenvalues
.iter()
.zip(x.iter())
.map(|(&lambda, &xi)| lambda * xi)
.collect()
};
let result = cg_solve(&matvec, &b, 1.0e-14, eigenvalues.len() + 2);
let got = result.kappa_hat.expect("Lanczos kappa");
let want = eigenvalues[eigenvalues.len() - 1] / eigenvalues[0];
assert!(
(got - want).abs() <= 1.0e-8 * want,
"Lanczos κ̂ {got} must match true condition {want}"
);
}
#[test]
fn returned_ev_is_fresh_code_ev_no_stale_gap() {
let (n, p, k) = (60usize, 6usize, 8usize);
let mut x = Array2::<f32>::zeros((n, p));
for i in 0..n {
for c in 0..p {
x[[i, c]] = (((i * 3 + c * 7 + 1) % 11) as f32 - 5.0) / 5.0;
}
}
let config = SparseDictConfig {
n_atoms: k,
active: 2, minibatch: 16,
max_epochs: 25,
score_tile: 8,
code_ridge: 1.0e-6,
decoder_ridge: 1.0e-6,
tolerance: 1.0e-9,
score_mode: gam_gpu::GpuMode::Off,
};
let fit = fit_sparse_dictionary(x.view(), &config).expect("fit");
let s = fit.active;
assert!(s > 1, "test must run the coupled s>1 lane");
let scorer = TileScorer::new(s, config.score_tile);
let codes = route_and_code_all(
x.view(),
fit.decoder.view(),
&scorer,
s,
config.code_ridge,
config.minibatch,
config.score_mode,
None,
)
.expect("fresh route");
let fresh_ev = explained_variance(x.view(), &codes, fit.decoder.view());
assert!(
(fresh_ev - fit.explained_variance).abs() < 1.0e-6,
"returned EV {} must equal fresh-code EV {fresh_ev} (no stale-code gap)",
fit.explained_variance
);
}
}