use super::codes::{SparseCode, solve_row_codes};
use super::scoring::{ScoreRoutePath, ScoreRouteStats, TileScorer};
use super::{SparseDictConfig, SparseDictConvergence, SparseDictFit};
use ndarray::{Array2, ArrayView2, Axis};
use rayon::prelude::*;
use std::collections::HashMap;
use std::fmt;
use std::time::Instant;
#[derive(Clone, Debug)]
pub enum SparseDictionaryError {
InvalidInput {
reason: String,
},
NumericalFailure {
reason: String,
},
InnerNonConvergence {
epochs: usize,
explained_variance: f64,
ev_residual: f64,
tolerance: f64,
accepted_births: usize,
decoder_fixed_point_residual: f64,
routing_residual: f64,
solve_residual: f64,
solve_tolerance: f64,
decoder_nonconverged_columns: usize,
decoder_factorization_failures: usize,
},
TraceNonConvergence {
rho: f64,
probe: usize,
iterations: usize,
residual: f64,
tolerance: f64,
},
InvalidRemlEvidence {
reason: String,
},
}
impl SparseDictionaryError {
fn invalid_input(reason: impl Into<String>) -> Self {
Self::InvalidInput {
reason: reason.into(),
}
}
}
impl From<String> for SparseDictionaryError {
fn from(reason: String) -> Self {
Self::NumericalFailure { reason }
}
}
impl fmt::Display for SparseDictionaryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidInput { reason } | Self::NumericalFailure { reason } => {
f.write_str(reason)
}
Self::InnerNonConvergence {
epochs,
explained_variance,
ev_residual,
tolerance,
accepted_births,
decoder_fixed_point_residual,
routing_residual,
solve_residual,
solve_tolerance,
decoder_nonconverged_columns,
decoder_factorization_failures,
} => write!(
f,
"fit_sparse_dictionary did not converge after {epochs} epochs: EV \
{explained_variance:.6}, EV residual {ev_residual:.3e} (tolerance \
{tolerance:.3e}), decoder fixed-point residual \
{decoder_fixed_point_residual:.3e}, routing residual {routing_residual:.3e}, \
accepted births {accepted_births}, linear-solve residual \
{solve_residual:.3e} (tolerance {solve_tolerance:.3e}), nonconverged decoder \
columns {decoder_nonconverged_columns}, dense factorization failures \
{decoder_factorization_failures}"
),
Self::TraceNonConvergence {
rho,
probe,
iterations,
residual,
tolerance,
} => write!(
f,
"fit_sparse_dictionary REML trace solve did not converge at rho={rho:.6e}, \
probe {probe}, after {iterations} iterations: relative residual \
{residual:.3e} exceeds {tolerance:.3e}"
),
Self::InvalidRemlEvidence { reason } => {
write!(
f,
"fit_sparse_dictionary REML evidence is invalid: {reason}"
)
}
}
}
}
impl std::error::Error for SparseDictionaryError {}
impl From<SparseDictionaryError> for String {
fn from(error: SparseDictionaryError) -> Self {
error.to_string()
}
}
#[derive(Clone, Debug)]
pub(crate) struct SparseDictIterate {
pub(crate) decoder: Array2<f32>,
pub(crate) indices: Array2<u32>,
pub(crate) codes: Array2<f32>,
pub(crate) explained_variance: f64,
pub(crate) epochs: usize,
pub(crate) active: usize,
pub(crate) score_route_stats: ScoreRouteStats,
pub(crate) decoder_solve_stats: DecoderSolveStats,
inner_ev_residual: f64,
decoder_fixed_point_residual: f64,
routing_residual: f64,
inner_tolerance: f64,
}
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::GpuPolicy,
mut score_route_stats: Option<&mut ScoreRouteStats>,
) -> Result<Vec<SparseCode>, String> {
let n = x.nrows();
let batch = minibatch.max(1);
if n == 0 {
return Ok(Vec::new());
}
let first_end = batch.min(n);
let first_block = x.slice(ndarray::s![0..first_end, ..]);
let first_routed = scorer.route_minibatch_with_mode(first_block, decoder, score_mode)?;
let path = first_routed.path;
if let Some(stats) = score_route_stats.as_deref_mut() {
stats.record_result(&first_routed);
}
let first_active = first_routed.selections;
let mut codes: Vec<SparseCode> = first_block
.axis_iter(Axis(0))
.into_par_iter()
.zip(first_active.into_par_iter())
.map(|(row, active)| solve_row_codes(row, decoder, &active, s, code_ridge))
.collect();
if path == ScoreRoutePath::Cpu {
let plan = gam_gpu::DictionaryScoreRoutePlan::default_for_shape(
batch,
decoder.nrows(),
decoder.ncols(),
);
if first_end < n {
let rest = x.slice(ndarray::s![first_end.., ..]);
let chunk_codes: Vec<Vec<SparseCode>> = rest
.axis_chunks_iter(Axis(0), batch)
.into_par_iter()
.map(|chunk| {
let routed = scorer.route_minibatch(chunk, decoder);
chunk
.axis_iter(Axis(0))
.zip(routed.into_iter())
.map(|(row, active)| solve_row_codes(row, decoder, &active, s, code_ridge))
.collect::<Vec<SparseCode>>()
})
.collect();
for chunk in chunk_codes {
if let Some(stats) = score_route_stats.as_deref_mut() {
stats.record(plan, ScoreRoutePath::Cpu);
}
codes.extend(chunk);
}
}
} else {
let mut start = first_end;
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)
}
fn decoder_fixed_point_residual(previous: &Array2<f32>, next: &Array2<f32>) -> f64 {
previous
.axis_iter(Axis(0))
.zip(next.axis_iter(Axis(0)))
.map(|(left, right)| {
let left_norm2 = left.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>();
let right_norm2 = right.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>();
if left_norm2 <= DEAD_DENOM && right_norm2 <= DEAD_DENOM {
return 0.0;
}
if left_norm2 <= DEAD_DENOM || right_norm2 <= DEAD_DENOM {
return 1.0;
}
let dot = left
.iter()
.zip(right.iter())
.map(|(&a, &b)| (a as f64) * (b as f64))
.sum::<f64>();
(1.0 - dot * dot / (left_norm2 * right_norm2)).clamp(0.0, 1.0)
})
.fold(0.0, f64::max)
}
fn routing_fixed_point_residual(
x: ArrayView2<'_, f32>,
previous_decoder: ArrayView2<'_, f32>,
previous: &[SparseCode],
next_decoder: ArrayView2<'_, f32>,
next: &[SparseCode],
) -> f64 {
let mut code_delta2 = 0.0f64;
let mut code_scale2 = 0.0f64;
let mut reconstruction_delta2 = 0.0f64;
let mut data_scale2 = 0.0f64;
for row in 0..x.nrows() {
let old = &previous[row];
let new = &next[row];
for (slot, &atom) in old.indices.iter().enumerate() {
let old_value = old.codes[slot] as f64;
if old_value == 0.0 {
continue;
}
let new_value = new
.indices
.iter()
.zip(new.codes.iter())
.filter(|(candidate, _)| **candidate == atom)
.map(|(_, &value)| value as f64)
.sum::<f64>();
let delta = new_value - old_value;
code_delta2 += delta * delta;
code_scale2 += old_value * old_value + new_value * new_value;
}
for (slot, &atom) in new.indices.iter().enumerate() {
let new_value = new.codes[slot] as f64;
if new_value == 0.0
|| old
.indices
.iter()
.zip(old.codes.iter())
.any(|(&candidate, &value)| candidate == atom && value != 0.0)
{
continue;
}
code_delta2 += new_value * new_value;
code_scale2 += new_value * new_value;
}
for column in 0..x.ncols() {
let old_value = old
.indices
.iter()
.zip(old.codes.iter())
.map(|(&atom, &code)| {
(code as f64) * previous_decoder[[atom as usize, column]] as f64
})
.sum::<f64>();
let new_value = new
.indices
.iter()
.zip(new.codes.iter())
.map(|(&atom, &code)| (code as f64) * next_decoder[[atom as usize, column]] as f64)
.sum::<f64>();
let delta = new_value - old_value;
reconstruction_delta2 += delta * delta;
let observed = x[[row, column]] as f64;
data_scale2 += observed * observed;
}
}
let code_residual = if code_scale2 > 0.0 {
code_delta2 / code_scale2
} else {
0.0
};
let reconstruction_residual = if data_scale2 > 0.0 {
reconstruction_delta2 / data_scale2
} else if reconstruction_delta2 == 0.0 {
0.0
} else {
f64::INFINITY
};
code_residual.max(reconstruction_residual)
}
pub(super) fn run(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
) -> Result<SparseDictIterate, SparseDictionaryError> {
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 fit_start = Instant::now();
let mut decoder = seed_decoder(x, k);
unit_norm_rows(&mut decoder);
log::warn!(
"[SAE sparse_dict] seeded decoder N={n} P={p} K={k} s={s} \
seed_s={:.1} (route + refresh follow)",
fit_start.elapsed().as_secs_f64(),
);
let scorer = TileScorer::new(s, config.score_tile);
let mut score_route_stats = ScoreRouteStats::default();
let mut epochs_run = 0usize;
let mut decoder_solve_stats = DecoderSolveStats::default();
let mut ev_residual = f64::INFINITY;
let mut decoder_residual = f64::INFINITY;
let mut routing_residual = f64::INFINITY;
let mut accepted_births = 0usize;
let initial_route_start = Instant::now();
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),
)?;
log::warn!(
"[SAE sparse_dict] initial route done: minibatches={} device={} cpu={} \
route_s={:.1} elapsed_s={:.1}",
score_route_stats.minibatches,
score_route_stats.device_minibatches,
score_route_stats.cpu_minibatches,
initial_route_start.elapsed().as_secs_f64(),
fit_start.elapsed().as_secs_f64(),
);
let mut current_ev = explained_variance(x, &codes, decoder.view());
for epoch in 0..config.max_epochs {
epochs_run = epoch + 1;
let epoch_start = Instant::now();
let certified_decoder = decoder.clone();
let certified_codes = codes.clone();
let certified_ev = current_ev;
let mut normal_eq = DecoderNormalEq::zeros(k, p);
normal_eq.accumulate(x, &certified_codes);
let sigma = residual_scale(x, &codes, decoder.view());
let (stats, _gate) = solve_decoder_with_routability_gate(
&mut decoder,
&normal_eq,
config.decoder_ridge as f64,
sigma,
);
decoder_solve_stats = stats;
let refresh_secs = epoch_start.elapsed().as_secs_f64();
unit_norm_rows(&mut decoder);
let revived_atoms = revive_dead_atoms(x, &codes, &mut decoder);
if !revived_atoms.is_empty() {
unit_norm_rows(&mut decoder);
}
let mut next_codes = route_and_code_all(
x,
decoder.view(),
&scorer,
s,
config.code_ridge,
config.minibatch,
config.score_mode,
Some(&mut score_route_stats),
)?;
let route_secs = epoch_start.elapsed().as_secs_f64() - refresh_secs;
let next_ev = explained_variance(x, &next_codes, decoder.view());
let improve = next_ev - certified_ev;
let mut revived_mask = vec![false; k];
for &atom in &revived_atoms {
revived_mask[atom] = true;
}
let mut accepted_mask = vec![false; k];
for code in &next_codes {
for (slot, &atom) in code.indices.iter().enumerate() {
let atom = atom as usize;
if code.codes[slot] != 0.0 && revived_mask[atom] {
accepted_mask[atom] = true;
}
}
}
accepted_births = accepted_mask
.into_iter()
.filter(|accepted| *accepted)
.count();
for &atom in &revived_atoms {
let accepted = next_codes.iter().any(|code| {
code.indices
.iter()
.zip(code.codes.iter())
.any(|(&candidate, &value)| candidate as usize == atom && value != 0.0)
});
if !accepted {
decoder.row_mut(atom).fill(0.0);
}
}
ev_residual = improve.abs();
decoder_residual = decoder_fixed_point_residual(&certified_decoder, &decoder);
routing_residual = routing_fixed_point_residual(
x,
certified_decoder.view(),
&certified_codes,
decoder.view(),
&next_codes,
);
log::warn!(
"[SAE epoch {}/{}] ev={:.6} improve={:.3e} revived={} refresh_s={:.2} \
route_s={:.2} elapsed_s={:.1} max_component={} cg_columns={} cg_nonconverged={} \
cg_kappa_bound={:?} cg_relative_residual={:.3e}",
epochs_run,
config.max_epochs,
next_ev,
improve,
revived_atoms.len(),
refresh_secs,
route_secs,
fit_start.elapsed().as_secs_f64(),
decoder_solve_stats.max_component_size,
decoder_solve_stats.cg_columns,
decoder_solve_stats.cg_nonconverged_columns,
decoder_solve_stats.cg_kappa_bound,
decoder_solve_stats.cg_relative_residual,
);
if accepted_births == 0
&& decoder_solve_stats.cg_nonconverged_columns == 0
&& decoder_solve_stats.dense_factorization_failures == 0
&& decoder_solve_stats.cg_relative_residual
<= decoder_solve_stats.cg_residual_stop.max(f64::MIN_POSITIVE)
&& ev_residual <= config.tolerance
&& decoder_residual <= config.tolerance
&& routing_residual <= config.tolerance
{
let (indices, code_mat) = pack_codes(&certified_codes, n, s);
return Ok(SparseDictIterate {
decoder: certified_decoder,
indices,
codes: code_mat,
explained_variance: certified_ev,
epochs: epochs_run,
active: s,
score_route_stats,
decoder_solve_stats,
inner_ev_residual: ev_residual,
decoder_fixed_point_residual: decoder_residual,
routing_residual,
inner_tolerance: config.tolerance,
});
}
codes = std::mem::take(&mut next_codes);
current_ev = next_ev;
}
Err(SparseDictionaryError::InnerNonConvergence {
epochs: epochs_run,
explained_variance: current_ev,
ev_residual,
tolerance: config.tolerance,
accepted_births,
decoder_fixed_point_residual: decoder_residual,
routing_residual,
solve_residual: decoder_solve_stats.cg_relative_residual,
solve_tolerance: decoder_solve_stats.cg_residual_stop,
decoder_nonconverged_columns: decoder_solve_stats.cg_nonconverged_columns,
decoder_factorization_failures: decoder_solve_stats.dense_factorization_failures,
})
}
pub(crate) fn run_linear_fast_kernel(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
shared_rho: f64,
) -> Result<SparseDictIterate, SparseDictionaryError> {
let mut unified = *config;
unified.code_ridge = shared_rho as f32;
unified.decoder_ridge = shared_rho as f32;
run(x, &unified)
}
#[derive(Clone, Copy, Debug)]
pub struct LinearBlockRemlStats {
pub gram_edof: f64,
pub p_cols: usize,
pub penalty_energy: f64,
pub rss: f64,
pub n_obs: usize,
}
pub fn linear_shared_rho_fs_step(
stats: &LinearBlockRemlStats,
rho: f64,
) -> Result<f64, SparseDictionaryError> {
if !(rho.is_finite() && rho > 0.0) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!("rho must be finite and positive; got {rho}"),
});
}
let gamma_tot = (stats.p_cols as f64) * stats.gram_edof;
let total_obs = (stats.n_obs.saturating_mul(stats.p_cols)) as f64;
if !(gamma_tot.is_finite() && gamma_tot > 0.0 && gamma_tot < total_obs) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!(
"pooled effective dof must lie strictly inside (0, {total_obs}); got {gamma_tot}"
),
});
}
if !(stats.rss.is_finite() && stats.rss >= 0.0) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!("RSS must be finite and non-negative; got {}", stats.rss),
});
}
if !(stats.penalty_energy.is_finite() && stats.penalty_energy > 0.0) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!(
"decoder penalty energy must be finite and positive; got {}",
stats.penalty_energy
),
});
}
let resid_dof = total_obs - gamma_tot;
let sigma2 = stats.rss / resid_dof;
let rho_new = gamma_tot * sigma2 / stats.penalty_energy;
if !(rho_new.is_finite() && rho_new > 0.0) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!("Fellner-Schall update produced invalid rho {rho_new}"),
});
}
Ok(rho_new)
}
const EDOF_TRACE_VARIANCE_PER_UNIT_TRACE: f64 = 0.05;
fn hutchinson_gram_edof(
diag: &[f64],
off: &HashMap<(u32, u32), f64>,
rho: f64,
k: usize,
) -> Result<f64, SparseDictionaryError> {
if !(rho.is_finite() && rho > 0.0) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!("trace ridge must be finite and positive; got {rho}"),
});
}
if k == 0 {
return Ok(0.0);
}
let mut neigh: Vec<Vec<(u32, f64)>> = vec![Vec::new(); k];
for (&(a, b), &val) in 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 matvec = |v: &[f64]| -> Vec<f64> {
let mut y = vec![0.0f64; k];
for a in 0..k {
let mut acc = (diag[a] + rho) * v[a];
for &(nb, val) in &neigh[a] {
acc += val * v[nb as usize];
}
y[a] = acc;
}
y
};
let mut lambda_max_bound = 0.0f64;
for a in 0..k {
let mut off_abs = 0.0f64;
for &(_, val) in &neigh[a] {
off_abs += val.abs();
}
lambda_max_bound = lambda_max_bound.max(diag[a] + rho + off_abs);
}
let lambda_min = rho.max(DEAD_DENOM);
let kappa_bound = (lambda_max_bound / lambda_min).max(1.0);
let root = kappa_bound.sqrt();
let residual_tolerance = decoder_solve_relative_tolerance();
let chebyshev = 0.5 * root * (2.0 * root / residual_tolerance).ln();
let cap = (chebyshev.max(0.0).ceil() as usize).min(k).max(1);
let m_probes = (2.0 / EDOF_TRACE_VARIANCE_PER_UNIT_TRACE).ceil() as usize;
let m_probes = m_probes.max(1);
let mut base_seed = gam_linalg::utils::splitmix64_hash(k as u64);
base_seed = gam_linalg::utils::splitmix64_hash(base_seed ^ (off.len() as u64).wrapping_add(1));
for &d in diag.iter() {
base_seed = gam_linalg::utils::splitmix64_hash(base_seed ^ d.to_bits());
}
let mut complementary_trace_acc = 0.0f64;
for probe in 0..m_probes {
let probe_salt =
gam_linalg::utils::splitmix64_hash(base_seed ^ (probe as u64).wrapping_add(1));
let mut z = vec![0.0f64; k];
for (a, zi) in z.iter_mut().enumerate() {
let h = gam_linalg::utils::splitmix64_hash(probe_salt ^ (a as u64).wrapping_add(1));
*zi = if h >> 63 == 0 { 1.0 } else { -1.0 };
}
let result = cg_solve(&matvec, &z, residual_tolerance, cap);
if result.stop != CgStop::Converged {
return Err(SparseDictionaryError::TraceNonConvergence {
rho,
probe,
iterations: result.iterations,
residual: result.relative_residual,
tolerance: residual_tolerance,
});
}
let zt_minv_z: f64 = z.iter().zip(result.x.iter()).map(|(zi, wi)| zi * wi).sum();
complementary_trace_acc += rho * zt_minv_z;
}
let complementary_trace = complementary_trace_acc / m_probes as f64;
let edof = k as f64 - complementary_trace;
if !(edof.is_finite() && (0.0..=k as f64).contains(&edof)) {
return Err(SparseDictionaryError::InvalidRemlEvidence {
reason: format!(
"Hutchinson effective dof must lie in [0, {k}]; got {edof} without clamping"
),
});
}
Ok(edof)
}
fn code_gram_from_routing(
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
k: usize,
) -> (Vec<f64>, HashMap<(u32, u32), f64>) {
let mut diag = vec![0.0f64; k];
let mut off: HashMap<(u32, u32), f64> = HashMap::new();
let s = indices.ncols();
for i in 0..indices.nrows() {
for a in 0..s {
let ca = codes[[i, a]] as f64;
if ca == 0.0 {
continue;
}
let ka = indices[[i, a]];
diag[ka as usize] += ca * ca;
for b in (a + 1)..s {
let cb = codes[[i, b]] as f64;
if cb == 0.0 {
continue;
}
let kb = indices[[i, b]];
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;
}
}
}
(diag, off)
}
fn reconstruction_rss_from_parts(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
) -> f64 {
let p = x.ncols();
let s = indices.ncols();
let mut rss = 0.0f64;
let mut recon = vec![0.0f64; p];
for i in 0..x.nrows() {
for r in recon.iter_mut() {
*r = 0.0;
}
for a in 0..s {
let cj = codes[[i, a]] as f64;
if cj == 0.0 {
continue;
}
let drow = decoder.row(indices[[i, a]] as usize);
for (c, r) in recon.iter_mut().enumerate() {
*r += cj * drow[c] as f64;
}
}
let xi = x.row(i);
for c in 0..p {
let d = xi[c] as f64 - recon[c];
rss += d * d;
}
}
rss
}
pub fn linear_block_reml_stats(
x: ArrayView2<'_, f32>,
fit: &SparseDictFit,
rho: f64,
) -> Result<LinearBlockRemlStats, SparseDictionaryError> {
linear_block_reml_stats_from_parts(
x,
fit.decoder.view(),
fit.indices.view(),
fit.codes.view(),
rho,
)
}
fn linear_block_reml_stats_from_parts(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
rho: f64,
) -> Result<LinearBlockRemlStats, SparseDictionaryError> {
let k = decoder.nrows();
let (diag, off) = code_gram_from_routing(indices, codes, k);
let gram_edof = hutchinson_gram_edof(&diag, &off, rho, k)?;
let penalty_energy: f64 = decoder.iter().map(|&d| (d as f64) * (d as f64)).sum();
let rss = reconstruction_rss_from_parts(x, decoder, indices, codes);
Ok(LinearBlockRemlStats {
gram_edof,
p_cols: x.ncols(),
penalty_energy,
rss,
n_obs: x.nrows(),
})
}
fn reml_schedule_rho_log_tol(inner_tolerance: f64) -> f64 {
inner_tolerance.sqrt().max(f64::EPSILON.sqrt())
}
pub fn run_linear_reml_schedule(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
) -> Result<SparseDictFit, SparseDictionaryError> {
validate(x, config)?;
if config.code_ridge != config.decoder_ridge {
return Err(SparseDictionaryError::invalid_input(format!(
"fit_sparse_dictionary has one shared REML ridge, so code_ridge ({}) and \
decoder_ridge ({}) must be equal",
config.code_ridge, config.decoder_ridge
)));
}
let data_energy = x
.iter()
.map(|&value| (value as f64) * (value as f64))
.sum::<f64>();
if data_energy == 0.0 {
let active = config.active.min(config.n_atoms).max(1);
let tolerance = reml_schedule_rho_log_tol(config.tolerance);
return Ok(SparseDictFit {
decoder: Array2::<f32>::zeros((config.n_atoms, x.ncols())),
indices: Array2::<u32>::zeros((x.nrows(), active)),
codes: Array2::<f32>::zeros((x.nrows(), active)),
explained_variance: 1.0,
epochs: 0,
convergence: SparseDictConvergence {
inner_ev_residual: 0.0,
inner_tolerance: config.tolerance,
decoder_residual: 0.0,
decoder_tolerance: config.tolerance,
routing_residual: 0.0,
routing_tolerance: config.tolerance,
outer_rho_residual: 0.0,
outer_tolerance: tolerance,
selected_rho: f64::INFINITY,
outer_iterations: 0,
},
active,
score_route_stats: ScoreRouteStats::default(),
decoder_solve_stats: DecoderSolveStats::default(),
});
}
let mut rho = config.decoder_ridge as f64;
let mut fit = run_linear_fast_kernel(x, config, rho)?;
let tol = reml_schedule_rho_log_tol(config.tolerance);
let mut outer_iterations = 0usize;
loop {
outer_iterations += 1;
let stats = linear_block_reml_stats_from_parts(
x,
fit.decoder.view(),
fit.indices.view(),
fit.codes.view(),
rho,
)?;
let rho_new = linear_shared_rho_fs_step(&stats, rho)?;
let log_change = (rho_new.ln() - rho.ln()).abs();
log::warn!(
"[SAE reml-schedule iter {}] rho={:.6e} rho_new={:.6e} log_change={:.3e} \
edof={:.2} rss={:.6e} penalty_energy={:.6e} tol={:.3e}",
outer_iterations,
rho,
rho_new,
log_change,
stats.gram_edof,
stats.rss,
stats.penalty_energy,
tol,
);
if log_change <= tol {
let convergence = SparseDictConvergence {
inner_ev_residual: fit.inner_ev_residual,
inner_tolerance: fit.inner_tolerance,
decoder_residual: fit.decoder_fixed_point_residual,
decoder_tolerance: fit.inner_tolerance,
routing_residual: fit.routing_residual,
routing_tolerance: fit.inner_tolerance,
outer_rho_residual: log_change,
outer_tolerance: tol,
selected_rho: rho,
outer_iterations,
};
return Ok(SparseDictFit {
decoder: fit.decoder,
indices: fit.indices,
codes: fit.codes,
explained_variance: fit.explained_variance,
epochs: fit.epochs,
convergence,
active: fit.active,
score_route_stats: fit.score_route_stats,
decoder_solve_stats: fit.decoder_solve_stats,
});
}
rho = rho_new;
fit = run_linear_fast_kernel(x, config, rho)?;
}
}
fn validate(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
) -> Result<(), SparseDictionaryError> {
if x.nrows() == 0 || x.ncols() == 0 {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary requires a non-empty N×P matrix",
));
}
if !x.iter().all(|v| v.is_finite()) {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary input must be finite",
));
}
if config.n_atoms == 0 {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary requires K >= 1",
));
}
if config.active == 0 {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary requires active (top_s) >= 1",
));
}
if config.max_epochs == 0 {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary requires max_epochs >= 1",
));
}
if !(config.code_ridge.is_finite() && config.code_ridge > 0.0) {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary code_ridge must be finite and positive",
));
}
if !(config.decoder_ridge.is_finite() && config.decoder_ridge > 0.0) {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary decoder_ridge must be finite and positive",
));
}
if !(config.tolerance.is_finite() && config.tolerance >= 0.0) {
return Err(SparseDictionaryError::invalid_input(
"fit_sparse_dictionary tolerance must be finite and non-negative",
));
}
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);
let chosen = if atom < n {
let (bi, _bv) = min_dist2
.par_iter_mut()
.enumerate()
.map(|(i, md)| {
let xi = x.row(i);
let mut d2 = 0.0f32;
for c in 0..p {
let d = xi[c] - prev[c];
d2 += d * d;
}
if d2 < *md {
*md = d2;
}
(i, *md)
})
.reduce(
|| (usize::MAX, f32::NEG_INFINITY),
|a, b| {
if b.1 > a.1 || (b.1 == a.1 && b.0 < a.0) {
b
} else {
a
}
},
);
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;
fn decoder_solve_relative_tolerance() -> f64 {
f64::EPSILON.sqrt()
}
pub(super) fn direct_solve_size_threshold(k: usize) -> usize {
if k == 0 {
return 0;
}
(k as f64).powf(2.0 / 3.0).ceil() as usize
}
#[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,
pub cg_nonconverged_columns: usize,
pub dense_factorization_failures: usize,
pub cg_kappa_bound: Option<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,
cg_nonconverged_columns: 0,
dense_factorization_failures: 0,
cg_kappa_bound: None,
}
}
}
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 result.stop != CgStop::Converged {
self.cg_nonconverged_columns += 1;
}
if let Some(kappa) = result.kappa_hat {
self.cg_kappa_hat = Some(self.cg_kappa_hat.map_or(kappa, |old| old.max(kappa)));
}
}
fn record_kappa_bound(&mut self, bound: f64) {
self.cg_kappa_bound = Some(self.cg_kappa_bound.map_or(bound, |old| old.max(bound)));
}
}
#[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).max(0.0)
} else {
0.0
};
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>,
) -> Vec<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 Vec::new();
}
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 = Vec::new();
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.push(atom);
}
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: decoder_solve_relative_tolerance(),
..DecoderSolveStats::default()
};
let direct_threshold = direct_solve_size_threshold(k);
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,
direct_threshold,
&mut stats,
);
}
if k > 0 {
stats.giant_component_fraction = stats.max_component_size as f64 / k as f64;
}
log::debug!(
"[SAE percolation] K={k} mean_degree={:.4} giant_fraction={:.4} \
components={} max_component={} direct_threshold={direct_threshold} \
cg_columns={} cg_iterations={} cg_kappa_hat={:?} cg_kappa_bound={:?} \
cg_nonconverged_columns={} cg_relative_residual={:.3e} cg_residual_stop={:.3e}",
stats.mean_cofiring_degree,
stats.giant_component_fraction,
stats.component_count,
stats.max_component_size,
stats.cg_columns,
stats.cg_iterations,
stats.cg_kappa_hat,
stats.cg_kappa_bound,
stats.cg_nonconverged_columns,
stats.cg_relative_residual,
stats.cg_residual_stop,
);
stats
}
fn solve_component(
decoder: &mut Array2<f32>,
eq: &DecoderNormalEq,
ridge: f64,
comp: &[usize],
neigh: &[Vec<(u32, f64)>],
p: usize,
direct_threshold: 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 <= direct_threshold {
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 Some(sol) = cholesky_solve_block(&mat, &rhs) else {
stats.dense_factorization_failures += 1;
return;
};
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 residual_tolerance = decoder_solve_relative_tolerance();
let mut lambda_max_bound = 0.0f64;
let mut lambda_min_bound = f64::INFINITY;
for &a in comp {
let mut off_abs = 0.0f64;
for &(nb, val) in &neigh[a] {
if local.contains_key(&(nb as usize)) {
off_abs += val.abs();
}
}
let center = eq.diag[a] + ridge;
lambda_max_bound = lambda_max_bound.max(center + off_abs);
lambda_min_bound = lambda_min_bound.min(center - off_abs);
}
let lambda_min = lambda_min_bound.max(ridge).max(DEAD_DENOM);
let kappa_bound = (lambda_max_bound / lambda_min).max(1.0);
stats.record_kappa_bound(kappa_bound);
let root = kappa_bound.sqrt();
let chebyshev = 0.5 * root * (2.0 * root / residual_tolerance).ln();
let cap = (chebyshev.max(0.0).ceil() as usize).min(m).max(1);
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];
}
if bnorm2.sqrt() <= DEAD_DENOM {
for &a in comp {
decoder[[a, c]] = 0.0;
}
continue;
}
let result = cg_solve(&matvec, &bvec, residual_tolerance, cap);
stats.record_cg(&result);
if result.stop == CgStop::Converged {
for (i, &a) in comp.iter().enumerate() {
decoder[[a, c]] = result.x[i] as f32;
}
} else {
log::warn!(
"[SAE CG] component size={m} did not converge: stop={:?} iters={} \
rel_residual={:.3e} residual_tolerance={:.3e} \
kappa_bound={:.3e} cap={cap}",
result.stop,
result.iterations,
result.relative_residual,
residual_tolerance,
kappa_bound,
);
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum CgStop {
Converged,
Breakdown,
CapReached,
}
struct CgSolveResult {
x: Vec<f64>,
iterations: usize,
relative_residual: f64,
kappa_hat: Option<f64>,
stop: CgStop,
}
fn cg_solve<F>(matvec: &F, b: &[f64], residual_tolerance: f64, cap: usize) -> CgSolveResult
where
F: Fn(&[f64]) -> Vec<f64>,
{
use gam_linalg::pcg::{DotReduction, PcgStop, pcg_core};
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,
stop: CgStop::Converged,
};
}
let rhs = ndarray::Array1::from_vec(b.to_vec());
let precond = ndarray::Array1::<f64>::from_elem(n, 1.0);
let mut solution = ndarray::Array1::<f64>::zeros(n);
let apply = |v: &ndarray::Array1<f64>, out: &mut ndarray::Array1<f64>| {
let av = matvec(v.as_slice().expect("pcg direction vector is contiguous"));
out.assign(&ndarray::Array1::from_vec(av));
};
let result = pcg_core(
apply,
&rhs.view(),
&precond.view(),
residual_tolerance,
cap,
0,
true,
DotReduction::Serial,
&mut solution.view_mut(),
);
let relative_residual = if result.rhs_norm > 0.0 {
result.final_residual_norm / result.rhs_norm
} else {
0.0
};
let kappa_hat = result
.diagnostics
.as_ref()
.and_then(|d| kappa_from_cg_tridiagonal(&d.alpha, &d.beta));
let stop = match result.stop {
PcgStop::Converged => CgStop::Converged,
PcgStop::MaxIters => CgStop::CapReached,
PcgStop::Breakdown | PcgStop::BadPreconditioner => CgStop::Breakdown,
};
CgSolveResult {
x: solution.to_vec(),
iterations: result.iterations,
relative_residual,
kappa_hat,
stop,
}
}
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>) -> Option<Array2<f64>> {
use faer::Side;
use gam_linalg::faer_ndarray::FaerCholesky;
let factor = mat.cholesky(Side::Lower).ok()?;
Some(factor.solve_mat(rhs))
}
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::{
CgStop, 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 direct_solve_threshold_tracks_percolation_scale_not_a_constant() {
use super::direct_solve_size_threshold;
assert_eq!(direct_solve_size_threshold(0), 0);
assert_eq!(direct_solve_size_threshold(1), 1);
for &k in &[8usize, 12, 64, 1024, 100_000] {
let tau = direct_solve_size_threshold(k);
let want = (k as f64).powf(2.0 / 3.0).ceil() as usize;
assert_eq!(tau, want, "threshold must equal ⌈K^{{2/3}}⌉ for K={k}");
assert!(
tau < k,
"a giant (size-K) component must exceed the dense threshold at K={k} (got {tau})"
);
}
assert!(direct_solve_size_threshold(100_000) > direct_solve_size_threshold(12));
}
#[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 cg_reports_cap_reached_when_iterations_exhausted() {
let eigenvalues = [1.0f64, 5.0, 25.0, 125.0, 625.0];
let b = vec![1.0f64; eigenvalues.len()];
let matvec = |x: &[f64]| -> Vec<f64> {
eigenvalues
.iter()
.zip(x.iter())
.map(|(&l, &xi)| l * xi)
.collect()
};
let result = cg_solve(&matvec, &b, 1.0e-12, 1);
assert_eq!(result.stop, CgStop::CapReached);
assert_eq!(result.iterations, 1);
assert!(result.x.iter().all(|v| v.is_finite()));
}
#[test]
fn cg_reports_breakdown_on_indefinite_operator() {
let eigenvalues = [1.0f64, -3.0, 2.0];
let b = vec![1.0f64, 1.0, 1.0];
let matvec = |x: &[f64]| -> Vec<f64> {
eigenvalues
.iter()
.zip(x.iter())
.map(|(&l, &xi)| l * xi)
.collect()
};
let result = cg_solve(&matvec, &b, 1.0e-12, 64);
assert_eq!(result.stop, CgStop::Breakdown);
assert!(result.iterations <= 64);
assert!(result.x.iter().all(|v| v.is_finite()));
}
#[test]
fn near_singular_giant_component_terminates_with_typed_failure() {
let k = 200usize;
let p = 2usize;
let diag = vec![1.0f64; k];
let mut off = HashMap::new();
for a in 0..(k - 1) {
off.insert((a as u32, (a + 1) as u32), 0.5);
}
let mut b = Array2::<f64>::zeros((k, p));
for i in 0..k {
b[[i, 0]] = ((i * 7 + 3) % 11) as f64 - 5.0;
b[[i, 1]] = ((i * 5 + 1) % 13) as f64 - 6.0;
}
let eq = DecoderNormalEq {
diag,
b,
off,
firings: vec![4; k],
amplitude_sum: vec![4.0; k],
};
let mut decoder = Array2::<f32>::zeros((k, p));
let ridge = 1.0e-9f64;
let stats = solve_decoder(&mut decoder, &eq, ridge);
assert_eq!(
stats.max_component_size, k,
"path graph is one giant component"
);
let kappa_bound = stats.cg_kappa_bound.expect("a-priori kappa bound recorded");
assert!(
kappa_bound > 1.0e6,
"near-singular block must report a large a-priori kappa bound, got {kappa_bound}"
);
assert!(
stats.cg_nonconverged_columns >= 1,
"an under-resolved column must be a TYPED non-convergence, not a silent spin"
);
assert!(
stats.cg_iterations <= k * p,
"iterations must be bounded by the derived cap, got {}",
stats.cg_iterations
);
assert!(
decoder.iter().all(|v| v.is_finite()),
"failed columns must leave the prior finite decoder untouched"
);
}
#[test]
fn shared_rho_fs_step_matches_closed_form_evidence_fixed_point() {
use super::{LinearBlockRemlStats, linear_shared_rho_fs_step};
let stats = LinearBlockRemlStats {
gram_edof: 2.5,
p_cols: 3,
penalty_energy: 4.0,
rss: 10.0,
n_obs: 8,
};
let rho_new = linear_shared_rho_fs_step(&stats, 1.0e-3).expect("valid FS evidence");
assert!(
(rho_new - 1.136_363_636_363_636_5).abs() < 1.0e-12,
"FS step must match the closed-form evidence fixed point, got {rho_new}"
);
let zero_energy = LinearBlockRemlStats {
penalty_energy: 0.0,
..stats
};
assert!(linear_shared_rho_fs_step(&zero_energy, 7.0e-4).is_err());
let zero_edof = LinearBlockRemlStats {
gram_edof: 0.0,
..stats
};
assert!(linear_shared_rho_fs_step(&zero_edof, 7.0e-4).is_err());
let saturated = LinearBlockRemlStats {
gram_edof: 100.0,
p_cols: 3,
penalty_energy: 4.0,
rss: 10.0,
n_obs: 8,
};
assert!(linear_shared_rho_fs_step(&saturated, 1.0e-3).is_err());
}
fn next_unit(state: &mut u64) -> f64 {
let h = gam_linalg::utils::splitmix64(state);
(h >> 11) as f64 / (1u64 << 53) as f64
}
fn densify_gram(diag: &[f64], off: &HashMap<(u32, u32), f64>, k: usize) -> Array2<f64> {
let mut a = Array2::<f64>::zeros((k, k));
for i in 0..k {
a[[i, i]] = diag[i];
}
for (&(r, c), &v) in off.iter() {
a[[r as usize, c as usize]] = v;
a[[c as usize, r as usize]] = v;
}
a
}
fn exact_gram_edof(a: &Array2<f64>, rho: f64) -> f64 {
use faer::Side;
use gam_linalg::faer_ndarray::FaerCholesky;
let k = a.nrows();
let mut m = a.clone();
for i in 0..k {
m[[i, i]] += rho;
}
let y = m.cholesky(Side::Lower).expect("A+ρI is SPD").solve_mat(a);
(0..k).map(|i| y[[i, i]]).sum()
}
#[test]
fn hutchinson_gram_edof_matches_exact_dense_trace() {
use super::{
EDOF_TRACE_VARIANCE_PER_UNIT_TRACE, code_gram_from_routing, hutchinson_gram_edof,
};
let (k, s, n) = (32usize, 3usize, 400usize);
let mut indices = Array2::<u32>::zeros((n, s));
let mut codes = Array2::<f32>::zeros((n, s));
let mut rng = 0x51E2_D3C4_A5B6_9788u64;
for i in 0..n {
for j in 0..s {
let atom = ((i * (j + 1) * 7 + j * 5 + 1) % k) as u32;
indices[[i, j]] = atom;
codes[[i, j]] = (next_unit(&mut rng) as f32 - 0.5) * 2.0;
}
}
let (diag, off) = code_gram_from_routing(indices.view(), codes.view(), k);
let a_dense = densify_gram(&diag, &off, k);
for &rho in &[1.0e-3_f64, 1.0e-1, 1.0] {
let exact = exact_gram_edof(&a_dense, rho);
let approx = hutchinson_gram_edof(&diag, &off, rho, k)
.expect("every trace probe must reach its residual certificate");
let probes = (2.0 / EDOF_TRACE_VARIANCE_PER_UNIT_TRACE).ceil();
let c = (k as f64 - exact).max(0.0);
let sd_bound = (2.0 * c / probes).sqrt();
let tol = 6.0 * sd_bound + 1.0e-6;
assert!(
(approx - exact).abs() <= tol,
"Hutchinson edof {approx} vs exact {exact} at rho={rho} exceeds derived \
6σ tolerance {tol} (c={c}, probes={probes})"
);
assert!(
approx >= 0.0 && approx <= k as f64 + 1.0e-9,
"edof {approx} must lie in [0, K]"
);
}
}
#[test]
fn shared_rho_fixed_point_converges_and_tracks_planted_noise() {
use super::{
linear_block_reml_stats_from_parts, linear_shared_rho_fs_step,
reml_schedule_rho_log_tol, run_linear_fast_kernel,
};
fn planted_noisy(n: usize, p: usize, k: usize, noise: f32, seed: u64) -> Array2<f32> {
let mut atoms = Array2::<f32>::zeros((k, p));
for atom in 0..k {
let mut norm = 0.0f64;
for c in 0..p {
let v = (((atom * 13 + c * 7 + 3) % 17) as f32 - 8.0) / 8.0;
atoms[[atom, c]] = v;
norm += (v as f64) * (v as f64);
}
let inv = 1.0 / norm.sqrt().max(1.0e-12) as f32;
for c in 0..p {
atoms[[atom, c]] *= inv;
}
}
let mut rng = seed;
let mut x = Array2::<f32>::zeros((n, p));
for i in 0..n {
let a0 = (i % k) as usize;
let a1 = ((i / k + 1) % k) as usize;
let c0 = 0.6 + 0.4 * next_unit(&mut rng) as f32;
let c1 = 0.2 + 0.3 * next_unit(&mut rng) as f32;
for c in 0..p {
let clean = c0 * atoms[[a0, c]] + c1 * atoms[[a1, c]];
let eps = noise * (next_unit(&mut rng) as f32 - 0.5) * 2.0;
x[[i, c]] = clean + eps;
}
}
x
}
let (n, p, k) = (300usize, 12usize, 24usize);
let config = SparseDictConfig {
n_atoms: k,
active: 2,
minibatch: 64,
max_epochs: 40,
score_tile: 12,
code_ridge: 1.0e-6,
decoder_ridge: 1.0e-6,
tolerance: 1.0e-9,
score_mode: gam_gpu::GpuPolicy::Off,
};
let fixed_point = |x: ArrayView2<'_, f32>| -> (f64, f64) {
let mut rho = config.decoder_ridge as f64;
let mut last_rel = f64::INFINITY;
for _ in 0..16 {
let fit = run_linear_fast_kernel(x, &config, rho).expect("kernel fit");
let stats = linear_block_reml_stats_from_parts(
x,
fit.decoder.view(),
fit.indices.view(),
fit.codes.view(),
rho,
)
.expect("trace evidence");
let rho_new = linear_shared_rho_fs_step(&stats, rho).expect("valid FS evidence");
last_rel = (rho_new.ln() - rho.ln()).abs();
rho = rho_new;
}
(rho, last_rel)
};
let x_low = planted_noisy(n, p, k, 0.03, 0x1111_2222_3333_4444);
let x_high = planted_noisy(n, p, k, 0.40, 0x1111_2222_3333_4444);
let (rho_low, rel_low) = fixed_point(x_low.view());
let (rho_high, rel_high) = fixed_point(x_high.view());
assert!(
rho_low.is_finite() && rho_low > 0.0 && rho_high.is_finite() && rho_high > 0.0,
"shared ρ* must be finite and positive (low={rho_low}, high={rho_high})"
);
let band = reml_schedule_rho_log_tol(config.tolerance);
assert!(
rel_low <= band && rel_high <= band,
"FS fixed point must settle within the derived stopping band {band}: \
last relative moves low={rel_low} high={rel_high}"
);
assert!(
rho_high > rho_low,
"shared ρ* must grow with planted noise: high-noise ρ*={rho_high} \
must exceed low-noise ρ*={rho_low}"
);
}
#[test]
fn reml_schedule_held_out_ev_matches_or_beats_magic_ridge() {
use super::{run_linear_fast_kernel, run_linear_reml_schedule};
use crate::sparse_dict::codes::solve_row_codes;
fn held_out_ev(
decoder: ArrayView2<'_, f32>,
x_test: ArrayView2<'_, f32>,
s: usize,
tile: usize,
code_ridge: f32,
) -> f64 {
let n = x_test.nrows();
let p = x_test.ncols();
let scorer = TileScorer::new(s, tile);
let mut means = vec![0.0f64; p];
for i in 0..n {
for c in 0..p {
means[c] += x_test[[i, c]] as f64;
}
}
for m in means.iter_mut() {
*m /= n as f64;
}
let mut rss = 0.0f64;
let mut tss = 0.0f64;
for i in 0..n {
let row = x_test.row(i);
let active = scorer.route_row(row, decoder);
let code = solve_row_codes(row, decoder, &active, s, code_ridge);
let mut recon = vec![0.0f64; p];
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;
}
}
for c in 0..p {
let r = x_test[[i, c]] as f64 - recon[c];
rss += r * r;
let t = x_test[[i, 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
}
}
let (k, p, n) = (24usize, 12usize, 500usize);
let mut atoms = Array2::<f32>::zeros((k, p));
for atom in 0..k {
let mut norm = 0.0f64;
for c in 0..p {
let v = (((atom * 11 + c * 5 + 2) % 13) as f32 - 6.0) / 6.0;
atoms[[atom, c]] = v;
norm += (v as f64) * (v as f64);
}
let inv = 1.0 / norm.sqrt().max(1.0e-12) as f32;
for c in 0..p {
atoms[[atom, c]] *= inv;
}
}
let mut rng = 0x0BAD_C0FF_EE12_3456u64;
let mut x = Array2::<f32>::zeros((n, p));
for i in 0..n {
let a0 = i % k;
let a1 = (i / k + 1) % k;
let c0 = 0.6 + 0.4 * next_unit(&mut rng) as f32;
let c1 = 0.2 + 0.3 * next_unit(&mut rng) as f32;
for c in 0..p {
let clean = c0 * atoms[[a0, c]] + c1 * atoms[[a1, c]];
let eps = 0.15 * (next_unit(&mut rng) as f32 - 0.5) * 2.0;
x[[i, c]] = clean + eps;
}
}
let mut train_rows = Vec::new();
let mut test_rows = Vec::new();
for i in 0..n {
if i % 5 == 0 {
test_rows.push(i);
} else {
train_rows.push(i);
}
}
let mut x_train = Array2::<f32>::zeros((train_rows.len(), p));
for (r, &i) in train_rows.iter().enumerate() {
x_train.row_mut(r).assign(&x.row(i));
}
let mut x_test = Array2::<f32>::zeros((test_rows.len(), p));
for (r, &i) in test_rows.iter().enumerate() {
x_test.row_mut(r).assign(&x.row(i));
}
let s = 2usize;
let tile = 12usize;
let config = SparseDictConfig {
n_atoms: k,
active: s,
minibatch: 128,
max_epochs: 60,
score_tile: tile,
code_ridge: 1.0e-6,
decoder_ridge: 1.0e-6,
tolerance: 1.0e-9,
score_mode: gam_gpu::GpuPolicy::Off,
};
let magic = run_linear_fast_kernel(x_train.view(), &config, config.decoder_ridge as f64)
.expect("magic-ridge fit");
let reml = run_linear_reml_schedule(x_train.view(), &config).expect("reml schedule fit");
let magic_ev = held_out_ev(
magic.decoder.view(),
x_test.view(),
s,
tile,
config.code_ridge,
);
let reml_ev = held_out_ev(
reml.decoder.view(),
x_test.view(),
s,
tile,
config.code_ridge,
);
assert!(
reml_ev + 1.0e-3 >= magic_ev,
"REML-selected shared ρ held-out EV {reml_ev} must match-or-beat the \
magic-ridge baseline {magic_ev}"
);
}
#[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::GpuPolicy::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
);
}
}