use super::*;
use super::outer_objective::ProbeRefusalKind;
use crate::chart_coordinate_solve::PeriodicCurveExtrema;
use opt::{BacktrackConfig, RidgeSchedule, backtracking_line_search, escalate_ridge};
const SAE_MANIFOLD_ROW_RIDGE_MAX_ATTEMPTS: usize = 12;
const SAE_MANIFOLD_LM_RATIO_LOW: f64 = 0.25;
const SAE_MANIFOLD_LM_RATIO_HIGH: f64 = 0.75;
const SAE_MANIFOLD_LM_RIDGE_FACTOR: f64 = 4.0;
impl InnerGlobalizationHint {
pub(crate) fn cold(step_size: f64, ridge_ext_coord: f64, ridge_beta: f64) -> Self {
Self {
warm_step: step_size,
lm_ridge_t: ridge_ext_coord,
lm_ridge_b: ridge_beta,
}
}
pub(crate) fn resume(
carried: Option<Self>,
step_size: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) -> Self {
carried.map_or_else(
|| Self::cold(step_size, ridge_ext_coord, ridge_beta),
|hint| Self {
warm_step: hint.warm_step.max(step_size),
lm_ridge_t: hint.lm_ridge_t.max(ridge_ext_coord),
lm_ridge_b: hint.lm_ridge_b.max(ridge_beta),
},
)
}
pub(crate) fn record_accepted_step(
&mut self,
accepted_step: f64,
warm_growth: f64,
unit_step_ceiling: f64,
gain_ratio: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) {
let clean_acceptance = accepted_step >= self.warm_step;
self.warm_step = if clean_acceptance {
(self.warm_step * warm_growth).min(unit_step_ceiling)
} else {
(accepted_step * warm_growth).min(unit_step_ceiling)
};
if gain_ratio < SAE_MANIFOLD_LM_RATIO_LOW {
self.lm_ridge_t *= SAE_MANIFOLD_LM_RIDGE_FACTOR;
self.lm_ridge_b *= SAE_MANIFOLD_LM_RIDGE_FACTOR;
} else if gain_ratio > SAE_MANIFOLD_LM_RATIO_HIGH {
self.lm_ridge_t =
(self.lm_ridge_t / SAE_MANIFOLD_LM_RIDGE_FACTOR).max(ridge_ext_coord);
self.lm_ridge_b =
(self.lm_ridge_b / SAE_MANIFOLD_LM_RIDGE_FACTOR).max(ridge_beta);
}
}
pub(crate) fn reset(
&mut self,
step_size: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) {
*self = Self::cold(step_size, ridge_ext_coord, ridge_beta);
}
}
const ARD_SPREAD_FLOOR: f64 = 1.0e-12;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum JointFitTermination {
Frozen,
Heuristic,
NoStrictDecrease,
IterationGrantExhausted,
}
pub(crate) struct JointFitOutcome {
pub(crate) loss: SaeManifoldLoss,
pub(crate) termination: JointFitTermination,
pub(crate) state_moved: bool,
pub(crate) moved_at: Option<StateMoveSite>,
}
pub(crate) struct EvidenceJointFitOutcome {
pub(crate) loss: SaeManifoldLoss,
pub(crate) fixed_point: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum StateMoveSite {
EntryBlockSweep,
TemperatureSchedule,
AcceptedNewtonStep,
ProximalCorrectionStep,
FrameRefresh,
InnerIncumbentRestore,
ExitBlockSweep,
ExitWarrantyRestore,
GaugeOrbitDescent,
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub(crate) struct GaugeOrbitDescent {
pub(crate) rounds: usize,
pub(crate) objective_decrease: f64,
pub(crate) dimension: usize,
pub(crate) max_directional_derivative: f64,
pub(crate) evaluations: usize,
pub(crate) entry_objective: Option<f64>,
pub(crate) exit_objective: Option<f64>,
}
impl GaugeOrbitDescent {
pub(crate) fn moved(&self) -> bool {
self.rounds > 0 && self.objective_decrease > 0.0
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum EvidenceFixedPointGap {
None,
NotASettledRoot,
StateMoved(StateMoveSite),
StateNotRecurred,
TemperatureStillAnnealing,
}
#[inline]
fn canonicalize_softmax_logit_row(logits: &mut [f64]) {
match logits.len() {
0 => {}
1 => logits[0] = 0.0,
k => {
let reference = logits[k - 1];
for logit in &mut logits[..k - 1] {
*logit -= reference;
}
logits[k - 1] = 0.0;
}
}
}
fn axis_coordinate_spread(coords: ArrayView2<'_, f64>, axis: usize, period: Option<f64>) -> f64 {
let n = coords.nrows();
if n == 0 || axis >= coords.ncols() {
return f64::NAN;
}
match period {
Some(p) if p.is_finite() && p > 0.0 => {
let w = std::f64::consts::TAU / p;
let mut cs = 0.0;
let mut sn = 0.0;
for row in 0..n {
let ang = coords[[row, axis]] * w;
cs += ang.cos();
sn += ang.sin();
}
let r = (cs * cs + sn * sn).sqrt() / n as f64;
(1.0 - r).max(0.0)
}
_ => {
let mut mean = 0.0;
for row in 0..n {
mean += coords[[row, axis]];
}
mean /= n as f64;
let mut var = 0.0;
for row in 0..n {
let d = coords[[row, axis]] - mean;
var += d * d;
}
var / n as f64
}
}
}
pub(crate) struct TargetCenteredColStats {
col_means: Vec<f64>,
ss_tot: f64,
}
impl TargetCenteredColStats {
pub(crate) fn compute(target: ArrayView2<'_, f64>) -> Self {
let n = target.nrows();
let p = target.ncols();
let mut col_means = vec![0.0_f64; p];
let mut ss_tot = 0.0_f64;
for col in 0..p {
let mut mean = 0.0_f64;
for (count, row) in (0..n).enumerate() {
let x = target[[row, col]];
mean += (x - mean) / (count as f64 + 1.0);
}
for row in 0..n {
let dev = target[[row, col]] - mean;
ss_tot += dev * dev;
}
col_means[col] = mean;
}
Self { col_means, ss_tot }
}
pub(crate) fn ss_tot(&self) -> f64 {
self.ss_tot
}
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct StructuralCoCollapseEvidence {
pub(crate) decoder_span_rank: usize,
pub(crate) target_reach: f64,
pub(crate) random_subspace_null: f64,
}
#[derive(Clone, Debug)]
pub(crate) struct DictionaryCollapseVerdict {
pub(crate) explained_variance: f64,
pub(crate) decoder_vanishing: super::construction::VanishedAtomsProof,
pub(crate) structural_collapse: Option<StructuralCoCollapseEvidence>,
}
impl DictionaryCollapseVerdict {
pub(crate) fn all_decoders_vanished(&self, atom_count: usize) -> bool {
self.decoder_vanishing.all_atoms_vanished(atom_count)
}
pub(crate) fn structurally_collapsed(&self) -> bool {
self.structural_collapse.is_some()
}
pub(crate) fn degenerate(&self, atom_count: usize) -> bool {
self.all_decoders_vanished(atom_count) || self.structurally_collapsed()
}
pub(crate) fn proof_unavailable_reason(&self) -> Option<&str> {
self.decoder_vanishing.unavailable_reason()
}
pub(crate) fn collapse_event_floor(&self, atom_count: usize) -> Option<f64> {
if self.all_decoders_vanished(atom_count) {
self.decoder_vanishing.signal_vanish_boundary()
} else {
self.structural_collapse
.map(|evidence| evidence.random_subspace_null)
}
}
}
pub(crate) fn ambient_sphere_killing_directions(u: [f64; 3]) -> [[f64; 3]; 3] {
[
[0.0, -u[2], u[1]],
[u[2], 0.0, -u[0]],
[-u[1], u[0], 0.0],
]
}
impl SaeManifoldTerm {
pub(crate) fn dictionary_collapse_verdict(
&self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
target_col_stats: Option<&TargetCenteredColStats>,
) -> Result<DictionaryCollapseVerdict, String> {
let fitted = self.try_fitted_target_aware(target, Some(rho))?;
let (n, p) = target.dim();
let residual = &fitted - ⌖
let residual_energy = self.residual_energy_for_vanishing(residual.view())?;
let explained_variance = self.dictionary_reconstruction_ev_from_residual_sum_squares(
residual_energy.sum_squares(),
target,
target_col_stats,
)?;
let mut grams = self.empty_decoder_gram_accumulator();
self.accumulate_decoder_gram(&mut grams)?;
let n_eff = self.per_atom_effective_sample_size();
let decoder_vanishing = self.vanished_atoms_from_signal_upper_bound(
&grams,
&n_eff,
residual_energy.mean_square(),
)?;
let k = self.k_atoms();
let frames_sc = (0..k)
.map(|atom| crate::manifold::certificate::certificate_output_frame(self, atom))
.collect::<Result<Vec<_>, String>>()?;
let r_dec_sc = union_output_frame_rank(&frames_sc, p);
let structural_collapse = if r_dec_sc == 0 || r_dec_sc >= k || p == 0 || r_dec_sc >= p {
None
} else {
let owned_stats;
let stats = match target_col_stats {
Some(stats) => stats,
None => {
owned_stats = TargetCenteredColStats::compute(target);
&owned_stats
}
};
if !(stats.ss_tot > 0.0) {
None
} else {
let q = union_output_frame_basis(&frames_sc, p);
let qc = q.ncols();
let mut in_span = 0.0_f64;
for row in 0..n {
for c in 0..qc {
let mut projection = 0.0_f64;
for output in 0..p {
projection += q[[output, c]]
* (target[[row, output]] - stats.col_means[output]);
}
in_span += projection * projection;
}
}
let target_reach = in_span / stats.ss_tot;
let random_subspace_null = r_dec_sc as f64 / p as f64;
(target_reach <= random_subspace_null).then_some(
StructuralCoCollapseEvidence {
decoder_span_rank: r_dec_sc,
target_reach,
random_subspace_null,
},
)
}
};
Ok(DictionaryCollapseVerdict {
explained_variance,
decoder_vanishing,
structural_collapse,
})
}
pub fn apply_newton_step(
&mut self,
delta_ext_coord: ArrayView1<'_, f64>,
delta_beta: ArrayView1<'_, f64>,
step_size: f64,
) -> Result<(), String> {
self.apply_newton_step_impl(delta_ext_coord, delta_beta, step_size, true)
}
pub(crate) fn snapshot_mutable_state(&self) -> SaeManifoldMutableState {
let atoms = self
.atoms
.iter()
.map(|atom| SaeManifoldAtomSnapshot {
decoder_coefficients: atom.decoder_coefficients().clone(),
decoder_frame: atom.decoder_frame.clone(),
smooth_penalty: atom.smooth_penalty().clone(),
basis_evaluator: atom.basis_evaluator.clone(),
basis_second_jet: atom.basis_second_jet.clone(),
homotopy_eta: atom.homotopy_eta,
chart_canonicalized: atom.chart_canonicalized,
reduced_column_map: atom.reduced_column_map.clone(),
caller_managed_basis: atom.basis_evaluator.is_none().then(|| {
(
atom.basis_values.clone(),
atom.basis_jacobian.clone(),
)
}),
})
.collect();
SaeManifoldMutableState {
atoms,
logits: self.assignment.logits.clone(),
coords: self.assignment.coords.clone(),
last_row_layout: self.last_row_layout.clone(),
}
}
pub(crate) fn matches_mutable_state(&self, snapshot: &SaeManifoldMutableState) -> bool {
let atoms_match = self.atoms.len() == snapshot.atoms.len()
&& self
.atoms
.iter()
.zip(snapshot.atoms.iter())
.all(|(atom, saved)| {
let evaluator_matches = match (&atom.basis_evaluator, &saved.basis_evaluator) {
(Some(current), Some(expected)) => Arc::ptr_eq(current, expected),
(None, None) => true,
_ => false,
};
let second_jet_matches = match (&atom.basis_second_jet, &saved.basis_second_jet)
{
(Some(current), Some(expected)) => Arc::ptr_eq(current, expected),
(None, None) => true,
_ => false,
};
let decoder_frame_matches = match (&atom.decoder_frame, &saved.decoder_frame) {
(Some(current), Some(expected)) => {
current.frame() == expected.frame()
&& current.gauge_singular_values()
== expected.gauge_singular_values()
}
(None, None) => true,
_ => false,
};
atom.decoder_coefficients() == &saved.decoder_coefficients
&& atom.chart_canonicalized == saved.chart_canonicalized
&& atom.reduced_column_map == saved.reduced_column_map
&& match &saved.caller_managed_basis {
Some((basis, jet)) => {
atom.basis_evaluator.is_none()
&& atom.basis_values == *basis
&& atom.basis_jacobian == *jet
}
None => atom.basis_evaluator.is_some(),
}
&& decoder_frame_matches
&& atom.smooth_penalty() == &saved.smooth_penalty
&& atom.homotopy_eta.to_bits() == saved.homotopy_eta.to_bits()
&& evaluator_matches
&& second_jet_matches
});
let coords_match = self.assignment.coords.len() == snapshot.coords.len()
&& self
.assignment
.coords
.iter()
.zip(snapshot.coords.iter())
.all(|(current, expected)| {
current.latent_id() == expected.latent_id()
&& current.as_flat() == expected.as_flat()
});
let row_layout_matches = match (&self.last_row_layout, &snapshot.last_row_layout) {
(Some(current), Some(expected)) => {
current.active_atoms == expected.active_atoms
&& current.coord_starts == expected.coord_starts
&& current.coord_offsets_full == expected.coord_offsets_full
&& current.coord_dims == expected.coord_dims
}
(None, None) => true,
_ => false,
};
atoms_match
&& self.assignment.logits == snapshot.logits
&& coords_match
&& row_layout_matches
}
pub(crate) fn restore_mutable_state(
&mut self,
snapshot: &SaeManifoldMutableState,
) -> Result<(), SaeMutableStateRestoreError> {
let incompatible = |
component: &'static str,
expected: Vec<usize>,
observed: Vec<usize>,
| SaeMutableStateRestoreError::IncompatibleCardinality {
component,
expected,
observed,
};
let k = self.atoms.len();
if snapshot.atoms.len() != k {
return Err(incompatible(
"atom",
vec![k],
vec![snapshot.atoms.len()],
));
}
if self.assignment.coords.len() != k {
return Err(incompatible(
"live coordinate-block",
vec![k],
vec![self.assignment.coords.len()],
));
}
if snapshot.coords.len() != k {
return Err(incompatible(
"snapshot coordinate-block",
vec![k],
vec![snapshot.coords.len()],
));
}
if snapshot.logits.dim() != self.assignment.logits.dim() {
return Err(incompatible(
"logit shape",
self.assignment.logits.shape().to_vec(),
snapshot.logits.shape().to_vec(),
));
}
let n = snapshot.logits.nrows();
for atom_idx in 0..k {
let atom = &self.atoms[atom_idx];
let coord = &snapshot.coords[atom_idx];
if atom.n_obs() != n {
return Err(incompatible(
"atom row",
vec![n],
vec![atom.n_obs()],
));
}
if coord.n_obs() != n || coord.latent_dim() != atom.latent_dim() {
return Err(incompatible(
"coordinate shape",
vec![n, atom.latent_dim()],
vec![coord.n_obs(), coord.latent_dim()],
));
}
}
if let Some(layout) = snapshot.last_row_layout.as_ref() {
if layout.active_atoms.len() != n || layout.coord_starts.len() != n {
return Err(incompatible(
"row-layout row",
vec![n, n],
vec![layout.active_atoms.len(), layout.coord_starts.len()],
));
}
if layout.coord_dims.len() != k || layout.coord_offsets_full.len() != k {
return Err(incompatible(
"row-layout atom",
vec![k, k],
vec![layout.coord_dims.len(), layout.coord_offsets_full.len()],
));
}
let mut full_offset = 0usize;
for atom_idx in 0..k {
let expected_dim = snapshot.coords[atom_idx].latent_dim();
if layout.coord_dims[atom_idx] != expected_dim
|| layout.coord_offsets_full[atom_idx] != full_offset
{
return Err(SaeMutableStateRestoreError::InvalidTopology {
component: "row layout".to_string(),
detail: format!(
"atom {atom_idx} has dim/offset ({}, {}), expected ({expected_dim}, {full_offset})",
layout.coord_dims[atom_idx],
layout.coord_offsets_full[atom_idx]
),
});
}
full_offset += expected_dim;
}
for row in 0..n {
let active = &layout.active_atoms[row];
let starts = &layout.coord_starts[row];
if starts.len() != active.len() {
return Err(incompatible(
"row-layout active block",
vec![active.len()],
vec![starts.len()],
));
}
let mut compact_offset = 0usize;
let mut previous = None;
for (&atom_idx, &start) in active.iter().zip(starts.iter()) {
if atom_idx >= k || previous.is_some_and(|prior| prior >= atom_idx) {
return Err(SaeMutableStateRestoreError::InvalidTopology {
component: "row layout".to_string(),
detail: format!(
"row {row} active atoms are not strict sorted indices below {k}: {active:?}"
),
});
}
if start != compact_offset {
return Err(SaeMutableStateRestoreError::InvalidTopology {
component: "row layout".to_string(),
detail: format!(
"row {row}, atom {atom_idx} starts at {start}, expected {compact_offset}"
),
});
}
compact_offset += layout.coord_dims[atom_idx];
previous = Some(atom_idx);
}
}
}
let restored_logits = snapshot.logits.clone();
let restored_coords = snapshot.coords.clone();
let restored_row_layout = snapshot.last_row_layout.clone();
let mut prepared = Vec::with_capacity(k);
for atom_idx in 0..k {
let coords = snapshot.coords[atom_idx].as_matrix();
prepared.push(
self.atoms[atom_idx]
.prepare_mutable_state_restore(&snapshot.atoms[atom_idx], coords.view())
.map_err(|detail| SaeMutableStateRestoreError::InvalidTopology {
component: format!("atom {atom_idx}"),
detail,
})?,
);
}
for (atom_idx, restored) in prepared.into_iter().enumerate() {
self.atoms[atom_idx].commit_prepared_mutable_state(restored);
}
self.assignment.logits = restored_logits;
self.assignment.coords = restored_coords;
self.last_row_layout = restored_row_layout;
Ok(())
}
pub(crate) fn refresh_basis_from_current_coords(&mut self) -> Result<(), String> {
let parallel = self.n_obs() >= SAE_LOSS_PARALLEL_ROW_MIN
&& self.k_atoms() > 1
&& rayon::current_thread_index().is_none();
self.refresh_basis_from_current_coords_with_parallelism(parallel)
}
pub(crate) fn refresh_basis_from_current_coords_with_parallelism(
&mut self,
parallel: bool,
) -> Result<(), String> {
if self.atoms.len() != self.assignment.coords.len() {
return Err(format!(
"SaeManifoldTerm::refresh_basis_from_current_coords: {} atoms but {} coordinate blocks",
self.atoms.len(),
self.assignment.coords.len()
));
}
if parallel {
use rayon::prelude::*;
let outcomes: Vec<Result<(), String>> = self
.atoms
.par_iter_mut()
.zip(self.assignment.coords.par_iter())
.map(|(atom, coord)| {
with_nested_parallel(|| {
let coords = coord.as_matrix();
atom.refresh_basis(coords.view())
})
})
.collect();
for outcome in outcomes {
outcome?;
}
} else {
for (atom, coord) in self.atoms.iter_mut().zip(&self.assignment.coords) {
let coords = coord.as_matrix();
atom.refresh_basis(coords.view())?;
}
}
Ok(())
}
fn apply_coordinate_step_from_rows<F>(
&mut self,
n: usize,
q: usize,
coord_offsets: &[usize],
step_size: f64,
delta_at: F,
refresh_basis: bool,
parallel: bool,
) -> Result<(), String>
where
F: Fn(usize, usize) -> f64 + Sync,
{
if self.atoms.len() != self.assignment.coords.len() {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: {} atoms but {} coordinate blocks",
self.atoms.len(),
self.assignment.coords.len()
));
}
if coord_offsets.len() != self.atoms.len() {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: {} coordinate offsets for {} atoms",
coord_offsets.len(),
self.atoms.len()
));
}
let update_atom = |atom_idx: usize,
atom: &mut SaeManifoldAtom,
coord: &mut LatentCoordValues,
nested_parallel: bool|
-> Result<(), String> {
let d = coord.latent_dim();
let mut delta_coord = Array1::<f64>::zeros(n * d);
for row in 0..n {
let row_base = row * q + coord_offsets[atom_idx];
for axis in 0..d {
delta_coord[row * d + axis] = step_size * delta_at(row, row_base + axis);
}
}
coord.retract_flat_delta(delta_coord.view());
if refresh_basis {
let coords = coord.as_matrix();
if nested_parallel {
with_nested_parallel(|| atom.refresh_basis(coords.view()))?;
} else {
atom.refresh_basis(coords.view())?;
}
}
Ok(())
};
if parallel {
use rayon::prelude::*;
let outcomes: Vec<Result<(), String>> = self
.atoms
.par_iter_mut()
.zip(self.assignment.coords.par_iter_mut())
.enumerate()
.map(|(atom_idx, (atom, coord))| update_atom(atom_idx, atom, coord, true))
.collect();
for outcome in outcomes {
outcome?;
}
} else {
for (atom_idx, (atom, coord)) in self
.atoms
.iter_mut()
.zip(self.assignment.coords.iter_mut())
.enumerate()
{
update_atom(atom_idx, atom, coord, false)?;
}
}
Ok(())
}
fn apply_decoder_step_from_flat(
&mut self,
delta_beta: ArrayView1<'_, f64>,
step_size: f64,
parallel: bool,
) -> Result<(), String> {
let expected = self.beta_dim();
if delta_beta.len() != expected {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: full decoder step length {} != expected {expected}",
delta_beta.len()
));
}
let p = self.output_dim();
let offsets = self.beta_offsets();
let update_atom = |atom_idx: usize, atom: &mut SaeManifoldAtom| {
let m = atom.basis_size();
let offset = offsets[atom_idx];
for basis_col in 0..m {
for out_col in 0..p {
let flat_idx = offset + basis_col * p + out_col;
atom.decoder_coefficients_mut()[[basis_col, out_col]] +=
step_size * delta_beta[flat_idx];
}
}
};
if parallel {
use rayon::prelude::*;
self.atoms
.par_iter_mut()
.enumerate()
.for_each(|(atom_idx, atom)| update_atom(atom_idx, atom));
} else {
for (atom_idx, atom) in self.atoms.iter_mut().enumerate() {
update_atom(atom_idx, atom);
}
}
Ok(())
}
pub(crate) fn canonicalize_affine_gauge_after_accept(
&mut self,
rho: Option<&SaeManifoldRho>,
) -> Result<(), String> {
for atom_idx in 0..self.k_atoms() {
if !matches!(
self.atoms[atom_idx].basis_kind(),
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::EuclideanPatch
| SaeAtomBasisKind::Duchon
| SaeAtomBasisKind::Poincare
) {
continue;
}
self.canonicalize_atom_affine_gauge(atom_idx, rho)?;
}
Ok(())
}
pub(crate) fn canonicalize_atom_affine_gauge(
&mut self,
atom_idx: usize,
rho: Option<&SaeManifoldRho>,
) -> Result<(), String> {
let n = self.n_obs();
let d = self.assignment.coords[atom_idx].latent_dim();
if n == 0 || d == 0 {
return Ok(());
}
let Some(evaluator) = self.atoms[atom_idx].basis_evaluator.as_ref() else {
return Ok(());
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let weights = self.atom_affine_gauge_weights(atom_idx, rho)?;
let weight_sum: f64 = weights.iter().sum();
if !(weight_sum.is_finite() && weight_sum > 0.0) {
return Ok(());
}
let mut shift = vec![0.0_f64; d];
for row in 0..n {
let w = weights[row];
for axis in 0..d {
shift[axis] += w * coords[[row, axis]];
}
}
for value in &mut shift {
*value /= weight_sum;
}
let mut scale = vec![1.0_f64; d];
let mut changed = false;
for axis in 0..d {
let mut var = 0.0_f64;
for row in 0..n {
let centered = coords[[row, axis]] - shift[axis];
var += weights[row] * centered * centered;
}
let rms = (var / weight_sum).sqrt();
if rms.is_finite() && rms > 1.0e-12 {
scale[axis] = rms;
}
if shift[axis].abs() > 1.0e-12 || (scale[axis] - 1.0).abs() > 1.0e-12 {
changed = true;
}
}
if !changed {
return Ok(());
}
let Some(new_evaluator) = evaluator.affine_transformed_evaluator(
&shift,
&scale,
self.atoms[atom_idx].basis_size(),
)?
else {
return Ok(());
};
let mut new_coords = coords.clone();
for row in 0..n {
for axis in 0..d {
new_coords[[row, axis]] = (coords[[row, axis]] - shift[axis]) / scale[axis];
}
}
let (new_phi, new_jet) = if self.atoms[atom_idx].homotopy_eta == 1.0 {
new_evaluator.evaluate(new_coords.view())?
} else {
let evaluated = new_evaluator
.evaluate_phi_eta(new_coords.view(), self.atoms[atom_idx].homotopy_eta)?;
(evaluated.phi, evaluated.jet)
};
let old_phi = self.atoms[atom_idx].basis_values.clone();
if new_phi.dim() != old_phi.dim() {
return Err(format!(
"SaeManifoldTerm::canonicalize_atom_affine_gauge: transformed basis shape {:?} != {:?}",
new_phi.dim(),
old_phi.dim()
));
}
let transport = solve_basis_transport(new_phi.view(), old_phi.view())?;
let old_decoder = self.atoms[atom_idx].decoder_coefficients().clone();
let old_smooth_penalty = self.atoms[atom_idx].smooth_penalty().clone();
let new_decoder = fast_ab(&transport, &old_decoder);
let old_fit = fast_ab(&old_phi, &old_decoder);
let new_fit = fast_ab(&new_phi, &new_decoder);
let fit_scale = old_fit
.iter()
.chain(new_fit.iter())
.fold(1.0_f64, |acc, &v| acc.max(v.abs()));
let max_abs = old_fit
.iter()
.zip(new_fit.iter())
.fold(0.0_f64, |acc, (&a, &b)| acc.max((a - b).abs()));
if max_abs > 1.0e-8 * fit_scale {
return Ok(());
}
let flat = Array1::from_iter(new_coords.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let atom = &mut self.atoms[atom_idx];
let base: Arc<dyn SaeBasisEvaluator> = new_evaluator.clone();
atom.basis_evaluator = Some(base);
atom.basis_second_jet = Some(new_evaluator);
let transported_penalty =
transport_smooth_penalty_for_decoder(transport.view(), old_smooth_penalty.view())?;
atom.install_reparameterized_basis(
new_phi,
new_jet,
new_decoder,
transported_penalty,
)?;
Ok(())
}
pub(crate) fn atom_affine_gauge_weights(
&self,
atom_idx: usize,
rho: Option<&SaeManifoldRho>,
) -> Result<Array1<f64>, String> {
let n = self.n_obs();
let mut weights = Array1::<f64>::zeros(n);
let mut scratch = vec![0.0_f64; self.k_atoms()];
for row in 0..n {
match rho {
Some(_) => self
.assignment
.try_assignments_row_into(row, &mut scratch)?,
None => {
let a = self.assignment.try_assignments_row(row)?;
scratch.copy_from_slice(a.as_slice().expect("contiguous assignment row"));
}
};
let assignments = &scratch;
let mut w = assignments[atom_idx].max(0.0);
if let Some(row_weights) = self.row_loss_weights.as_ref() {
w *= row_weights[row].max(0.0);
}
weights[row] = if w.is_finite() { w } else { 0.0 };
}
Ok(weights)
}
pub fn canonicalize_charts_post_fit(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
) -> Result<(), String> {
use crate::chart_canonicalization::{CHART_RECOMPOSITION_REL_TOL, CanonicalChartTopology};
self.assignment.validate_rho_domain(rho)?;
let ard_pre_spread: Vec<Vec<f64>> = (0..self.k_atoms())
.map(|k| {
let coords = self.assignment.coords[k].as_matrix();
let periods = self.assignment.coords[k].effective_axis_periods();
(0..coords.ncols())
.map(|axis| {
axis_coordinate_spread(
coords.view(),
axis,
periods.get(axis).copied().flatten(),
)
})
.collect()
})
.collect();
enum ChartPlan {
UnitSpeed(CanonicalChartTopology),
TorusFlow { period: f64 },
PatchFlow,
SphereFlow,
}
let mut eligible: Vec<(usize, ChartPlan)> = Vec::new();
for atom_idx in 0..self.k_atoms() {
let atom = &self.atoms[atom_idx];
if atom.basis_evaluator.is_none()
|| atom.homotopy_eta != 1.0
|| self.assignment.coords[atom_idx].latent_dim() != atom.latent_dim()
{
continue;
}
let plan = match (atom.basis_kind(), atom.latent_dim()) {
(SaeAtomBasisKind::Periodic | SaeAtomBasisKind::Torus, 1) => {
ChartPlan::UnitSpeed(CanonicalChartTopology::Circle { period: 1.0 })
}
(
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::Duchon
| SaeAtomBasisKind::EuclideanPatch,
1,
) => ChartPlan::UnitSpeed(CanonicalChartTopology::Interval),
(SaeAtomBasisKind::Torus, 2) => ChartPlan::TorusFlow { period: 1.0 },
(
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::Duchon
| SaeAtomBasisKind::EuclideanPatch,
2,
) => ChartPlan::PatchFlow,
(SaeAtomBasisKind::Sphere, 2) => ChartPlan::SphereFlow,
_ => continue,
};
eligible.push((atom_idx, plan));
}
if !eligible.is_empty() {
let snapshot = self.snapshot_mutable_state();
let pre_total = self.penalized_objective_total(target, rho, analytic_penalties, 1.0)?;
let mut any_changed = false;
for (atom_idx, plan) in &eligible {
let outcome = match plan {
ChartPlan::UnitSpeed(topology) => {
self.canonicalize_atom_unit_speed_chart(*atom_idx, topology)
}
ChartPlan::TorusFlow { period } => {
self.canonicalize_atom_torus_flow_chart(*atom_idx, *period)
}
ChartPlan::PatchFlow => self.canonicalize_atom_patch_flow_chart(*atom_idx),
ChartPlan::SphereFlow => self.canonicalize_atom_sphere_flow_chart(*atom_idx),
};
match outcome {
Ok(changed) => any_changed |= changed,
Err(err) => {
self.restore_mutable_state(&snapshot)?;
return Err(err);
}
}
}
if any_changed {
let canonical_total =
self.penalized_objective_total(target, rho, analytic_penalties, 1.0);
let keep = match canonical_total {
Ok(total) => {
total.is_finite()
&& total
<= pre_total + CHART_RECOMPOSITION_REL_TOL * (1.0 + pre_total.abs())
}
Err(_) => false,
};
if !keep {
self.restore_mutable_state(&snapshot)?;
}
}
}
let loao_ev = self
.per_atom_loao_explained_variance(target, rho)
.unwrap_or_else(|err| {
log::warn!("[#1026] per-atom LOAO EV unavailable: {err}");
vec![None; self.k_atoms()]
});
for atom_idx in 0..self.k_atoms() {
let atom = &self.atoms[atom_idx];
if atom.latent_dim() != 1
|| atom.homotopy_eta != 1.0
|| self.assignment.coords[atom_idx].latent_dim() != atom.latent_dim()
{
continue;
}
let Some(evaluator) = atom.basis_evaluator.as_ref().cloned() else {
continue;
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let row_coords = coords.column(0);
let dev = loao_ev
.get(atom_idx)
.copied()
.flatten()
.map_or_else(|| "unavailable".to_string(), |d| format!("{d:.6e}"));
match crate::chart_canonicalization::d1_atom_fitted_turning(
evaluator.as_ref(),
atom.decoder_coefficients().view(),
row_coords,
) {
Ok(Some(theta)) => log::info!(
"[#1026] atom '{}' fitted turning Θ = {theta:.6e} rad, \
training LOAO ΔEV = {dev} \
(∫κ ds; 0 = linear-tail direction, 2π = full curved loop; \
Θ≈0 + large ΔEV = linear direction, high-Θ + large ΔEV = \
genuine curved family — the hybrid-vs-shatter signal)",
atom.name
),
Ok(None) => log::info!(
"[#1026] atom '{}' fitted turning unavailable, training LOAO ΔEV = {dev} \
(no analytic second jet or degenerate curve)",
atom.name
),
Err(err) => {
log::warn!("[#1026] atom '{}' fitted turning errored: {err}", atom.name)
}
}
}
match self.compute_hybrid_split_report(rho, Some(target)) {
Ok(report) => {
if let Some(report) = &report {
log::info!(
"[#1026] hybrid split: {} curved / {} linear atoms (Σ NLE = {:.6e})",
report.selection.curved_atom_count,
report.selection.linear_atom_count(),
report.selection.total_negative_log_evidence,
);
}
self.hybrid_split_report = report;
}
Err(err) => {
log::warn!("[#1026] hybrid split report unavailable: {err}");
self.hybrid_split_report = None;
}
}
let ard_precisions = self.validated_ard_precisions(rho)?;
for atom_idx in 0..self.k_atoms() {
let log_ard = &rho.log_ard[atom_idx];
self.atoms[atom_idx].ard_precisions = if log_ard.is_empty() {
None
} else {
let coords = self.assignment.coords[atom_idx].as_matrix();
let periods = self.assignment.coords[atom_idx].effective_axis_periods();
let pre = &ard_pre_spread[atom_idx];
let stamped: Array1<f64> = (0..log_ard.len())
.map(|axis| {
let alpha = ard_precisions[atom_idx][axis];
let sp_pre = pre.get(axis).copied().unwrap_or(f64::NAN);
let sp_post = axis_coordinate_spread(
coords.view(),
axis,
periods.get(axis).copied().flatten(),
);
if sp_pre.is_finite()
&& sp_post.is_finite()
&& sp_pre > ARD_SPREAD_FLOOR
&& sp_post > ARD_SPREAD_FLOOR
{
alpha * sp_pre / sp_post
} else {
alpha
}
})
.collect();
Some(stamped)
};
}
Ok(())
}
pub(crate) fn canonicalize_atom_unit_speed_chart(
&mut self,
atom_idx: usize,
topology: &crate::chart_canonicalization::CanonicalChartTopology,
) -> Result<bool, String> {
use crate::chart_canonicalization::{CHART_RECOMPOSITION_REL_TOL, unit_speed_retraction};
let n = self.n_obs();
if n == 0 {
return Ok(false);
}
let Some(evaluator) = self.atoms[atom_idx].basis_evaluator.as_ref().cloned() else {
return Ok(false);
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let row_coords = coords.column(0).to_owned();
let Some(repar) = unit_speed_retraction(
evaluator.as_ref(),
self.atoms[atom_idx].decoder_coefficients().view(),
row_coords.view(),
topology,
)?
else {
return Ok(false);
};
let mut new_coords = Array2::<f64>::zeros((n, 1));
for row in 0..n {
new_coords[[row, 0]] = repar.new_row_coords[row];
}
let (new_phi, new_jet) = if self.atoms[atom_idx].homotopy_eta == 1.0 {
evaluator.evaluate(new_coords.view())?
} else {
let evaluated =
evaluator.evaluate_phi_eta(new_coords.view(), self.atoms[atom_idx].homotopy_eta)?;
(evaluated.phi, evaluated.jet)
};
if new_phi.dim() != self.atoms[atom_idx].basis_values.dim()
|| new_jet.dim() != self.atoms[atom_idx].basis_jacobian.dim()
{
return Err(format!(
"SaeManifoldTerm::canonicalize_atom_unit_speed_chart: canonical basis {:?} / jet {:?} must match the fitted shapes {:?} / {:?}",
new_phi.dim(),
new_jet.dim(),
self.atoms[atom_idx].basis_values.dim(),
self.atoms[atom_idx].basis_jacobian.dim()
));
}
let old_fit = fast_ab(
&self.atoms[atom_idx].basis_values,
self.atoms[atom_idx].decoder_coefficients(),
);
let new_fit = fast_ab(&new_phi, &repar.new_decoder);
let mut fit_scale = 0.0_f64;
let mut max_abs = 0.0_f64;
for (a, b) in old_fit.iter().zip(new_fit.iter()) {
fit_scale = fit_scale.max(a.abs()).max(b.abs());
max_abs = max_abs.max((a - b).abs());
}
if !(fit_scale.is_finite() && max_abs.is_finite()) {
return Ok(false);
}
if fit_scale > 0.0 && max_abs > CHART_RECOMPOSITION_REL_TOL * fit_scale {
return Ok(false);
}
let old_smooth_penalty = self.atoms[atom_idx].smooth_penalty().clone();
let flat = Array1::from_iter(new_coords.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let atom = &mut self.atoms[atom_idx];
let transported_penalty = transport_smooth_penalty_for_decoder(
repar.decoder_transport.view(),
old_smooth_penalty.view(),
)?;
atom.install_reparameterized_basis(
new_phi,
new_jet,
repar.new_decoder,
transported_penalty,
)?;
atom.chart_canonicalized = true;
Ok(true)
}
pub(crate) fn d1_unit_speed_topology(
&self,
atom_idx: usize,
) -> Option<crate::chart_canonicalization::CanonicalChartTopology> {
use crate::chart_canonicalization::CanonicalChartTopology;
let atom = &self.atoms[atom_idx];
if atom.basis_evaluator.is_none()
|| atom.homotopy_eta != 1.0
|| self.assignment.coords[atom_idx].latent_dim() != atom.latent_dim()
{
return None;
}
match (atom.basis_kind(), atom.latent_dim()) {
(SaeAtomBasisKind::Periodic | SaeAtomBasisKind::Torus, 1) => {
Some(CanonicalChartTopology::Circle { period: 1.0 })
}
(
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::Duchon
| SaeAtomBasisKind::EuclideanPatch,
1,
) => Some(CanonicalChartTopology::Interval),
_ => None,
}
}
pub(crate) fn retract_unit_speed_charts_in_loop(&mut self) -> Result<usize, String> {
let mut retracted = 0usize;
for atom_idx in 0..self.atoms.len() {
let Some(topology) = self.d1_unit_speed_topology(atom_idx) else {
continue;
};
if self.canonicalize_atom_unit_speed_chart(atom_idx, &topology)? {
retracted += 1;
}
}
Ok(retracted)
}
pub(crate) fn canonicalize_atom_torus_flow_chart(
&mut self,
atom_idx: usize,
period: f64,
) -> Result<bool, String> {
use crate::chart_canonicalization::{
CHART_RECOMPOSITION_REL_TOL, torus_isometry_flow_reparameterization,
};
let n = self.n_obs();
if n == 0 {
return Ok(false);
}
let Some(evaluator) = self.atoms[atom_idx].basis_evaluator.as_ref().cloned() else {
return Ok(false);
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let Some(repar) = torus_isometry_flow_reparameterization(
evaluator.as_ref(),
self.atoms[atom_idx].decoder_coefficients().view(),
coords.view(),
period,
)?
else {
return Ok(false);
};
let new_coords = repar.new_row_coords.clone();
let (new_phi, new_jet) = evaluator.evaluate(new_coords.view())?;
if new_phi.dim() != self.atoms[atom_idx].basis_values.dim()
|| new_jet.dim() != self.atoms[atom_idx].basis_jacobian.dim()
{
return Err(format!(
"SaeManifoldTerm::canonicalize_atom_torus_flow_chart: canonical basis {:?} / jet {:?} must match the fitted shapes {:?} / {:?}",
new_phi.dim(),
new_jet.dim(),
self.atoms[atom_idx].basis_values.dim(),
self.atoms[atom_idx].basis_jacobian.dim()
));
}
let old_fit = fast_ab(
&self.atoms[atom_idx].basis_values,
self.atoms[atom_idx].decoder_coefficients(),
);
let new_fit = fast_ab(&new_phi, &repar.new_decoder);
let mut fit_scale = 0.0_f64;
let mut max_abs = 0.0_f64;
for (a, b) in old_fit.iter().zip(new_fit.iter()) {
fit_scale = fit_scale.max(a.abs()).max(b.abs());
max_abs = max_abs.max((a - b).abs());
}
if !(fit_scale.is_finite() && max_abs.is_finite()) {
return Ok(false);
}
if fit_scale > 0.0 && max_abs > CHART_RECOMPOSITION_REL_TOL * fit_scale {
return Ok(false);
}
let old_smooth_penalty = self.atoms[atom_idx].smooth_penalty().clone();
let flat = Array1::from_iter(new_coords.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let atom = &mut self.atoms[atom_idx];
let transported_penalty = transport_smooth_penalty_for_decoder(
repar.decoder_transport.view(),
old_smooth_penalty.view(),
)?;
atom.install_reparameterized_basis(
new_phi,
new_jet,
repar.new_decoder,
transported_penalty,
)?;
atom.chart_canonicalized = true;
Ok(true)
}
pub(crate) fn canonicalize_atom_patch_flow_chart(
&mut self,
atom_idx: usize,
) -> Result<bool, String> {
use crate::chart_canonicalization::{
CHART_RECOMPOSITION_REL_TOL, patch_isometry_flow_reparameterization,
};
let n = self.n_obs();
if n == 0 {
return Ok(false);
}
let Some(evaluator) = self.atoms[atom_idx].basis_evaluator.as_ref().cloned() else {
return Ok(false);
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let Some(repar) = patch_isometry_flow_reparameterization(
evaluator.as_ref(),
self.atoms[atom_idx].decoder_coefficients().view(),
coords.view(),
)?
else {
return Ok(false);
};
let new_coords = repar.new_row_coords.clone();
let (new_phi, new_jet) = evaluator.evaluate(new_coords.view())?;
if new_phi.dim() != self.atoms[atom_idx].basis_values.dim()
|| new_jet.dim() != self.atoms[atom_idx].basis_jacobian.dim()
{
return Err(format!(
"SaeManifoldTerm::canonicalize_atom_patch_flow_chart: canonical basis {:?} / jet {:?} must match the fitted shapes {:?} / {:?}",
new_phi.dim(),
new_jet.dim(),
self.atoms[atom_idx].basis_values.dim(),
self.atoms[atom_idx].basis_jacobian.dim()
));
}
let old_fit = fast_ab(
&self.atoms[atom_idx].basis_values,
self.atoms[atom_idx].decoder_coefficients(),
);
let new_fit = fast_ab(&new_phi, &repar.new_decoder);
let mut fit_scale = 0.0_f64;
let mut max_abs = 0.0_f64;
for (a, b) in old_fit.iter().zip(new_fit.iter()) {
fit_scale = fit_scale.max(a.abs()).max(b.abs());
max_abs = max_abs.max((a - b).abs());
}
if !(fit_scale.is_finite() && max_abs.is_finite()) {
return Ok(false);
}
if fit_scale > 0.0 && max_abs > CHART_RECOMPOSITION_REL_TOL * fit_scale {
return Ok(false);
}
let old_smooth_penalty = self.atoms[atom_idx].smooth_penalty().clone();
let flat = Array1::from_iter(new_coords.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let atom = &mut self.atoms[atom_idx];
let transported_penalty = transport_smooth_penalty_for_decoder(
repar.decoder_transport.view(),
old_smooth_penalty.view(),
)?;
atom.install_reparameterized_basis(
new_phi,
new_jet,
repar.new_decoder,
transported_penalty,
)?;
atom.chart_canonicalized = true;
Ok(true)
}
pub(crate) fn canonicalize_atom_sphere_flow_chart(
&mut self,
atom_idx: usize,
) -> Result<bool, String> {
use crate::chart_canonicalization::{
CHART_RECOMPOSITION_REL_TOL, sphere_isometry_flow_reparameterization,
};
let n = self.n_obs();
if n == 0 {
return Ok(false);
}
let Some(evaluator) = self.atoms[atom_idx].basis_evaluator.as_ref().cloned() else {
return Ok(false);
};
let coords = self.assignment.coords[atom_idx].as_matrix();
let Some(repar) = sphere_isometry_flow_reparameterization(
evaluator.as_ref(),
self.atoms[atom_idx].decoder_coefficients().view(),
coords.view(),
)?
else {
return Ok(false);
};
let new_coords = repar.new_row_coords.clone();
let (new_phi, new_jet) = evaluator.evaluate(new_coords.view())?;
if new_phi.dim() != self.atoms[atom_idx].basis_values.dim()
|| new_jet.dim() != self.atoms[atom_idx].basis_jacobian.dim()
{
return Err(format!(
"SaeManifoldTerm::canonicalize_atom_sphere_flow_chart: canonical basis {:?} / jet {:?} must match the fitted shapes {:?} / {:?}",
new_phi.dim(),
new_jet.dim(),
self.atoms[atom_idx].basis_values.dim(),
self.atoms[atom_idx].basis_jacobian.dim()
));
}
let old_fit = fast_ab(
&self.atoms[atom_idx].basis_values,
self.atoms[atom_idx].decoder_coefficients(),
);
let new_fit = fast_ab(&new_phi, &repar.new_decoder);
let mut fit_scale = 0.0_f64;
let mut max_abs = 0.0_f64;
for (a, b) in old_fit.iter().zip(new_fit.iter()) {
fit_scale = fit_scale.max(a.abs()).max(b.abs());
max_abs = max_abs.max((a - b).abs());
}
if !(fit_scale.is_finite() && max_abs.is_finite()) {
return Ok(false);
}
if fit_scale > 0.0 && max_abs > CHART_RECOMPOSITION_REL_TOL * fit_scale {
return Ok(false);
}
let old_smooth_penalty = self.atoms[atom_idx].smooth_penalty().clone();
let flat = Array1::from_iter(new_coords.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let atom = &mut self.atoms[atom_idx];
let transported_penalty = transport_smooth_penalty_for_decoder(
repar.decoder_transport.view(),
old_smooth_penalty.view(),
)?;
atom.install_reparameterized_basis(
new_phi,
new_jet,
repar.new_decoder,
transported_penalty,
)?;
atom.chart_canonicalized = true;
Ok(true)
}
pub(crate) fn inner_iterate_scale(&self) -> f64 {
let mut iterate_norm_sq = 0.0_f64;
for &v in self.assignment.logits.iter() {
iterate_norm_sq += v * v;
}
for coords in &self.assignment.coords {
let matrix = coords.as_matrix();
for &v in matrix.iter() {
iterate_norm_sq += v * v;
}
}
for atom in &self.atoms {
for &v in atom.decoder_coefficients().iter() {
iterate_norm_sq += v * v;
}
}
1.0 + iterate_norm_sq.sqrt()
}
pub(crate) fn inner_iterate_max(&self) -> Result<f64, SaeInnerKktScaleError> {
let mut max_abs = 0.0_f64;
for (component, &value) in self.assignment.logits.iter().enumerate() {
if !value.is_finite() {
return Err(SaeInnerKktScaleError::NonFiniteIterate {
family: "assignment-logit",
group: 0,
component,
value,
});
}
max_abs = max_abs.max(value.abs());
}
for (atom, coords) in self.assignment.coords.iter().enumerate() {
let matrix = coords.as_matrix();
for (component, &value) in matrix.iter().enumerate() {
if !value.is_finite() {
return Err(SaeInnerKktScaleError::NonFiniteIterate {
family: "coordinate",
group: atom,
component,
value,
});
}
max_abs = max_abs.max(value.abs());
}
}
for (atom, manifold_atom) in self.atoms.iter().enumerate() {
for (component, &value) in manifold_atom.decoder_coefficients().iter().enumerate() {
if !value.is_finite() {
return Err(SaeInnerKktScaleError::NonFiniteIterate {
family: "decoder",
group: atom,
component,
value,
});
}
max_abs = max_abs.max(value.abs());
}
}
let scale = 1.0 + max_abs;
if !scale.is_finite() {
return Err(SaeInnerKktScaleError::IterateScaleOverflow { max_abs });
}
Ok(scale)
}
fn machine_null_eigenvectors(
mut operator: Array2<f64>,
context: &str,
) -> Result<Vec<Array1<f64>>, String> {
let dim = operator.nrows();
if operator.ncols() != dim {
return Err(format!(
"{context}: decoder operator must be square, got {:?}",
operator.dim()
));
}
if dim == 0 {
return Ok(Vec::new());
}
for row in 0..dim {
for col in 0..row {
let sym = 0.5 * (operator[[row, col]] + operator[[col, row]]);
operator[[row, col]] = sym;
operator[[col, row]] = sym;
}
}
if !operator.iter().all(|value| value.is_finite()) {
return Err(format!("{context}: non-finite decoder operator entry"));
}
let (evals, evecs) = operator
.eigh(Side::Lower)
.map_err(|err| format!("{context}: eigh failed: {err}"))?;
let operator_norm = evals
.iter()
.fold(0.0_f64, |scale, &value| scale.max(value.abs()));
let d_eps = dim as f64 * f64::EPSILON;
let gamma_d = if d_eps < 1.0 {
d_eps / (1.0 - d_eps)
} else {
return Err(format!(
"{context}: decoder operator dimension {dim} exceeds the f64 backward-error domain"
));
};
let null_floor = gamma_d * operator_norm;
let mut out = Vec::new();
for eig_idx in 0..dim {
let eigenvalue = evals[eig_idx];
if !(eigenvalue.is_finite() && eigenvalue.abs() <= null_floor) {
continue;
}
let direction = evecs.column(eig_idx).to_owned();
let applied = operator.dot(&direction);
let rayleigh = direction.dot(&applied).abs();
let residual_norm = applied
.iter()
.map(|value| value * value)
.sum::<f64>()
.sqrt();
if rayleigh.is_finite()
&& residual_norm.is_finite()
&& rayleigh <= null_floor
&& residual_norm <= null_floor
{
out.push(direction);
}
}
Ok(out)
}
pub(crate) fn joint_decoder_beta_null_directions(
&self,
penalized_gram_scale: &[f64],
) -> Result<Vec<Array1<f64>>, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let p = self.output_dim();
let k_atoms = self.k_atoms();
if penalized_gram_scale.len() != k_atoms {
return Err(format!(
"joint_decoder_beta_null_directions: {} smooth scales for {k_atoms} atoms",
penalized_gram_scale.len()
));
}
let basis_sizes: Vec<usize> = self.atoms.iter().map(|atom| atom.basis_size()).collect();
let mut basis_offsets = Vec::with_capacity(k_atoms);
let mut basis_dim = 0usize;
for &m in &basis_sizes {
basis_offsets.push(basis_dim);
basis_dim += m;
}
let border_dim = self.factored_border_dim();
if p == 0 || basis_dim == 0 || border_dim == 0 {
return Ok(Vec::new());
}
let assignments = self.assignment.assignments();
let mut joint_basis = Array2::<f64>::zeros((basis_dim, basis_dim));
let mut weighted_basis = vec![0.0_f64; basis_dim];
for row in 0..n {
for atom_idx in 0..k_atoms {
let atom = &self.atoms[atom_idx];
let off = basis_offsets[atom_idx];
let weight = assignments[[row, atom_idx]];
for basis_col in 0..basis_sizes[atom_idx] {
weighted_basis[off + basis_col] = weight * atom.basis_values[[row, basis_col]];
}
}
for col in 0..basis_dim {
let value = weighted_basis[col];
if value == 0.0 {
continue;
}
for row_idx in 0..basis_dim {
joint_basis[[row_idx, col]] += weighted_basis[row_idx] * value;
}
}
}
let joint_data_basis = joint_basis.clone();
for atom_idx in 0..k_atoms {
let m = basis_sizes[atom_idx];
let penalty = self.atoms[atom_idx].smooth_penalty();
if penalty.dim() != (m, m) {
return Err(format!(
"joint_decoder_beta_null_directions: atom {atom_idx} penalty shape {:?} != ({m}, {m})",
penalty.dim()
));
}
let off = basis_offsets[atom_idx];
let scale = penalized_gram_scale[atom_idx];
if !scale.is_finite() || scale < 0.0 {
return Err(format!(
"joint_decoder_beta_null_directions: atom {atom_idx} smooth scale must be finite and nonnegative, got {scale}"
));
}
for row in 0..m {
for col in 0..m {
joint_basis[[off + row, off + col]] += scale * penalty[[row, col]];
}
}
}
let coord_base = n * q;
if !self.any_frame_active() {
let null_basis = Self::machine_null_eigenvectors(
joint_basis,
"joint_decoder_beta_null_directions(full-B)",
)?;
let mut out = Vec::with_capacity(null_basis.len() * p);
for basis_direction in null_basis {
for out_col in 0..p {
let mut direction = Array1::<f64>::zeros(coord_base + border_dim);
for basis_col in 0..basis_dim {
direction[coord_base + basis_col * p + out_col] =
basis_direction[basis_col];
}
out.push(direction);
}
}
return Ok(out);
}
let border_offsets = self.factored_border_offsets();
let frame_ranks: Vec<usize> = self
.atoms
.iter()
.map(SaeManifoldAtom::border_frame_rank)
.collect();
let mut joint_border = Array2::<f64>::zeros((border_dim, border_dim));
let whitens_likelihood = self
.row_metric
.as_ref()
.is_some_and(|metric| metric.whitens_likelihood());
if whitens_likelihood {
let metric = self
.row_metric
.as_ref()
.expect("whitens_likelihood implies a row metric");
let metric_rank = metric.metric_rank();
let frames: Vec<Array2<f64>> = (0..k_atoms)
.map(|atom_idx| self.frame_output_matrix(atom_idx))
.collect();
let mut whitened_jacobian = Array2::<f64>::zeros((metric_rank, border_dim));
for row in 0..n {
whitened_jacobian.fill(0.0);
for atom_idx in 0..k_atoms {
let m = basis_sizes[atom_idx];
let rank = frame_ranks[atom_idx];
let border_off = border_offsets[atom_idx];
let weight = assignments[[row, atom_idx]];
for frame_col in 0..rank {
for metric_col in 0..metric_rank {
let mut projected = 0.0_f64;
for out_col in 0..p {
projected += frames[atom_idx][[out_col, frame_col]]
* metric.factor_entry(row, out_col, metric_col);
}
if projected == 0.0 {
continue;
}
for basis_col in 0..m {
whitened_jacobian
[[metric_col, border_off + basis_col * rank + frame_col]] =
weight
* self.atoms[atom_idx].basis_values[[row, basis_col]]
* projected;
}
}
}
}
for col in 0..border_dim {
for row_idx in 0..border_dim {
let mut value = 0.0_f64;
for metric_col in 0..metric_rank {
value += whitened_jacobian[[metric_col, row_idx]]
* whitened_jacobian[[metric_col, col]];
}
joint_border[[row_idx, col]] += value;
}
}
}
} else {
for atom_j in 0..k_atoms {
let mj = basis_sizes[atom_j];
let rj = frame_ranks[atom_j];
let basis_j = basis_offsets[atom_j];
let border_j = border_offsets[atom_j];
for atom_k in 0..k_atoms {
let mk = basis_sizes[atom_k];
let rk = frame_ranks[atom_k];
let basis_k = basis_offsets[atom_k];
let border_k = border_offsets[atom_k];
let frame_overlap = self.frame_cross_factor(atom_j, atom_k);
for col_j in 0..mj {
for col_k in 0..mk {
let gram = joint_data_basis[[basis_j + col_j, basis_k + col_k]];
for channel_j in 0..rj {
for channel_k in 0..rk {
joint_border[[
border_j + col_j * rj + channel_j,
border_k + col_k * rk + channel_k,
]] += gram * frame_overlap[[channel_j, channel_k]];
}
}
}
}
}
}
}
for atom_idx in 0..k_atoms {
let m = basis_sizes[atom_idx];
let rank = frame_ranks[atom_idx];
let off = border_offsets[atom_idx];
let penalty = self.atoms[atom_idx].smooth_penalty();
let scale = penalized_gram_scale[atom_idx];
for basis_row in 0..m {
for basis_col in 0..m {
let value = scale * penalty[[basis_row, basis_col]];
for channel in 0..rank {
joint_border[[
off + basis_row * rank + channel,
off + basis_col * rank + channel,
]] += value;
}
}
}
}
let null_border = Self::machine_null_eigenvectors(
joint_border,
"joint_decoder_beta_null_directions(factored)",
)?;
Ok(null_border
.into_iter()
.map(|beta_direction| {
let mut direction = Array1::<f64>::zeros(coord_base + border_dim);
direction
.slice_mut(s![coord_base..])
.assign(&beta_direction);
direction
})
.collect())
}
pub(crate) fn decoder_channel_null_directions(&self) -> Result<Vec<Array1<f64>>, String> {
let p = self.output_dim();
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let border_dim = self.factored_border_dim();
let total_len = n * q + border_dim;
if p == 0 || border_dim == 0 {
return Ok(Vec::new());
}
let beta_offsets = self.factored_border_offsets();
let mut out = Vec::new();
for atom_idx in 0..self.k_atoms() {
let atom = &self.atoms[atom_idx];
if atom.decoder_frame.is_some() {
continue;
}
let m = atom.basis_size();
if m == 0 {
continue;
}
let (_u, sv, vt_opt) = match atom.decoder_coefficients().svd(false, true) {
Ok(parts) => parts,
Err(_) => continue,
};
let Some(vt) = vt_opt else {
continue;
};
let max_sv = sv.iter().fold(0.0_f64, |acc, &v| acc.max(v));
let sv_floor = SAE_DECODER_BETA_NULL_RELATIVE_FLOOR.sqrt() * max_sv;
let beta_base = n * q + beta_offsets[atom_idx];
for c_idx in 0..vt.nrows() {
let realised = c_idx < sv.len() && max_sv > 0.0 && sv[c_idx] > sv_floor;
if realised {
continue;
}
let channel = vt.row(c_idx);
if channel.len() != p {
continue;
}
let norm_sq = channel.iter().map(|v| v * v).sum::<f64>();
if !(norm_sq.is_finite() && norm_sq > 1.0e-24) {
continue;
}
for col in 0..m {
let mut dir = Array1::<f64>::zeros(total_len);
for out_col in 0..p {
dir[beta_base + col * p + out_col] = channel[out_col];
}
out.push(dir);
}
}
}
Ok(out)
}
pub(crate) fn quotient_newton_step_norm_sq(
&self,
delta_ext_coord: ArrayView1<'_, f64>,
delta_beta: ArrayView1<'_, f64>,
raw_step_norm_sq: f64,
penalized_gram_scale: &[f64],
) -> Result<f64, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let border_dim = self.factored_border_dim();
if delta_ext_coord.len() != n * q || delta_beta.len() != border_dim {
return Ok(raw_step_norm_sq);
}
let mut residual = Array1::<f64>::zeros(delta_ext_coord.len() + delta_beta.len());
for i in 0..delta_ext_coord.len() {
residual[i] = delta_ext_coord[i];
}
let beta_base = delta_ext_coord.len();
for i in 0..delta_beta.len() {
residual[beta_base + i] = delta_beta[i];
}
let quotient = self.quotient_residual_norm_sq(residual, penalized_gram_scale)?;
Ok(if quotient.is_finite() {
quotient.max(0.0).min(raw_step_norm_sq)
} else {
raw_step_norm_sq
})
}
pub(crate) fn quotient_residual_norm_sq(
&self,
mut residual: Array1<f64>,
penalized_gram_scale: &[f64],
) -> Result<f64, String> {
for basis in self.posterior_null_quotient_basis(penalized_gram_scale)? {
if basis.len() != residual.len() {
continue;
}
let coeff = residual.dot(&basis);
for i in 0..residual.len() {
residual[i] -= coeff * basis[i];
}
}
Ok(residual.iter().map(|v| v * v).sum::<f64>())
}
pub(crate) fn decoder_channel_null_quotient_directions(
&self,
penalized_gram_scale: &[f64],
) -> Result<Vec<Array1<f64>>, String> {
let k_atoms = self.k_atoms();
if penalized_gram_scale.len() != k_atoms {
return Err(format!(
"decoder_channel_null_quotient_directions: {} smooth scales for {k_atoms} atoms",
penalized_gram_scale.len()
));
}
let slope_bound = SAE_MANIFOLD_INNER_GRAD_REL_TOL * self.inner_iterate_scale();
let p = self.output_dim();
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let border_dim = self.factored_border_dim();
let total_len = n * q + border_dim;
if p == 0 || border_dim == 0 {
return Ok(Vec::new());
}
let beta_offsets = self.factored_border_offsets();
let mut out = Vec::new();
for atom_idx in 0..k_atoms {
let atom = &self.atoms[atom_idx];
if atom.decoder_frame.is_some() {
continue;
}
let m = atom.basis_size();
if m == 0 {
continue;
}
let lambda = penalized_gram_scale[atom_idx];
if !lambda.is_finite() || lambda < 0.0 {
return Err(format!(
"decoder_channel_null_quotient_directions: atom {atom_idx} smooth scale must \
be finite and nonnegative, got {lambda}"
));
}
let (_u, sv, vt_opt) = match atom.decoder_coefficients().svd(false, true) {
Ok(parts) => parts,
Err(_) => continue,
};
let Some(vt) = vt_opt else {
continue;
};
let max_sv = sv.iter().fold(0.0_f64, |acc, &v| acc.max(v));
let sv_floor = SAE_DECODER_BETA_NULL_RELATIVE_FLOOR.sqrt() * max_sv;
let beta_base = n * q + beta_offsets[atom_idx];
let s_gram = atom.smooth_penalty();
if s_gram.dim() != (m, m) {
return Err(format!(
"decoder_channel_null_quotient_directions: atom {atom_idx} penalty shape \
{:?} != ({m}, {m})",
s_gram.dim()
));
}
for c_idx in 0..vt.nrows() {
let realised = c_idx < sv.len() && max_sv > 0.0 && sv[c_idx] > sv_floor;
if realised {
continue;
}
let channel = vt.row(c_idx);
if channel.len() != p {
continue;
}
let norm_sq = channel.iter().map(|v| v * v).sum::<f64>();
if !(norm_sq.is_finite() && norm_sq > 1.0e-24) {
continue;
}
let mut b_c = Array1::<f64>::zeros(m);
for row in 0..m {
let mut acc = 0.0_f64;
for out_col in 0..p {
acc += atom.decoder_coefficients()[[row, out_col]] * channel[out_col];
}
b_c[row] = acc;
}
let mut s_b_c = Array1::<f64>::zeros(m);
for row in 0..m {
let mut acc = 0.0_f64;
for col in 0..m {
acc += s_gram[[row, col]] * b_c[col];
}
s_b_c[row] = acc;
}
for col in 0..m {
let slope = (lambda * s_b_c[col]).abs();
if slope > slope_bound {
continue;
}
let mut dir = Array1::<f64>::zeros(total_len);
for out_col in 0..p {
dir[beta_base + col * p + out_col] = channel[out_col];
}
out.push(dir);
}
}
}
Ok(out)
}
pub(crate) fn posterior_null_quotient_basis(
&self,
penalized_gram_scale: &[f64],
) -> Result<Vec<Array1<f64>>, String> {
Ok(Self::orthonormalized(
self.joint_decoder_beta_null_directions(penalized_gram_scale)?
.into_iter()
.chain(self.decoder_channel_null_quotient_directions(
penalized_gram_scale,
)?),
))
}
pub(crate) fn likelihood_flat_block_basis(
&self,
penalized_gram_scale: &[f64],
) -> Result<Vec<Array1<f64>>, String> {
Ok(Self::orthonormalized(
self.dense_step_gauge_vectors()?
.into_iter()
.chain(self.joint_decoder_beta_null_directions(penalized_gram_scale)?)
.chain(self.decoder_channel_null_directions()?),
))
}
fn orthonormalized(candidates: impl Iterator<Item = Array1<f64>>) -> Vec<Array1<f64>> {
let mut orthonormal: Vec<Array1<f64>> = Vec::new();
for mut candidate in candidates {
for basis in &orthonormal {
let coeff = candidate.dot(basis);
for i in 0..candidate.len() {
candidate[i] -= coeff * basis[i];
}
}
let norm_sq = candidate.iter().map(|v| v * v).sum::<f64>();
if norm_sq <= 1.0e-24 || !norm_sq.is_finite() {
continue;
}
let inv_norm = norm_sq.sqrt().recip();
for v in candidate.iter_mut() {
*v *= inv_norm;
}
orthonormal.push(candidate);
}
orthonormal
}
pub(crate) fn descend_gauge_orbit(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
registry: Option<&AnalyticPenaltyRegistry>,
penalized_gram_scale: &[f64],
max_rounds: usize,
) -> Result<GaugeOrbitDescent, String> {
let mut outcome = GaugeOrbitDescent::default();
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let dense_len = n.saturating_mul(q);
let border_dim = self.factored_border_dim();
if dense_len + border_dim == 0 {
return Ok(outcome);
}
for _ in 0..max_rounds {
let system = match self.assemble_arrow_schur(target, rho, registry) {
Ok(system) => system,
Err(_) => return Ok(outcome),
};
if system.rows.len() != n
|| system.row_offsets.len() != n + 1
|| system.gb.len() != border_dim
{
return Ok(outcome);
}
let mut gradient = Array1::<f64>::zeros(dense_len + border_dim);
for (row_index, row) in system.rows.iter().enumerate() {
let base = system.row_offsets[row_index];
let dim = system.row_dims[row_index];
if base + dim > dense_len || row.gt.len() < dim {
return Ok(outcome);
}
for axis in 0..dim {
gradient[base + axis] = row.gt[axis];
}
}
for (index, &value) in system.gb.iter().enumerate() {
gradient[dense_len + index] = value;
}
drop(system);
if !gradient.iter().all(|value| value.is_finite()) {
return Ok(outcome);
}
let basis = self.likelihood_flat_block_basis(penalized_gram_scale)?;
outcome.dimension = basis.len();
let mut direction = Array1::<f64>::zeros(gradient.len());
let mut max_directional = 0.0_f64;
for vector in &basis {
if vector.len() != gradient.len() {
continue;
}
let coeff = gradient.dot(vector);
max_directional = max_directional.max(coeff.abs());
for index in 0..direction.len() {
direction[index] -= coeff * vector[index];
}
}
outcome.max_directional_derivative = max_directional;
let slope = direction.dot(&direction).sqrt();
if !(slope.is_finite() && slope > 0.0) {
return Ok(outcome);
}
for value in direction.iter_mut() {
*value /= slope;
}
let base_objective = self.penalized_objective_total(target, rho, registry, 1.0)?;
if outcome.entry_objective.is_none() {
outcome.entry_objective = Some(base_objective);
}
if !base_objective.is_finite() {
return Ok(outcome);
}
let material_floor =
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL * (1.0 + base_objective.abs());
let far_alpha = self.inner_iterate_scale();
let near_alpha = material_floor / slope;
if !(far_alpha.is_finite() && far_alpha > 0.0 && near_alpha.is_finite()) {
return Ok(outcome);
}
let snapshot = self.snapshot_mutable_state();
let evaluate = |term: &mut Self, alpha: f64| -> Result<f64, String> {
if !(alpha.is_finite() && alpha > 0.0) {
return Ok(f64::INFINITY);
}
let value = term
.apply_newton_step(
direction.slice(s![..dense_len]),
direction.slice(s![dense_len..]),
alpha,
)
.and_then(|()| term.penalized_objective_total(target, rho, registry, 1.0))
.unwrap_or(f64::INFINITY);
term.restore_mutable_state(&snapshot).map_err(|err| {
format!(
"SaeManifoldTerm::descend_gauge_orbit: restoring the pre-round state \
after the speculative trial at alpha={alpha:.6e} failed: {err}"
)
})?;
Ok(if value.is_finite() {
value
} else {
f64::INFINITY
})
};
let mut best_alpha = 0.0_f64;
let mut best_value = base_objective;
let mut alpha = far_alpha;
while alpha >= near_alpha {
outcome.evaluations += 1;
let value = evaluate(self, alpha)?;
if value < best_value {
best_value = value;
best_alpha = alpha;
}
alpha *= 0.5;
}
if best_alpha > 0.0 {
const GOLDEN_RATIO_INVERSE: f64 = 0.618_033_988_749_894_9;
let mut low = best_alpha * 0.5;
let mut high = best_alpha * 2.0;
let resolution = f64::EPSILON.sqrt() * best_alpha;
let mut inner_low = high - GOLDEN_RATIO_INVERSE * (high - low);
let mut inner_high = low + GOLDEN_RATIO_INVERSE * (high - low);
let mut value_low = evaluate(self, inner_low)?;
let mut value_high = evaluate(self, inner_high)?;
outcome.evaluations += 2;
while high - low > resolution {
if value_low <= value_high {
high = inner_high;
inner_high = inner_low;
value_high = value_low;
inner_low = high - GOLDEN_RATIO_INVERSE * (high - low);
value_low = evaluate(self, inner_low)?;
} else {
low = inner_low;
inner_low = inner_high;
value_low = value_high;
inner_high = low + GOLDEN_RATIO_INVERSE * (high - low);
value_high = evaluate(self, inner_high)?;
}
outcome.evaluations += 1;
}
if value_low < best_value {
best_value = value_low;
best_alpha = inner_low;
}
if value_high < best_value {
best_value = value_high;
best_alpha = inner_high;
}
}
let trial_decrease = base_objective - best_value;
if !(best_alpha > 0.0 && trial_decrease > material_floor) {
self.restore_mutable_state(&snapshot)
.map_err(|err| format!("SaeManifoldTerm::descend_gauge_orbit: {err}"))?;
return Ok(outcome);
}
if let Err(err) = self.apply_newton_step(
direction.slice(s![..dense_len]),
direction.slice(s![dense_len..]),
best_alpha,
) {
self.restore_mutable_state(&snapshot).map_err(|restore_err| {
format!(
"SaeManifoldTerm::descend_gauge_orbit: committed step application failed \
({err}); restoring the pre-round state also failed ({restore_err})"
)
})?;
return Err(format!(
"SaeManifoldTerm::descend_gauge_orbit: committed step application: {err}"
));
}
let committed_objective =
match self.penalized_objective_total(target, rho, registry, 1.0) {
Ok(value) => value,
Err(err) => {
self.restore_mutable_state(&snapshot).map_err(|restore_err| {
format!(
"SaeManifoldTerm::descend_gauge_orbit: committed objective \
evaluation failed ({err}); restoring the pre-round state also \
failed ({restore_err})"
)
})?;
return Err(format!(
"SaeManifoldTerm::descend_gauge_orbit: committed objective \
evaluation: {err}"
));
}
};
let decrease = base_objective - committed_objective;
if !(committed_objective.is_finite() && decrease > material_floor) {
self.restore_mutable_state(&snapshot)
.map_err(|err| format!("SaeManifoldTerm::descend_gauge_orbit: {err}"))?;
return Ok(outcome);
}
outcome.rounds += 1;
outcome.objective_decrease += decrease;
outcome.exit_objective = Some(committed_objective);
log::debug!(
"SAE gauge-orbit descent: round {} committed {decrease:.6e} at α={best_alpha:.6e} \
(objective {base_objective:.9e} → {committed_objective:.9e}, span dim {}, \
maxᵢ|gᵀvᵢ|={max_directional:.6e}, floor {material_floor:.6e})",
outcome.rounds,
outcome.dimension,
);
}
Ok(outcome)
}
pub(crate) fn quotient_gradient_norm_sq(
&self,
grad_ext_coord: ArrayView1<'_, f64>,
grad_beta: ArrayView1<'_, f64>,
raw_grad_norm_sq: f64,
penalized_gram_scale: &[f64],
) -> Result<f64, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let border_dim = self.factored_border_dim();
if grad_ext_coord.len() != n * q || grad_beta.len() != border_dim {
return Ok(raw_grad_norm_sq);
}
let mut residual = Array1::<f64>::zeros(grad_ext_coord.len() + grad_beta.len());
for i in 0..grad_ext_coord.len() {
residual[i] = grad_ext_coord[i];
}
let beta_base = grad_ext_coord.len();
for i in 0..grad_beta.len() {
residual[beta_base + i] = grad_beta[i];
}
let quotient = self.quotient_residual_norm_sq(residual, penalized_gram_scale)?;
Ok(if quotient.is_finite() {
quotient.max(0.0).min(raw_grad_norm_sq)
} else {
raw_grad_norm_sq
})
}
pub(crate) fn quotient_gradient_norm_from_system(
&self,
sys: &ArrowSchurSystem,
raw_grad_norm_sq: f64,
penalized_gram_scale: &[f64],
) -> f64 {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let dense_len = n.saturating_mul(q);
let mut grad_ext_coord = Array1::<f64>::zeros(dense_len);
let mut dense_layout_ok = sys.rows.len() == n && sys.row_offsets.len() == n + 1;
if dense_layout_ok {
for (row_idx, row) in sys.rows.iter().enumerate() {
let base = sys.row_offsets[row_idx];
let di = sys.row_dims[row_idx];
if base + di > dense_len || row.gt.len() < di {
dense_layout_ok = false;
break;
}
for axis in 0..di {
grad_ext_coord[base + axis] = row.gt[axis];
}
}
}
let raw_grad_norm = raw_grad_norm_sq.sqrt();
if dense_layout_ok {
self.quotient_gradient_norm_sq(
grad_ext_coord.view(),
sys.gb.view(),
raw_grad_norm_sq,
penalized_gram_scale,
)
.map(|v| v.sqrt())
.unwrap_or(raw_grad_norm)
} else {
raw_grad_norm
}
}
pub(crate) fn dense_step_gauge_vectors(&self) -> Result<Vec<Array1<f64>>, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let p = self.output_dim();
let coord_offsets = self.assignment.coord_offsets();
let beta_offsets = self.factored_border_offsets();
let total_len = n * q + self.factored_border_dim();
let mut out = Vec::new();
for atom_idx in 0..self.k_atoms() {
let d = self.assignment.coords[atom_idx].latent_dim();
let coords = self.assignment.coords[atom_idx].as_matrix();
match self.atoms[atom_idx].basis_kind() {
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::EuclideanPatch
| SaeAtomBasisKind::Poincare => {
for axis in 0..d {
let mut field = Array2::<f64>::zeros((n, d));
field.column_mut(axis).fill(1.0);
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
for axis in 0..d {
let mut field = Array2::<f64>::zeros((n, d));
for row in 0..n {
field[[row, axis]] = coords[[row, axis]];
}
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
}
SaeAtomBasisKind::Duchon => {
for axis in 0..d {
let mut field = Array2::<f64>::zeros((n, d));
field.column_mut(axis).fill(1.0);
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
for axis in 0..d {
let mut field = Array2::<f64>::zeros((n, d));
for row in 0..n {
field[[row, axis]] = coords[[row, axis]];
}
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
}
SaeAtomBasisKind::Periodic | SaeAtomBasisKind::Torus => {
for axis in 0..d {
let mut field = Array2::<f64>::zeros((n, d));
field.column_mut(axis).fill(1.0);
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
}
SaeAtomBasisKind::KleinBottle => {
if d != 2 {
return Err(format!(
"dense_step_gauge_vectors: Klein atom {atom_idx} requires latent dimension 2, got {d}"
));
}
let mut field = Array2::<f64>::zeros((n, d));
field.column_mut(0).fill(1.0);
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
SaeAtomBasisKind::Sphere | SaeAtomBasisKind::ProjectivePlane => {
if d != 3 {
return Err(format!(
"dense_step_gauge_vectors: spherical atom {atom_idx} rides the ambient cover and requires latent dimension 3, got {d}"
));
}
let mut fields = [
Array2::<f64>::zeros((n, d)),
Array2::<f64>::zeros((n, d)),
Array2::<f64>::zeros((n, d)),
];
for row in 0..n {
let directions = ambient_sphere_killing_directions([
coords[[row, 0]],
coords[[row, 1]],
coords[[row, 2]],
]);
for generator in 0..3 {
for axis in 0..3 {
fields[generator][[row, axis]] = directions[generator][axis];
}
}
}
for field in fields {
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
}
SaeAtomBasisKind::Cylinder => {
let mut field = Array2::<f64>::zeros((n, d));
if d > 0 {
field.column_mut(0).fill(1.0);
}
if let Some(g) = self.dense_step_gauge_vector_from_field(
atom_idx,
field.view(),
&coord_offsets,
&beta_offsets,
total_len,
)? {
out.push(g);
}
}
SaeAtomBasisKind::Mobius
| SaeAtomBasisKind::FiniteSet
| SaeAtomBasisKind::Precomputed(_) => {}
}
}
if p == 0 {
return Ok(Vec::new());
}
Ok(out)
}
pub(crate) fn dense_joint_vector_in_arrow_layout(
&self,
dense: ArrayView1<'_, f64>,
row_offsets: &[usize],
border_dim: usize,
owner: &str,
) -> Result<Array1<f64>, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let declared_border = self.factored_border_dim();
if border_dim != declared_border {
return Err(format!(
"{owner}: arrow border dimension {border_dim} != term border dimension {declared_border}"
));
}
if row_offsets.len() != n + 1 || row_offsets.first() != Some(&0) {
return Err(format!(
"{owner}: arrow row offsets must have length {} and start at zero, got {:?}",
n + 1,
row_offsets
));
}
if let Some(layout) = self.last_row_layout.as_ref() {
return layout.restrict_dense_joint_vector(
dense,
q,
row_offsets,
border_dim,
owner,
);
}
for row in 0..n {
let span = row_offsets[row + 1]
.checked_sub(row_offsets[row])
.ok_or_else(|| format!("{owner}: arrow row offsets decrease at row {row}"))?;
if span != q {
return Err(format!(
"{owner}: dense arrow row {row} has width {span}, expected full chart width {q}"
));
}
}
let expected = row_offsets[n]
.checked_add(border_dim)
.ok_or_else(|| format!("{owner}: arrow joint length overflows usize"))?;
if dense.len() != expected {
return Err(format!(
"{owner}: dense joint vector has length {}, but the dense arrow layout has length {expected}",
dense.len()
));
}
Ok(dense.to_owned())
}
pub(crate) fn joint_chart_gauge_basis_for_arrow_layout(
&self,
row_offsets: &[usize],
border_dim: usize,
owner: &str,
) -> Result<Vec<Array1<f64>>, String> {
let mut basis = Vec::<Array1<f64>>::new();
if border_dim != self.factored_border_dim() {
return Ok(Vec::new());
}
for dense in self.dense_step_gauge_vectors()? {
let mut gauge = self.dense_joint_vector_in_arrow_layout(
dense.view(),
row_offsets,
border_dim,
owner,
)?;
let original_norm = gauge.dot(&gauge).max(0.0).sqrt();
if !(original_norm.is_finite() && original_norm > 0.0) {
continue;
}
for _ in 0..2 {
for kept in &basis {
let coefficient = gauge.dot(kept);
gauge.scaled_add(-coefficient, kept);
}
}
let residual_norm = gauge.dot(&gauge).max(0.0).sqrt();
if !(residual_norm.is_finite()
&& residual_norm > f64::EPSILON.sqrt() * original_norm)
{
continue;
}
gauge.mapv_inplace(|value| value / residual_norm);
basis.push(gauge);
}
Ok(basis)
}
pub(crate) fn closed_form_beta_gauge_directions(&self) -> Result<Vec<Array1<f64>>, String> {
let border = self.factored_border_dim();
if border == 0 {
return Ok(Vec::new());
}
let coord_len = self.n_obs() * self.assignment.row_block_dim();
let mut out = Vec::new();
let mut probe: Vec<Array1<f64>> = Vec::new();
for gauge in self.dense_step_gauge_vectors()? {
if gauge.len() != coord_len + border {
continue;
}
let beta_part = gauge.slice(s![coord_len..]).to_owned();
let norm_sq = beta_part.iter().map(|&v| v * v).sum::<f64>();
if !(norm_sq.is_finite() && norm_sq > 1.0e-24) {
continue;
}
let mut residual = beta_part.clone();
for basis in &probe {
let coefficient = residual.dot(basis);
residual.scaled_add(-coefficient, basis);
}
let residual_norm_sq = residual.dot(&residual);
if !(residual_norm_sq.is_finite() && residual_norm_sq > 0.0) {
continue;
}
residual *= residual_norm_sq.sqrt().recip();
probe.push(residual);
out.push(beta_part);
}
Ok(out)
}
pub(crate) fn row_gauge_deflation_for_layout(
&self,
row_layout: Option<&SaeRowLayout>,
) -> Result<Option<ArrowRowGaugeDeflation>, String> {
let n = self.n_obs();
let mut rows: Vec<Vec<Array1<f64>>> = Vec::with_capacity(n);
for row in 0..n {
let q_row = match row_layout {
Some(layout) => layout.row_q_active(row),
None => self.assignment.row_block_dim(),
};
rows.push(Vec::with_capacity(self.k_atoms().min(4)));
match row_layout {
Some(layout) => {
for (active_pos, &atom_idx) in layout.active_atoms[row].iter().enumerate() {
let start = layout.coord_starts[row][active_pos];
self.push_atom_row_gauge_deflations(
&mut rows[row],
row,
atom_idx,
start,
q_row,
)?;
}
}
None => {
let coord_offsets = self.assignment.coord_offsets();
for atom_idx in 0..self.k_atoms() {
self.push_atom_row_gauge_deflations(
&mut rows[row],
row,
atom_idx,
coord_offsets[atom_idx],
q_row,
)?;
}
}
}
}
if rows.iter().all(Vec::is_empty) {
Ok(None)
} else {
Ok(Some(ArrowRowGaugeDeflation::new(rows)))
}
}
pub(crate) fn push_atom_row_gauge_deflations(
&self,
row_dirs: &mut Vec<Array1<f64>>,
row: usize,
atom_idx: usize,
coord_start: usize,
q_row: usize,
) -> Result<(), String> {
let d = self.assignment.coords[atom_idx].latent_dim();
let mut tangent = vec![0.0_f64; self.output_dim()];
match self.atoms[atom_idx].basis_kind() {
SaeAtomBasisKind::Linear
| SaeAtomBasisKind::EuclideanPatch
| SaeAtomBasisKind::Duchon
| SaeAtomBasisKind::Poincare => {
for axis in 0..d {
self.atoms[atom_idx].fill_decoded_derivative_row(row, axis, &mut tangent);
if tangent.iter().map(|&v| v * v).sum::<f64>() <= 1.0e-24 {
continue;
}
let mut translation = Array1::<f64>::zeros(q_row);
translation[coord_start + axis] = 1.0;
row_dirs.push(translation);
let coord_value = self.assignment.coords[atom_idx].as_matrix()[[row, axis]];
let mut scale = Array1::<f64>::zeros(q_row);
scale[coord_start + axis] = coord_value;
row_dirs.push(scale);
}
}
SaeAtomBasisKind::Periodic | SaeAtomBasisKind::Torus => {
for axis in 0..d {
self.atoms[atom_idx].fill_decoded_derivative_row(row, axis, &mut tangent);
if tangent.iter().map(|&v| v * v).sum::<f64>() <= 1.0e-24 {
continue;
}
let mut phase = Array1::<f64>::zeros(q_row);
phase[coord_start + axis] = 1.0;
row_dirs.push(phase);
}
}
SaeAtomBasisKind::KleinBottle => {
if d != 2 {
return Err(format!(
"push_atom_row_gauge_deflations: Klein atom {atom_idx} requires latent dimension 2, got {d}"
));
}
self.atoms[atom_idx].fill_decoded_derivative_row(row, 0, &mut tangent);
if tangent.iter().map(|&v| v * v).sum::<f64>() > 1.0e-24 {
let mut phase = Array1::<f64>::zeros(q_row);
phase[coord_start] = 1.0;
row_dirs.push(phase);
}
}
SaeAtomBasisKind::Sphere | SaeAtomBasisKind::ProjectivePlane => {
if d != 3 {
return Err(format!(
"push_atom_row_gauge_deflations: spherical atom {atom_idx} rides the ambient cover and requires latent dimension 3, got {d}"
));
}
let coords = self.assignment.coords[atom_idx].as_matrix();
let directions = ambient_sphere_killing_directions([
coords[[row, 0]],
coords[[row, 1]],
coords[[row, 2]],
]);
for direction in directions {
let mut decoded_motion = vec![0.0_f64; self.output_dim()];
for axis in 0..3 {
self.atoms[atom_idx].fill_decoded_derivative_row(row, axis, &mut tangent);
for output in 0..decoded_motion.len() {
decoded_motion[output] += direction[axis] * tangent[output];
}
}
if decoded_motion.iter().map(|&v| v * v).sum::<f64>() <= 1.0e-24 {
continue;
}
let mut rotation = Array1::<f64>::zeros(q_row);
for axis in 0..2 {
rotation[coord_start + axis] = direction[axis];
}
row_dirs.push(rotation);
}
}
SaeAtomBasisKind::Cylinder => {
if d > 0 {
self.atoms[atom_idx].fill_decoded_derivative_row(row, 0, &mut tangent);
if tangent.iter().map(|&v| v * v).sum::<f64>() > 1.0e-24 {
let mut phase = Array1::<f64>::zeros(q_row);
phase[coord_start] = 1.0;
row_dirs.push(phase);
}
}
}
SaeAtomBasisKind::Mobius
| SaeAtomBasisKind::FiniteSet
| SaeAtomBasisKind::Precomputed(_) => {}
}
Ok(())
}
pub(crate) fn dense_step_gauge_vector_from_field(
&self,
atom_idx: usize,
field: ArrayView2<'_, f64>,
coord_offsets: &[usize],
beta_offsets: &[usize],
total_len: usize,
) -> Result<Option<Array1<f64>>, String> {
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let p = self.output_dim();
let atom = &self.atoms[atom_idx];
let m = atom.basis_size();
let d = self.assignment.coords[atom_idx].latent_dim();
if field.dim() != (n, d) {
return Err(format!(
"dense_step_gauge_vector_from_field: field shape {:?} != ({n}, {d})",
field.dim()
));
}
let mut design = Array2::<f64>::zeros((n, m));
let mut motion = Array2::<f64>::zeros((n, p));
for row in 0..n {
let assignments = self.assignment.try_assignments_row(row)?;
let a = assignments[atom_idx];
if a == 0.0 {
continue;
}
for col in 0..m {
design[[row, col]] = a * atom.basis_values[[row, col]];
}
for axis in 0..d {
let dt = field[[row, axis]];
if dt == 0.0 {
continue;
}
for col in 0..m {
let w = a * dt * atom.basis_jacobian[[row, col, axis]];
if w == 0.0 {
continue;
}
for out_col in 0..p {
motion[[row, out_col]] += w * atom.decoder_coefficients()[[col, out_col]];
}
}
}
}
let raw = motion.iter().map(|v| v * v).sum::<f64>();
if raw <= f64::MIN_POSITIVE || !raw.is_finite() {
return Ok(None);
}
motion.mapv_inplace(|v| -v);
let delta_b = solve_design_least_squares(design.view(), motion.view())?;
let mut gauge = Array1::<f64>::zeros(total_len);
for row in 0..n {
let row_base = row * q + coord_offsets[atom_idx];
for axis in 0..d {
gauge[row_base + axis] = field[[row, axis]];
}
}
let beta_base = n * q + beta_offsets[atom_idx];
let delta_border = match atom.decoder_frame.as_ref() {
Some(frame) => delta_b.dot(&frame.frame()),
None => delta_b,
};
let border_rank = delta_border.ncols();
for col in 0..m {
for channel in 0..border_rank {
gauge[beta_base + col * border_rank + channel] = delta_border[[col, channel]];
}
}
Ok(Some(gauge))
}
pub fn collapse_events(&self) -> &[CollapseEvent] {
&self.collapse_events
}
pub fn record_fit_data_collapse_if_needed(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
iteration: usize,
) -> Result<bool, String> {
let (n, p) = target.dim();
if n == 0 || p == 0 || self.k_atoms() == 0 {
return Ok(false);
}
let atom_count = self.k_atoms();
let verdict = self.dictionary_collapse_verdict(target, rho, None)?;
if let Some(reason) = verdict.proof_unavailable_reason() {
return Err(format!(
"SaeManifoldTerm::record_fit_data_collapse_if_needed: \
decoder-vanishing proof unavailable: {reason}"
));
}
if !verdict.degenerate(atom_count) {
return Ok(false);
}
let assignments = self.assignment.try_assignments()?;
let event_floor = verdict.collapse_event_floor(atom_count).ok_or_else(|| {
"collapse verdict was degenerate without a certified event boundary".to_string()
})?;
let mut collapsed_active_atom = false;
for atom in 0..atom_count {
let active_mass = assignments
.column(atom)
.iter()
.copied()
.fold(0.0_f64, f64::max);
if active_mass == 0.0 {
continue;
}
collapsed_active_atom = true;
let already_terminal = self
.collapse_events
.iter()
.any(|e| e.atom == atom && e.action == CollapseAction::Terminal);
if already_terminal {
continue;
}
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: verdict.explained_variance,
floor: event_floor,
action: CollapseAction::Terminal,
});
}
Ok(collapsed_active_atom)
}
pub fn curvature_walk_report(&self) -> Option<&CurvatureWalkReport> {
self.curvature_walk_report.as_ref()
}
pub(crate) fn reconstruction_residual(
&self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<Array2<f64>, String> {
let fitted = self.try_fitted_for_rho(rho)?;
if fitted.dim() != target.dim() {
return Err(format!(
"SaeManifoldTerm::reconstruction_residual: fitted {:?} != target {:?}",
fitted.dim(),
target.dim()
));
}
Ok(&fitted - &target)
}
pub(crate) fn enforce_active_mass_guard(
&mut self,
iteration: usize,
rho: Option<&SaeManifoldRho>,
) -> Result<(), String> {
if !self.guards_enabled {
return Ok(());
}
let n = self.n_obs();
let k = self.k_atoms();
if n == 0 || k == 0 {
return Ok(());
}
let mut max_mass = vec![0.0_f64; k];
let mut a = vec![0.0_f64; k];
for row in 0..n {
match rho {
Some(_) => self.assignment.try_assignments_row_into(row, &mut a),
None => self
.assignment
.try_assignments_row(row)
.map(|row_a| a.copy_from_slice(row_a.as_slice().expect("contiguous row"))),
}
.map_err(|e| format!("SaeManifoldTerm::enforce_active_mass_guard: {e}"))?;
for atom in 0..k {
if a[atom] > max_mass[atom] {
max_mass[atom] = a[atom];
}
}
}
let active_mass_floor = crate::inference::atom_lens::SAE_TRUST_ACTIVE_MASS_FLOOR;
for atom in 0..k {
if max_mass[atom] >= active_mass_floor {
continue;
}
let reseeds_used = self
.collapse_events
.iter()
.filter(|e| e.atom == atom && e.action == CollapseAction::Reseeded)
.count();
if reseeds_used < SAE_ATOM_COLLAPSE_RESEED_BUDGET {
self.reseed_collapsed_atom_logits(atom);
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: max_mass[atom],
floor: active_mass_floor,
action: CollapseAction::Reseeded,
});
} else {
let already_terminal = self
.collapse_events
.iter()
.any(|e| e.atom == atom && e.action == CollapseAction::Terminal);
if !already_terminal {
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: max_mass[atom],
floor: active_mass_floor,
action: CollapseAction::Terminal,
});
}
}
}
Ok(())
}
pub(crate) fn reseed_collapsed_atom_logits(&mut self, atom: usize) {
let n = self.n_obs();
match self.assignment.mode {
AssignmentMode::Softmax { .. } => {
for row in 0..n {
let row_max = self
.assignment
.logits
.row(row)
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
self.assignment.logits[[row, atom]] =
if row_max.is_finite() { row_max } else { 0.0 };
}
canonicalize_softmax_logits(&mut self.assignment.logits);
}
AssignmentMode::OrderedBetaBernoulli { .. } => {
for row in 0..n {
self.assignment.logits[[row, atom]] = 0.0;
}
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
for row in 0..n {
self.assignment.logits[[row, atom]] = threshold + temperature;
}
}
AssignmentMode::TopK { .. } => {
for row in 0..n {
let row_max = self
.assignment
.logits
.row(row)
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
self.assignment.logits[[row, atom]] =
if row_max.is_finite() { row_max } else { 0.0 };
}
}
}
}
pub(crate) fn fix_decoder_scale_gauge(&mut self) -> Result<(), String> {
let k = self.k_atoms();
if k < 2 || self.frames_active() {
return Ok(());
}
if !matches!(self.assignment.mode, AssignmentMode::Softmax { .. }) {
return Ok(());
}
let norms: Vec<f64> = self
.atoms
.iter()
.map(|atom| atom.contribution_frobenius_scale())
.collect();
let mut sorted = norms.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let median = if k % 2 == 1 {
sorted[k / 2]
} else {
0.5 * (sorted[k / 2 - 1] + sorted[k / 2])
};
if !(median > 0.0) {
return Ok(());
}
let ceiling = median / SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO;
let n = self.n_obs();
let mut any_regauged = false;
for atom in 0..k {
if !(norms[atom] > ceiling) {
continue;
}
let s = norms[atom] / median;
if !(s.is_finite() && s > 1.0) {
continue;
}
self.atoms[atom]
.decoder_coefficients_mut()
.mapv_inplace(|v| v / s);
let ln_s = s.ln();
for row in 0..n {
self.assignment.logits[[row, atom]] += ln_s;
}
any_regauged = true;
}
if any_regauged {
canonicalize_softmax_logits(&mut self.assignment.logits);
}
Ok(())
}
pub(crate) fn enforce_decoder_norm_guard(
&mut self,
target: ArrayView2<'_, f64>,
iteration: usize,
rho: &SaeManifoldRho,
target_col_stats: Option<&TargetCenteredColStats>,
) -> Result<(), String> {
if !self.guards_enabled {
return Ok(());
}
let n = self.n_obs();
let k = self.k_atoms();
if n == 0 || k < 2 {
return Ok(());
}
let norms: Vec<f64> = self
.atoms
.iter()
.map(|atom| atom.contribution_frobenius_scale())
.collect();
let mut sorted = norms.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let median = if k % 2 == 1 {
sorted[k / 2]
} else {
0.5 * (sorted[k / 2 - 1] + sorted[k / 2])
};
if !(median > 0.0) && iteration == 0 {
return Ok(());
}
let floor = SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO * median;
let mut breached: Vec<usize> = Vec::new();
for atom in 0..k {
if norms[atom] < floor {
breached.push(atom);
}
}
if breached.is_empty() {
if iteration == 0 {
return Ok(());
}
let verdict = self.dictionary_collapse_verdict(target, rho, target_col_stats)?;
if let Some(reason) = verdict.proof_unavailable_reason() {
return Err(format!(
"SaeManifoldTerm::enforce_decoder_norm_guard: \
decoder-vanishing proof unavailable: {reason}"
));
}
let ev = verdict.explained_variance;
let dictionary_degenerate = verdict.degenerate(k);
if !dictionary_degenerate
&& let Some((j, kk, coherence)) = self.structural_coherence_collapse_detected()?
{
log::warn!(
"SaeManifoldTerm: structural coherence collapse — atoms ({j}, {kk}) decode a \
shared output subspace (μ̂={coherence:.4} above the derived random-subspace \
null) at healthy EV={ev:.4}; diagnostic only, deferred to the structure search"
);
}
if !dictionary_degenerate {
return Ok(());
}
let collapse_event_floor = verdict.collapse_event_floor(k).ok_or_else(|| {
"degenerate dictionary verdict omitted its certified boundary".to_string()
})?;
let max_signal_upper_bound = verdict
.decoder_vanishing
.max_signal_upper_bound()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted signal bounds".to_string()
})?;
let residual_roundoff_floor = verdict
.decoder_vanishing
.residual_roundoff_floor()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its roundoff floor".to_string()
})?;
let residual_scale_upper = verdict
.decoder_vanishing
.residual_scale_upper()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its residual scale".to_string()
})?;
let signal_vanish_boundary = verdict
.decoder_vanishing
.signal_vanish_boundary()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its signal boundary".to_string()
})?;
let collapse_arm = if verdict.all_decoders_vanished(k) {
"all gated decoder signals vanished at floating-point resolution"
} else {
"the decoder union output span collapsed structurally"
};
let candidate_uniformity = self.coordinate_uniformity_aggregate();
let prefer = match self.best_cocollapse_incumbent.as_ref() {
None => ev.is_finite(),
Some((best_ev, best_uniformity, _)) => prefer_candidate_basin(
ev,
candidate_uniformity,
*best_ev,
*best_uniformity,
SAE_FINAL_EV_DEGRADATION_TOL,
),
};
if prefer {
self.best_cocollapse_incumbent =
Some((ev, candidate_uniformity, self.snapshot_mutable_state()));
}
if self.dictionary_cocollapse_reseeds >= SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET {
if let Some((best_ev, best_uniformity, best_state)) =
self.best_cocollapse_incumbent.take()
{
let current_uniformity = self.coordinate_uniformity_aggregate();
if prefer_candidate_basin(
best_ev,
best_uniformity,
ev,
current_uniformity,
SAE_FINAL_EV_DEGRADATION_TOL,
) {
self.restore_mutable_state(&best_state)?;
log::warn!(
"SaeManifoldTerm: dictionary co-collapse multi-start budget spent; \
restoring best basin (EV={best_ev:.4}) over last reseed (EV={ev:.4})"
);
}
}
for atom in 0..k {
let already_terminal = self
.collapse_events
.iter()
.any(|e| e.atom == atom && e.action == CollapseAction::Terminal);
if !already_terminal {
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: ev,
floor: collapse_event_floor,
action: CollapseAction::Terminal,
});
}
}
return Ok(());
}
self.dictionary_cocollapse_reseeds += 1;
log::warn!(
"SaeManifoldTerm: dictionary co-collapse ({collapse_arm}; EV telemetry={ev:.4}, \
max gated-signal upper bound={max_signal_upper_bound:.3e}, residual scale \
upper bound={residual_scale_upper:.3e}, residual roundoff floor=\
{residual_roundoff_floor:.3e}, derived signal boundary=\
{signal_vanish_boundary:.3e}) with no relative-norm breach; \
reseeding all {k} atoms onto distinct residual PCs (dictionary multi-start \
{}/{SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET}: total co-collapse, no atom \
carries material signal to anchor)",
self.dictionary_cocollapse_reseeds
);
let all: Vec<usize> = (0..k).collect();
let pc_pair_offset = self.dictionary_cocollapse_reseeds.saturating_sub(1);
self.reseed_atoms_onto_distinct_residual_pcs(&all, target, rho, pc_pair_offset)?;
for atom in 0..k {
self.reseed_collapsed_atom_logits(atom);
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: ev,
floor: collapse_event_floor,
action: CollapseAction::Reseeded,
});
}
self.refit_decoder_sequential_deflation(target)?;
self.anchor_logits_to_residual_ownership(target)?;
self.refit_decoder_sequential_deflation(target)?;
let revert_to_incumbent = if let Some((incumbent_ev, incumbent_uniformity, _)) =
self.best_cocollapse_incumbent.as_ref()
{
let incumbent_ev = *incumbent_ev;
let incumbent_uniformity = *incumbent_uniformity;
let reseeded_ev =
self.dictionary_reconstruction_ev_maybe(target, rho, target_col_stats)?;
let reseeded_uniformity = self.coordinate_uniformity_aggregate();
!prefer_candidate_basin(
reseeded_ev,
reseeded_uniformity,
incumbent_ev,
incumbent_uniformity,
SAE_FINAL_EV_DEGRADATION_TOL,
)
} else {
false
};
if revert_to_incumbent {
let incumbent = self.best_cocollapse_incumbent.take();
if let Some((_, _, ref state)) = incumbent {
self.restore_mutable_state(state)?;
}
self.best_cocollapse_incumbent = incumbent;
}
return Ok(());
}
let mut to_reseed: Vec<usize> = Vec::new();
for &atom in &breached {
let reseeds_used = self
.collapse_events
.iter()
.filter(|e| e.atom == atom && e.action == CollapseAction::Reseeded)
.count();
if reseeds_used < SAE_ATOM_COLLAPSE_RESEED_BUDGET {
to_reseed.push(atom);
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: norms[atom] / median,
floor: SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO,
action: CollapseAction::Reseeded,
});
} else {
let already_terminal = self
.collapse_events
.iter()
.any(|e| e.atom == atom && e.action == CollapseAction::Terminal);
if !already_terminal {
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: norms[atom] / median,
floor: SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO,
action: CollapseAction::Terminal,
});
}
}
}
if !to_reseed.is_empty() {
self.reseed_atoms_onto_distinct_residual_pcs(&to_reseed, target, rho, 0)?;
for &atom in &to_reseed {
self.reseed_collapsed_atom_logits(atom);
}
self.refit_decoder_least_squares_at_current_state(target, Some(rho))?;
}
Ok(())
}
pub(crate) fn dictionary_reconstruction_ev(
&self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<f64, String> {
self.dictionary_reconstruction_ev_maybe(target, rho, None)
}
pub(crate) fn dictionary_reconstruction_ev_maybe(
&self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
precomputed: Option<&TargetCenteredColStats>,
) -> Result<f64, String> {
let residual = self.reconstruction_residual(target, rho)?;
let residual_energy = self.residual_energy_for_vanishing(residual.view())?;
self.dictionary_reconstruction_ev_from_residual_sum_squares(
residual_energy.sum_squares(),
target,
precomputed,
)
}
fn dictionary_reconstruction_ev_from_residual_sum_squares(
&self,
ss_res: f64,
target: ArrayView2<'_, f64>,
precomputed: Option<&TargetCenteredColStats>,
) -> Result<f64, String> {
if !(ss_res.is_finite() && ss_res >= 0.0) {
return Err(format!(
"dictionary reconstruction residual sum of squares must be finite and \
non-negative; got {ss_res}"
));
}
let owned;
let ss_tot = match precomputed {
Some(stats) => stats.ss_tot,
None => {
owned = TargetCenteredColStats::compute(target);
owned.ss_tot
}
};
if !(ss_tot > 0.0) {
return Ok(if ss_res > 0.0 { 0.0 } else { 1.0 });
}
Ok(1.0 - ss_res / ss_tot)
}
pub(crate) fn reseed_atoms_onto_distinct_residual_pcs(
&mut self,
atoms: &[usize],
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
pc_pair_offset: usize,
) -> Result<(), String> {
if atoms.is_empty() {
return Ok(());
}
let residual = self.reconstruction_residual(target, rho)?;
let basis_kinds: Vec<SaeAtomBasisKind> = atoms
.iter()
.map(|&a| self.atoms[a].basis_kind().clone())
.collect();
let dims: Vec<usize> = atoms.iter().map(|&a| self.atoms[a].latent_dim()).collect();
let n = self.n_obs();
let pc_pairs = (residual.ncols().min(n)) / 2;
let data_row_reseed = self.data_row_reseed;
let all_flat = basis_kinds.iter().all(|k| {
matches!(
k,
SaeAtomBasisKind::EuclideanPatch | SaeAtomBasisKind::Linear
)
});
let seeded = if data_row_reseed && all_flat && n > 0 && pc_pair_offset >= pc_pairs.max(1) {
let anchor_rows: Vec<usize> = (0..atoms.len())
.map(|slot| (slot + pc_pair_offset.wrapping_mul(atoms.len().max(1))) % n)
.collect();
sae_data_row_anchored_euclidean_coords(residual.view(), &dims, &anchor_rows)?
} else {
sae_pca_seed_initial_coords_with_pc_offset(
residual.view(),
&basis_kinds,
&dims,
pc_pair_offset,
)?
};
for (slot, &atom) in atoms.iter().enumerate() {
let d = dims[slot];
let mut flat = Array1::<f64>::zeros(n * d);
for row in 0..n {
for axis in 0..d {
flat[row * d + axis] = seeded[[slot, row, axis]];
}
}
self.assignment.coords[atom].set_flat(flat.view());
let coords = self.assignment.coords[atom].as_matrix();
self.atoms[atom].refresh_basis(coords.view())?;
}
Ok(())
}
pub(crate) fn residual_has_uncovered_signal(
&self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<bool, String> {
let residual = self.reconstruction_residual(target, rho)?;
self.residual_view_has_uncovered_signal(residual.view())
}
fn residual_view_has_uncovered_signal(
&self,
residual: ArrayView2<'_, f64>,
) -> Result<bool, String> {
if residual.nrows() == 0 || residual.ncols() == 0 {
return Ok(false);
}
let gram = fast_atb(&residual, &residual);
let (_u, energies, _vt) = gram
.svd(false, false)
.map_err(|e| format!("residual_has_uncovered_signal: residual-Gram SVD failed: {e}"))?;
let energies: Vec<f64> = energies
.iter()
.copied()
.filter(|v| v.is_finite() && *v >= 0.0)
.collect();
Ok(leading_direction_above_noise_floor(&energies))
}
pub(crate) fn reseed_curved_atoms_sequential_deflation(
&mut self,
atoms: &[usize],
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<Vec<usize>, String> {
if atoms.is_empty() {
return Ok(Vec::new());
}
let n = self.n_obs();
let p = self.output_dim();
if n == 0 || p == 0 {
return Ok(Vec::new());
}
let mut residual = self.reconstruction_residual(target, rho)?;
let mut reseeded = Vec::new();
for &atom in atoms {
if !self.residual_view_has_uncovered_signal(residual.view())? {
break;
}
self.seed_atom_chart_coords(atom, n, residual.view(), None)?;
let m = self.atoms[atom].basis_size();
let mut design = Array2::<f64>::zeros((n, m));
for row in 0..n {
let assignments = self.assignment.try_assignments_row(row)?;
let gate = assignments[atom];
for col in 0..m {
design[[row, col]] = gate * self.atoms[atom].basis_values[[row, col]];
}
}
let beta = solve_design_least_squares(design.view(), residual.view())?;
if beta.dim() != (m, p) {
return Err(format!(
"SaeManifoldTerm::reseed_curved_atoms_sequential_deflation: atom {atom} \
beta shape {:?} != ({m}, {p})",
beta.dim()
));
}
let fit = design.dot(&beta);
residual = &residual - &fit;
for col in 0..m {
for out in 0..p {
self.atoms[atom].decoder_coefficients_mut()[[col, out]] = beta[[col, out]];
}
}
reseeded.push(atom);
}
Ok(reseeded)
}
pub(crate) fn refit_decoder_sequential_deflation(
&mut self,
target: ArrayView2<'_, f64>,
) -> Result<(), String> {
let n = self.n_obs();
let p = self.output_dim();
if target.dim() != (n, p) {
return Err(format!(
"SaeManifoldTerm::refit_decoder_sequential_deflation: target shape {:?} != ({n}, {p})",
target.dim()
));
}
let k = self.k_atoms();
if k == 0 || n == 0 {
return Ok(());
}
let mut gates = Array2::<f64>::zeros((n, k));
for row in 0..n {
let assignments = self.assignment.try_assignments_row(row)?;
for atom in 0..k {
gates[[row, atom]] = assignments[atom];
}
}
let gated_design = |slf: &Self, atom: usize| -> Array2<f64> {
let m = slf.atoms[atom].basis_size();
let mut d = Array2::<f64>::zeros((n, m));
for row in 0..n {
let w = gates[[row, atom]];
for col in 0..m {
d[[row, col]] = w * slf.atoms[atom].basis_values[[row, col]];
}
}
d
};
let mut residual = target.to_owned();
let mut remaining: Vec<usize> = (0..k).collect();
while !remaining.is_empty() {
let mut best_atom = remaining[0];
let mut best_energy = f64::NEG_INFINITY;
let mut best_beta: Option<Array2<f64>> = None;
for &atom in &remaining {
let d = gated_design(self, atom);
if d.iter().all(|&v| v == 0.0) {
return Err(format!(
"refit_decoder_sequential_deflation: atom {atom} is {}; the seed ρ \
leaves the reduced problem rank-deficient (recoverable \
infeasible-ρ probe)",
ProbeRefusalKind::all_zero_gated_design_marker()
));
}
let beta = solve_design_least_squares(d.view(), residual.view())?;
let fit = d.dot(&beta);
let energy: f64 = fit.iter().map(|v| v * v).sum();
if energy > best_energy {
best_energy = energy;
best_atom = atom;
best_beta = Some(beta);
}
}
let beta = best_beta.expect("remaining is non-empty so a best atom was chosen");
let m = self.atoms[best_atom].basis_size();
if beta.dim() != (m, p) {
return Err(format!(
"SaeManifoldTerm::refit_decoder_sequential_deflation: atom {best_atom} beta shape {:?} != ({m}, {p})",
beta.dim()
));
}
let d = gated_design(self, best_atom);
let fit = d.dot(&beta);
residual = &residual - &fit;
for col in 0..m {
for out in 0..p {
self.atoms[best_atom].decoder_coefficients_mut()[[col, out]] = beta[[col, out]];
}
}
remaining.retain(|&a| a != best_atom);
}
Ok(())
}
pub(crate) fn anchor_logits_to_residual_ownership(
&mut self,
target: ArrayView2<'_, f64>,
) -> Result<(), String> {
let n = self.n_obs();
let p = self.output_dim();
let k = self.k_atoms();
if n == 0 || k < 2 {
return Ok(());
}
let mut owner = vec![0usize; n];
for row in 0..n {
let mut best = 0usize;
let mut best_align = f64::NEG_INFINITY;
for atom in 0..k {
let m = self.atoms[atom].basis_size();
let phi = &self.atoms[atom].basis_values;
let b = self.atoms[atom].decoder_coefficients();
let mut align = 0.0_f64;
for out in 0..p {
let mut recon = 0.0_f64;
for col in 0..m {
recon += phi[[row, col]] * b[[col, out]];
}
align += recon * target[[row, out]];
}
if align > best_align {
best_align = align;
best = atom;
}
}
owner[row] = best;
}
let bias = self.assignment.mode.temperature().max(f64::MIN_POSITIVE);
for row in 0..n {
for atom in 0..k {
let delta = if atom == owner[row] { bias } else { -bias };
self.assignment.logits[[row, atom]] += delta;
}
}
if matches!(self.assignment.mode, AssignmentMode::Softmax { .. }) {
canonicalize_softmax_logits(&mut self.assignment.logits);
}
Ok(())
}
pub(crate) fn structural_coherence_collapse_detected(
&self,
) -> Result<Option<(usize, usize, f64)>, String> {
Ok(self
.structural_coherence_collapsed_pairs()?
.into_iter()
.max_by(|a, b| a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal))
.map(|(j, kk, coherence, _bar)| (j, kk, coherence)))
}
fn structural_coherence_collapsed_pairs(
&self,
) -> Result<Vec<(usize, usize, f64, f64)>, String> {
let k = self.k_atoms();
let p = self.output_dim();
if k < 2 || p == 0 {
return Ok(Vec::new());
}
let frames = (0..k)
.map(|atom| crate::manifold::certificate::certificate_output_frame(self, atom))
.collect::<Result<Vec<_>, String>>()?;
let effective_output_rank = union_output_frame_rank(&frames, p);
let overcomplete = k > effective_output_rank;
let mut candidates: Vec<(usize, usize)> = Vec::new();
for j in 0..k {
for kk in (j + 1)..k {
let rj = frames[j].ncols();
let rk = frames[kk].ncols();
if rj == 0 || rk == 0 {
continue;
}
if overcomplete {
candidates.push((j, kk));
continue;
}
let overlap = fast_atb(&frames[j], &frames[kk]);
let (_u, s, _vt) = overlap.svd(false, false).map_err(|e| {
format!("structural_coherence_collapse_detected: SVD failed ({j},{kk}): {e}")
})?;
let coherence = s.iter().copied().fold(0.0_f64, f64::max);
let a = rj as f64 / p as f64;
let b = rk as f64 / p as f64;
let mu_null = (a * (1.0 - b)).max(0.0).sqrt() + (b * (1.0 - a)).max(0.0).sqrt();
let frame_bar = 0.5 * (mu_null.min(1.0) + 1.0);
if coherence > frame_bar {
candidates.push((j, kk));
}
}
}
if candidates.is_empty() {
return Ok(Vec::new());
}
let gates = self.assignment.assignments();
let n = gates.nrows();
let mut in_candidate = vec![false; k];
for &(j, kk) in &candidates {
in_candidate[j] = true;
in_candidate[kk] = true;
}
let mut contribution: Vec<Option<Array2<f64>>> = vec![None; k];
for atom in 0..k {
if !in_candidate[atom] {
continue;
}
let phi = &self.atoms[atom].basis_values;
if phi.nrows() != n || n == 0 {
continue;
}
let mut y = fast_ab(phi, self.atoms[atom].decoder_coefficients());
for row in 0..n {
let g = gates[[row, atom]];
for col in 0..y.ncols() {
y[[row, col]] *= g;
}
}
contribution[atom] = Some(y);
}
let mut collapsed = Vec::with_capacity(candidates.len());
for (j, kk) in candidates {
let d_eff = (self.atoms[j].basis_size().max(1) * frames[j].ncols().max(1))
.min(self.atoms[kk].basis_size().max(1) * frames[kk].ncols().max(1))
as f64;
let e_null = (2.0 / (std::f64::consts::PI * d_eff)).sqrt();
let contribution_bar = 0.5 * (e_null.min(1.0) + 1.0);
let contribution_cos = match (&contribution[j], &contribution[kk]) {
(Some(yj), Some(yk)) => {
let mut dot = 0.0_f64;
let mut nj = 0.0_f64;
let mut nk = 0.0_f64;
for (a, b) in yj.iter().zip(yk.iter()) {
dot += a * b;
nj += a * a;
nk += b * b;
}
let denom = (nj * nk).sqrt();
if denom > 0.0 {
(dot / denom).abs()
} else {
0.0
}
}
_ => {
if overcomplete {
0.0
} else {
1.0
}
}
};
if contribution_cos > contribution_bar {
collapsed.push((j, kk, contribution_cos, contribution_bar));
}
}
Ok(collapsed)
}
pub(crate) fn enforce_structural_coherence_guard(
&mut self,
target: ArrayView2<'_, f64>,
iteration: usize,
rho: &SaeManifoldRho,
) -> Result<(), String> {
if !self.guards_enabled || self.k_atoms() < 2 {
return Ok(());
}
let mut pairs = self.structural_coherence_collapsed_pairs()?;
if pairs.is_empty() {
return Ok(());
}
pairs.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal));
let k = self.k_atoms();
let decoder_norms: Vec<f64> = self
.atoms
.iter()
.map(|atom| {
atom.decoder_coefficients()
.iter()
.map(|value| value * value)
.sum::<f64>()
.sqrt()
})
.collect();
let mut selected = vec![false; k];
let mut floor_by_atom = vec![0.0_f64; k];
let mut coherence_by_atom = vec![0.0_f64; k];
for &(j, kk, coherence, bar) in &pairs {
let atom = if selected[j] && selected[kk] {
continue;
} else if selected[j] {
kk
} else if selected[kk] {
j
} else if decoder_norms[j] < decoder_norms[kk] {
j
} else if decoder_norms[kk] < decoder_norms[j] {
kk
} else {
kk
};
selected[atom] = true;
floor_by_atom[atom] = bar;
coherence_by_atom[atom] = coherence;
}
let residual_uncovered = self.residual_has_uncovered_signal(target, rho)?;
let mut to_reseed = Vec::new();
for atom in 0..k {
if !selected[atom] {
continue;
}
let reseeds_used = self
.collapse_events
.iter()
.filter(|event| event.atom == atom && event.action == CollapseAction::Reseeded)
.count();
if residual_uncovered
&& reseeds_used < SAE_ATOM_COLLAPSE_RESEED_BUDGET
&& self.structural_cocollapse_reseeds < SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET
{
to_reseed.push(atom);
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: coherence_by_atom[atom],
floor: floor_by_atom[atom],
action: CollapseAction::Reseeded,
});
} else {
let already_terminal = self
.collapse_events
.iter()
.any(|event| event.atom == atom && event.action == CollapseAction::Terminal);
if !already_terminal {
self.collapse_events.push(CollapseEvent {
iteration,
atom,
max_active_mass: coherence_by_atom[atom],
floor: floor_by_atom[atom],
action: CollapseAction::Terminal,
});
}
}
}
if to_reseed.is_empty() {
return Ok(());
}
self.structural_cocollapse_reseeds += 1;
log::warn!(
"SaeManifoldTerm: structural coherence collapse — reseeding {} duplicate-output \
atom(s) onto residual PCs (structural multi-start \
{}/{SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET})",
to_reseed.len(),
self.structural_cocollapse_reseeds
);
let pc_pair_offset = self.structural_cocollapse_reseeds.saturating_sub(1);
let (curved, flat): (Vec<usize>, Vec<usize>) =
to_reseed.iter().copied().partition(|&atom| {
!matches!(
self.atoms[atom].basis_kind(),
SaeAtomBasisKind::EuclideanPatch | SaeAtomBasisKind::Linear
)
});
if !curved.is_empty() {
self.reseed_curved_atoms_sequential_deflation(&curved, target, rho)?;
}
if !flat.is_empty() {
self.reseed_atoms_onto_distinct_residual_pcs(&flat, target, rho, pc_pair_offset)?;
}
for &atom in &to_reseed {
self.reseed_collapsed_atom_logits(atom);
}
self.refit_decoder_sequential_deflation(target)?;
self.anchor_logits_to_residual_ownership(target)?;
self.refit_decoder_sequential_deflation(target)?;
Ok(())
}
pub(crate) fn seed_cold_start_disjoint_charts(
&mut self,
target: ArrayView2<'_, f64>,
) -> Result<(), String> {
self.seed_disjoint_charts(target)
}
pub(crate) fn place_reactive_entry_disjoint_charts(
&mut self,
target: ArrayView2<'_, f64>,
) -> Result<(), String> {
self.seed_disjoint_charts(target)
}
fn seed_disjoint_charts(&mut self, target: ArrayView2<'_, f64>) -> Result<(), String> {
let n = self.n_obs();
let p = self.output_dim();
let k = self.k_atoms();
if n == 0 || k == 0 {
return Ok(());
}
let joint_planes: Vec<super::isa_seed::IsaPlaneCandidate> =
match super::isa_seed::capture_signal_span(target, k)? {
Some(parts) => super::isa_seed::isa_extract_certified_planes(
target,
&parts,
k,
&super::isa_seed::IsaSeedConfig::default(),
),
None => Vec::new(),
};
let mut next_isa_plane = joint_planes.into_iter();
let mut residual = target.to_owned();
for atom in 0..k {
let dim = self.atoms[atom].latent_dim();
let is_periodic = matches!(self.atoms[atom].basis_kind(), SaeAtomBasisKind::Periodic);
let isa_plane = if dim > 0 && is_periodic {
next_isa_plane.next()
} else {
None
};
self.seed_atom_chart_coords(atom, n, residual.view(), isa_plane)?;
let m = self.atoms[atom].basis_size();
let mut d = Array2::<f64>::zeros((n, m));
for row in 0..n {
let assignments = self.assignment.try_assignments_row(row)?;
let w = assignments[atom];
for col in 0..m {
d[[row, col]] = w * self.atoms[atom].basis_values[[row, col]];
}
}
let beta = solve_design_least_squares(d.view(), residual.view())?;
if beta.dim() != (m, p) {
return Err(format!(
"SaeManifoldTerm::seed_cold_start_disjoint_charts: atom {atom} beta shape {:?} != ({m}, {p})",
beta.dim()
));
}
let fit = d.dot(&beta);
residual = &residual - &fit;
for col in 0..m {
for out in 0..p {
self.atoms[atom].decoder_coefficients_mut()[[col, out]] = beta[[col, out]];
}
}
}
Ok(())
}
pub(crate) fn refit_reactive_entry_decoders_at_smooth_face(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<(), String> {
self.assignment.validate_rho_domain(rho)?;
let n = self.n_obs();
let p = self.output_dim();
let k = self.k_atoms();
if target.dim() != (n, p) {
return Err(format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: target shape {:?} != ({n}, {p})",
target.dim()
));
}
if rho.log_lambda_smooth.len() != k {
return Err(format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: rho smoothness length {} != K {k}",
rho.log_lambda_smooth.len()
));
}
if n == 0 || k == 0 {
return Ok(());
}
let mut residual = target.to_owned();
for atom in 0..k {
let m = self.atoms[atom].basis_size();
if self.atoms[atom].smooth_penalty().dim() != (m, m) {
return Err(format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: atom {atom} smooth penalty shape {:?} != ({m}, {m})",
self.atoms[atom].smooth_penalty().dim()
));
}
let mut weighted_design = Array2::<f64>::zeros((n, m));
let mut weighted_residual = residual.clone();
for row in 0..n {
let assignments = self.assignment.try_assignments_row(row)?;
let honesty_weight = self
.row_loss_weights
.as_ref()
.map_or(1.0, |weights| weights[row]);
if !(honesty_weight.is_finite() && honesty_weight >= 0.0) {
return Err(format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: row {row} has invalid design-honesty weight {honesty_weight}"
));
}
let root_weight = honesty_weight.sqrt();
let gate = assignments[atom];
for basis_col in 0..m {
weighted_design[[row, basis_col]] =
root_weight * gate * self.atoms[atom].basis_values[[row, basis_col]];
}
for output in 0..p {
weighted_residual[[row, output]] *= root_weight;
}
}
let mut normal = fast_atb(&weighted_design, &weighted_design);
let lambda = rho.lambda_smooth_for(atom)?;
for left in 0..m {
for right in 0..m {
let smooth = 0.5
* (self.atoms[atom].smooth_penalty()[[left, right]]
+ self.atoms[atom].smooth_penalty()[[right, left]]);
normal[[left, right]] += lambda * smooth;
}
}
let rhs = fast_atb(&weighted_design, &weighted_residual);
let factor = normal.cholesky(Side::Lower).map_err(|error| {
format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: atom {atom} penalized normal equation is not positive definite at lambda={lambda:.6e}: {error}"
)
})?;
let beta = factor.solve_mat(&rhs);
if beta.dim() != (m, p) || !beta.iter().all(|value| value.is_finite()) {
return Err(format!(
"SaeManifoldTerm::refit_reactive_entry_decoders_at_smooth_face: atom {atom} solve produced invalid beta shape {:?}",
beta.dim()
));
}
let mut design = weighted_design;
for row in 0..n {
let honesty_weight = self
.row_loss_weights
.as_ref()
.map_or(1.0, |weights| weights[row]);
let root_weight = honesty_weight.sqrt();
if root_weight > 0.0 {
for basis_col in 0..m {
design[[row, basis_col]] /= root_weight;
}
} else {
let assignments = self.assignment.try_assignments_row(row)?;
for basis_col in 0..m {
design[[row, basis_col]] =
assignments[atom] * self.atoms[atom].basis_values[[row, basis_col]];
}
}
}
let fit = design.dot(&beta);
residual = &residual - &fit;
self.atoms[atom].decoder_coefficients_mut().assign(&beta);
}
Ok(())
}
fn seed_atom_chart_coords(
&mut self,
atom: usize,
n: usize,
residual: ArrayView2<'_, f64>,
isa_plane: Option<super::isa_seed::IsaPlaneCandidate>,
) -> Result<(), String> {
let kind = self.atoms[atom].basis_kind().clone();
let dim = self.atoms[atom].latent_dim();
let mut flat = Array1::<f64>::zeros(n * dim);
if let Some(plane) = isa_plane {
for row in 0..n {
flat[row * dim] = plane.phases_turns[[row, 0]];
}
} else {
let seeded = sae_pca_seed_initial_coords(
residual,
std::slice::from_ref(&kind),
std::slice::from_ref(&dim),
)?;
for row in 0..n {
for axis in 0..dim {
flat[row * dim + axis] = seeded[[0, row, axis]];
}
}
}
self.assignment.coords[atom].set_flat(flat.view());
let coords = self.assignment.coords[atom].as_matrix();
self.atoms[atom].refresh_basis(coords.view())?;
Ok(())
}
pub(crate) fn apply_newton_step_impl(
&mut self,
delta_ext_coord: ArrayView1<'_, f64>,
delta_beta: ArrayView1<'_, f64>,
step_size: f64,
refresh_basis: bool,
) -> Result<(), String> {
self.apply_newton_step_impl_with_parallelism(
delta_ext_coord,
delta_beta,
step_size,
refresh_basis,
None,
)
}
pub(crate) fn apply_newton_step_impl_with_parallelism(
&mut self,
delta_ext_coord: ArrayView1<'_, f64>,
delta_beta: ArrayView1<'_, f64>,
step_size: f64,
refresh_basis: bool,
forced_parallelism: Option<bool>,
) -> Result<(), String> {
if !(step_size.is_finite() && step_size > 0.0) {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: step_size must be finite and positive; got {step_size}"
));
}
let n = self.n_obs();
let q = self.assignment.row_block_dim();
let k_atoms = self.k_atoms();
let assignment_dim = self.assignment.assignment_coord_dim();
let at_top_level = rayon::current_thread_index().is_none();
let parallel_rows =
forced_parallelism.unwrap_or(n >= SAE_LOSS_PARALLEL_ROW_MIN && at_top_level);
let parallel_atoms = forced_parallelism
.unwrap_or(n >= SAE_LOSS_PARALLEL_ROW_MIN && k_atoms > 1 && at_top_level);
let softmax = matches!(self.assignment.mode, AssignmentMode::Softmax { .. });
let expected_delta_len = if self.last_frames_active {
self.factored_border_dim()
} else {
self.beta_dim()
};
if delta_beta.len() != expected_delta_len {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: delta_beta length {} != expected {}",
delta_beta.len(),
expected_delta_len
));
}
if let Some(ref layout) = self.last_row_layout.clone() {
let row_dims: Vec<usize> = (0..n).map(|row| layout.row_q_active(row)).collect();
let mut compact_offsets = Vec::with_capacity(n + 1);
compact_offsets.push(0usize);
let mut compact_total = 0usize;
for &row_dim in &row_dims {
compact_total += row_dim;
compact_offsets.push(compact_total);
}
let total_len = compact_offsets[n];
if delta_ext_coord.len() != total_len {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: compact delta_ext_coord length {} != expected {}",
delta_ext_coord.len(),
total_len
));
}
let mut full_delta = vec![0.0_f64; n * q];
if parallel_rows && q > 0 {
use rayon::prelude::*;
full_delta
.par_chunks_mut(q)
.enumerate()
.for_each(|(row, full_row)| {
let compact_row: Vec<f64> = delta_ext_coord
.slice(ndarray::s![compact_offsets[row]..compact_offsets[row + 1]])
.iter()
.copied()
.collect();
layout.expand_row(row, &compact_row, full_row);
});
} else {
for row in 0..n {
let compact_row: Vec<f64> = delta_ext_coord
.slice(ndarray::s![compact_offsets[row]..compact_offsets[row + 1]])
.iter()
.copied()
.collect();
layout.expand_row(row, &compact_row, &mut full_delta[row * q..(row + 1) * q]);
}
}
let logit_step_cap =
SAE_ASSIGNMENT_LOGIT_STEP_CAP_TAUS * self.assignment.mode.temperature();
if parallel_rows {
use rayon::prelude::*;
self.assignment
.logits
.axis_iter_mut(ndarray::Axis(0))
.into_par_iter()
.enumerate()
.for_each(|(row, mut logits)| {
let row_base = row * q;
for atom_idx in 0..assignment_dim {
logits[atom_idx] += (step_size * full_delta[row_base + atom_idx])
.clamp(-logit_step_cap, logit_step_cap);
}
if softmax {
canonicalize_softmax_logit_row(
logits.as_slice_mut().expect("contiguous logit row"),
);
}
});
} else {
for row in 0..n {
let row_base = row * q;
let mut logits = self.assignment.logits.row_mut(row);
for atom_idx in 0..assignment_dim {
logits[atom_idx] += (step_size * full_delta[row_base + atom_idx])
.clamp(-logit_step_cap, logit_step_cap);
}
if softmax {
canonicalize_softmax_logit_row(
logits.as_slice_mut().expect("contiguous logit row"),
);
}
}
}
let coord_offsets = self.assignment.coord_offsets();
self.apply_coordinate_step_from_rows(
n,
q,
&coord_offsets,
step_size,
|_, flat_idx| full_delta[flat_idx],
refresh_basis,
parallel_atoms,
)?;
} else {
if delta_ext_coord.len() != n * q {
return Err(format!(
"SaeManifoldTerm::apply_newton_step: delta_ext_coord length {} != expected {}",
delta_ext_coord.len(),
n * q
));
}
let coord_offsets = self.assignment.coord_offsets();
let logit_step_cap =
SAE_ASSIGNMENT_LOGIT_STEP_CAP_TAUS * self.assignment.mode.temperature();
if parallel_rows {
use rayon::prelude::*;
self.assignment
.logits
.axis_iter_mut(ndarray::Axis(0))
.into_par_iter()
.enumerate()
.for_each(|(row, mut logits)| {
let row_base = row * q;
for atom_idx in 0..assignment_dim {
logits[atom_idx] += (step_size * delta_ext_coord[row_base + atom_idx])
.clamp(-logit_step_cap, logit_step_cap);
}
if softmax {
canonicalize_softmax_logit_row(
logits.as_slice_mut().expect("contiguous logit row"),
);
}
});
} else {
for row in 0..n {
let row_base = row * q;
let mut logits = self.assignment.logits.row_mut(row);
for atom_idx in 0..assignment_dim {
logits[atom_idx] += (step_size * delta_ext_coord[row_base + atom_idx])
.clamp(-logit_step_cap, logit_step_cap);
}
if softmax {
canonicalize_softmax_logit_row(
logits.as_slice_mut().expect("contiguous logit row"),
);
}
}
}
self.apply_coordinate_step_from_rows(
n,
q,
&coord_offsets,
step_size,
|_, flat_idx| delta_ext_coord[flat_idx],
refresh_basis,
parallel_atoms,
)?;
}
if self.last_frames_active {
let delta_b = FrameProjection::new(self).lift_border_vec(delta_beta);
self.apply_decoder_step_from_flat(delta_b.view(), step_size, parallel_atoms)?;
} else {
self.apply_decoder_step_from_flat(delta_beta, step_size, parallel_atoms)?;
}
Ok(())
}
pub(crate) fn solve_fixed_decoder_row_step(
h: ArrayView2<'_, f64>,
g: ArrayView1<'_, f64>,
base_ridge: f64,
) -> Result<Array1<f64>, String> {
let d = h.nrows();
if h.ncols() != d || g.len() != d {
return Err(format!(
"SaeManifoldTerm::solve_fixed_decoder_row_step: shape mismatch H={:?}, g={}",
h.dim(),
g.len()
));
}
if d == 0 {
return Ok(Array1::<f64>::zeros(0));
}
let mut last_err = String::new();
escalate_ridge(
RidgeSchedule {
initial: base_ridge.max(SAE_MANIFOLD_ROW_RIDGE_FLOOR),
growth: SAE_MANIFOLD_ROW_RIDGE_GROWTH,
max_escalations: SAE_MANIFOLD_ROW_RIDGE_MAX_ATTEMPTS,
},
|ridge| {
let mut a = h.to_owned();
for axis in 0..d {
a[[axis, axis]] += ridge;
}
match sae_cholesky_solve_neg_gradient(a.view(), g) {
Ok(delta) => Some(delta),
Err(err) => {
last_err = err;
None
}
}
},
)
.map(|success| success.value)
.map_err(|_| {
format!(
"SaeManifoldTerm::solve_fixed_decoder_row_step: row Hessian did not factor after LM escalation; last error: {last_err}"
)
})
}
pub(crate) fn fixed_decoder_step_from_rows(
sys: &ArrowSchurSystem,
ridge_ext_coord: f64,
) -> Result<Array1<f64>, String> {
let total = sys.row_offsets[sys.rows.len()];
let mut delta = Array1::<f64>::zeros(total);
for (row_idx, row) in sys.rows.iter().enumerate() {
let row_delta =
Self::solve_fixed_decoder_row_step(row.htt.view(), row.gt.view(), ridge_ext_coord)?;
let start = sys.row_offsets[row_idx];
let end = sys.row_offsets[row_idx + 1];
if row_delta.len() != end - start {
return Err(format!(
"SaeManifoldTerm::fixed_decoder_step_from_rows: row {row_idx} delta len {} != row span {}",
row_delta.len(),
end - start
));
}
delta.slice_mut(s![start..end]).assign(&row_delta);
}
Ok(delta)
}
pub(crate) fn enrichment_visit_order(&self) -> Vec<usize> {
let n = self.n_obs();
if self.row_metric.is_none() {
return (0..n).collect();
}
let metric = match self.diagnostic_metric() {
Ok(m) => m,
Err(_) => return (0..n).collect(),
};
let measure = gam_solve::row_sampling_measure::RowSamplingMeasure::from_metric(&metric);
let drawn = measure.enrichment_order(n, n as u64);
let mut order = Vec::with_capacity(n);
let mut seen = vec![false; n];
for row in drawn {
if row < n && !seen[row] {
seen[row] = true;
order.push(row);
}
}
for (row, &was_seen) in seen.iter().enumerate() {
if !was_seen {
order.push(row);
}
}
order
}
pub fn seed_coords_by_decoder_projection(
&mut self,
target: ArrayView2<'_, f64>,
) -> Result<(), String> {
let n = self.n_obs();
let p = self.output_dim();
if target.dim() != (n, p) {
return Err(format!(
"SaeManifoldTerm::seed_coords_by_decoder_projection: target shape {:?} != ({n}, {p})",
target.dim()
));
}
let visit_order = self.enrichment_visit_order();
for atom_idx in 0..self.k_atoms() {
let d = self.atoms[atom_idx].latent_dim();
if matches!(
self.atoms[atom_idx].basis_kind(),
SaeAtomBasisKind::Periodic | SaeAtomBasisKind::Torus
) && d == 1
{
let mut seeded = self.assignment.coords[atom_idx].as_matrix();
let mut decoder = self.atoms[atom_idx].full_width_decoder();
let eta = self.atoms[atom_idx].homotopy_eta;
if eta != 1.0 {
for basis in 3..decoder.nrows() {
for output in 0..decoder.ncols() {
decoder[[basis, output]] *= eta;
}
}
}
let gram = decoder.dot(&decoder.t());
let extrema = PeriodicCurveExtrema::from_gram(gram.view()).map_err(|error| {
format!(
"SaeManifoldTerm::seed_coords_by_decoder_projection: atom {atom_idx}: {error}"
)
})?;
for &row in &visit_order {
let linear = decoder.dot(&target.row(row));
let coefficients = linear.as_slice().ok_or_else(|| {
"SaeManifoldTerm::seed_coords_by_decoder_projection: Fourier coefficients are not contiguous".to_string()
})?;
seeded[[row, 0]] = extrema
.minimize_squared_distance(coefficients)
.map_err(|error| {
format!(
"SaeManifoldTerm::seed_coords_by_decoder_projection: row {row}, atom {atom_idx}: {error}"
)
})?
.coordinate;
}
let flat = Array1::from_iter(seeded.iter().copied());
self.assignment.coords[atom_idx].set_flat(flat.view());
let coords = self.assignment.coords[atom_idx].as_matrix();
self.atoms[atom_idx].refresh_basis(coords.view())?;
}
}
Ok(())
}
pub fn run_fixed_decoder_arrow_schur(
&mut self,
target: ArrayView2<'_, f64>,
rho: &mut SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
max_iter: usize,
step_size: f64,
ridge_ext_coord: f64,
) -> Result<SaeManifoldLoss, String> {
*rho = rho.clone().for_assignment(self.assignment.mode);
self.assignment.validate_rho_domain(rho)?;
if !(step_size.is_finite() && step_size > 0.0) {
return Err(format!(
"SaeManifoldTerm::run_fixed_decoder_arrow_schur: step_size must be finite and positive; got {step_size}"
));
}
let warm_growth = 1.0 / BacktrackConfig::default().contraction;
let unit_step_ceiling = step_size.max(1.0);
let mut warm_step = step_size;
if max_iter < 1 {
return Err(
"SaeManifoldTerm::run_fixed_decoder_arrow_schur: max_iter must be positive".into(),
);
}
let beta_zero = Array1::<f64>::zeros(self.beta_dim());
let mut last_loss = self.loss(target, rho)?;
for _ in 0..max_iter {
self.advance_temperature_schedule()?;
let pre_step_loss = self.loss(target, rho)?;
self.fixed_decoder_assembly = true;
let sys_result = self.assemble_arrow_schur(target, rho, analytic_penalties);
self.fixed_decoder_assembly = false;
let sys = sys_result
.map_err(|err| format!("SaeManifoldTerm::run_fixed_decoder_arrow_schur: {err}"))?;
let pre_step_total =
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)?;
let delta_ext_coord = Self::fixed_decoder_step_from_rows(&sys, ridge_ext_coord)?;
let directional_decrease = sae_manifold_newton_directional_decrease(
&sys,
delta_ext_coord.view(),
beta_zero.view(),
);
let grad_norm_sq: f64 = sys
.rows
.iter()
.flat_map(|row| row.gt.iter())
.map(|&v| v * v)
.sum();
let step_norm_sq: f64 = delta_ext_coord.iter().map(|&v| v * v).sum();
let directional_decrease_floor = SAE_MANIFOLD_DIRECTIONAL_DECREASE_REL_FLOOR
* grad_norm_sq.sqrt()
* step_norm_sq.sqrt();
let snapshot = self.snapshot_mutable_state();
if !(pre_step_total.is_finite()
&& directional_decrease.is_finite()
&& directional_decrease > 0.0
&& directional_decrease > directional_decrease_floor)
{
self.restore_mutable_state(&snapshot)?;
last_loss = pre_step_loss;
break;
}
let mut first_trial = true;
let accepted = backtracking_line_search::<_, String>(
BacktrackConfig {
initial_step: warm_step,
max_steps: SAE_MANIFOLD_MAX_LINESEARCH_HALVINGS + 1,
..BacktrackConfig::default()
},
|trial_step_size| {
if !std::mem::take(&mut first_trial) {
self.restore_mutable_state(&snapshot)?;
}
Ok(self
.apply_newton_step(
delta_ext_coord.view(),
beta_zero.view(),
trial_step_size,
)
.and_then(|()| {
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)
})
.ok()
.map(|post_step_total| (post_step_total, ())))
},
|trial_step_size, post_step_total| {
let armijo_bound = pre_step_total
- SAE_MANIFOLD_ARMIJO_C1 * trial_step_size * directional_decrease;
post_step_total.is_finite() && post_step_total <= armijo_bound
},
)?;
match accepted {
Some(step) => {
warm_step = (if step.step >= warm_step {
warm_step * warm_growth
} else {
step.step * warm_growth
})
.min(unit_step_ceiling);
last_loss = self.loss(target, rho)?;
}
None => {
self.restore_mutable_state(&snapshot)?;
last_loss = pre_step_loss;
break;
}
}
}
Ok(last_loss)
}
pub(crate) fn reduce_atoms_to_data_supported_rank(&mut self) -> Result<(), String> {
let p = self.output_dim();
if p == 0 || self.beta_dim() == 0 {
return Ok(());
}
let mut grams = self.empty_decoder_gram_accumulator();
self.accumulate_decoder_gram(&mut grams)?;
let plans: Vec<Option<Array2<f64>>> =
{
let atoms = &self.atoms;
let compute_plan =
|atom_idx: usize| -> Option<Array2<f64>> {
let m = atoms[atom_idx].basis_size();
if m == 0 || grams[atom_idx].dim() != (m, m) {
return None;
}
let mut data_gram = grams[atom_idx].clone();
for i in 0..m {
for j in 0..i {
let sym = 0.5 * (data_gram[[i, j]] + data_gram[[j, i]]);
data_gram[[i, j]] = sym;
data_gram[[j, i]] = sym;
}
}
let (evals, evecs) = match data_gram.eigh(Side::Lower) {
Ok(pair) => pair,
Err(_) => return None,
};
let max_eig = evals.iter().fold(0.0_f64, |acc, &v| {
if v.is_finite() { acc.max(v) } else { acc }
});
if !(max_eig > 0.0) {
return None;
}
let cutoff = SAE_MANIFOLD_SPECTRAL_RANK_CUTOFF * max_eig;
let kept: Vec<usize> = (0..evals.len())
.filter(|&idx| {
let lambda = evals[idx];
lambda.is_finite() && lambda > cutoff
})
.collect();
let r = kept.len();
if r == m || r == 0 {
return None;
}
if atoms[atom_idx].basis_second_jet.is_none() {
return None;
}
let mut q = Array2::<f64>::zeros((m, r));
for (col, &eig_idx) in kept.iter().enumerate() {
for row in 0..m {
q[[row, col]] = evecs[[row, eig_idx]];
}
}
Some(q)
};
let n_atoms = atoms.len();
let parallel =
n_atoms >= SAE_LOSS_PARALLEL_ROW_MIN && rayon::current_thread_index().is_none();
if parallel {
use rayon::prelude::*;
(0..n_atoms)
.into_par_iter()
.map(|atom_idx| with_nested_parallel(|| compute_plan(atom_idx)))
.collect()
} else {
(0..n_atoms).map(compute_plan).collect()
}
};
for (atom_idx, plan) in plans.into_iter().enumerate() {
let Some(q) = plan else { continue };
self.atoms[atom_idx]
.reduce_basis_to_subspace(&q)
.map_err(|err| {
format!(
"SaeManifoldTerm::reduce_atoms_to_data_supported_rank: atom {atom_idx}: {err}"
)
})?;
}
Ok(())
}
pub fn run_joint_fit_arrow_schur(
&mut self,
target: ArrayView2<'_, f64>,
rho: &mut SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
max_iter: usize,
step_size: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) -> Result<SaeManifoldLoss, String> {
self.run_joint_fit_arrow_schur_with_termination_policy(
target,
rho,
analytic_penalties,
max_iter,
step_size,
ridge_ext_coord,
ridge_beta,
true,
)
.map(|outcome| outcome.loss)
}
pub(crate) fn run_joint_fit_arrow_schur_for_quasi_laplace(
&mut self,
target: ArrayView2<'_, f64>,
rho: &mut SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
max_iter: usize,
step_size: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) -> Result<EvidenceJointFitOutcome, String> {
let entry_state = self.snapshot_mutable_state();
let entry_temperature = self.assignment.mode.temperature();
let outcome = self.run_joint_fit_arrow_schur_with_termination_policy(
target,
rho,
analytic_penalties,
max_iter,
step_size,
ridge_ext_coord,
ridge_beta,
false,
)?;
if matches!(outcome.termination, JointFitTermination::Heuristic) {
return Err(
"SaeManifoldTerm::run_joint_fit_arrow_schur_for_quasi_laplace: heuristic \
termination escaped the evidence policy"
.to_string(),
);
}
let entry_state_recurred = self.matches_mutable_state(&entry_state)
&& self.assignment.mode.temperature().to_bits() == entry_temperature.to_bits();
let temperature_stable_on_reentry =
self.temperature_schedule.as_ref().is_none_or(|schedule| {
schedule.current_tau(schedule.iter_count).to_bits()
== self.assignment.mode.temperature().to_bits()
});
let settled_root = matches!(
outcome.termination,
JointFitTermination::Frozen | JointFitTermination::NoStrictDecrease
);
let gap = if !settled_root {
EvidenceFixedPointGap::NotASettledRoot
} else if outcome.state_moved {
EvidenceFixedPointGap::StateMoved(
outcome
.moved_at
.unwrap_or(StateMoveSite::AcceptedNewtonStep),
)
} else if !entry_state_recurred {
EvidenceFixedPointGap::StateNotRecurred
} else if !temperature_stable_on_reentry {
EvidenceFixedPointGap::TemperatureStillAnnealing
} else {
EvidenceFixedPointGap::None
};
Ok(EvidenceJointFitOutcome {
loss: outcome.loss,
fixed_point: gap == EvidenceFixedPointGap::None,
})
}
pub(crate) fn run_joint_fit_arrow_schur_with_termination_policy(
&mut self,
target: ArrayView2<'_, f64>,
rho: &mut SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
max_iter: usize,
step_size: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
allow_heuristic_termination: bool,
) -> Result<JointFitOutcome, String> {
*rho = rho.clone().for_assignment(self.assignment.mode);
self.assignment.validate_rho_domain(rho)?;
if !(step_size.is_finite() && step_size > 0.0) {
return Err(format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: step_size must be finite and positive; got {step_size}"
));
}
let faer_sequential_inner_fit = gam_linalg::faer_ndarray::FaerSequentialScope::enter();
self.refresh_basis_from_current_coords()
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?;
if max_iter == 0 {
return self.loss(target, rho).map(|loss| JointFitOutcome {
loss,
termination: JointFitTermination::Frozen,
state_moved: false,
moved_at: None,
});
}
self.reduce_atoms_to_data_supported_rank()?;
self.ensure_decoder_frames_active_for_current_decoder()
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?;
if allow_heuristic_termination {
self.collapse_events.clear();
}
self.enforce_active_mass_guard(0, Some(rho))?;
self.enforce_decoder_norm_guard(target, 0, rho, None)?;
{
let mut grams = self.empty_decoder_gram_accumulator();
self.accumulate_decoder_gram(&mut grams)?;
self.finalize_decoder_identifiability_audit(&grams, self.n_obs())?;
}
let max_decoder_norm = self
.atoms
.iter()
.map(|atom| atom.decoder_coefficients().iter().map(|v| v * v).sum::<f64>())
.fold(0.0_f64, f64::max)
.sqrt();
if !(max_decoder_norm > 0.0) {
self.seed_cold_start_disjoint_charts(target)?;
}
let mut best_reconstruction_ev = self
.dictionary_reconstruction_ev(target, rho)
.unwrap_or(f64::NEG_INFINITY);
let mut best_reconstruction_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
let initial_reconstruction_is_structurally_healthy = best_reconstruction_ev.is_finite()
&& best_reconstruction_obj.is_finite()
&& self.structural_coherence_collapse_detected()?.is_none();
let mut best_reconstruction_uniformity = if initial_reconstruction_is_structurally_healthy {
self.coordinate_uniformity_aggregate()
} else {
None
};
let mut best_reconstruction_state = if initial_reconstruction_is_structurally_healthy {
Some(self.snapshot_mutable_state())
} else {
best_reconstruction_ev = f64::NEG_INFINITY;
best_reconstruction_obj = f64::INFINITY;
None
};
let mut warranty_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
let mut warranty_state = if warranty_obj.is_finite() {
Some(self.snapshot_mutable_state())
} else {
None
};
let mut previous_full_iterate_objective = f64::INFINITY;
let mut consecutive_objective_stalls = 0usize;
let target_col_stats = TargetCenteredColStats::compute(target);
let warm_growth = 1.0 / BacktrackConfig::default().contraction;
let mut globalization = InnerGlobalizationHint::resume(
self.inner_globalization_hint,
step_size,
ridge_ext_coord,
ridge_beta,
);
let mut termination = JointFitTermination::IterationGrantExhausted;
let mut gauge_block_armed = true;
let mut state_moved = false;
let mut moved_at: Option<StateMoveSite> = None;
if self.sweep_blocks_to_objective_fixed_point(
target,
rho,
analytic_penalties,
allow_heuristic_termination,
)? {
state_moved = true;
moved_at.get_or_insert(StateMoveSite::EntryBlockSweep);
}
for outer_iteration in 0..max_iter {
let temperature_before = self.assignment.mode.temperature();
if self
.advance_temperature_schedule()?
.is_some_and(|temperature| temperature.to_bits() != temperature_before.to_bits())
{
state_moved = true;
moved_at.get_or_insert(StateMoveSite::TemperatureSchedule);
}
let mut sys = self
.assemble_arrow_schur(target, rho, analytic_penalties)
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?;
let plan = self
.streaming_plan()
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?
.admitted_or_error(self.n_obs(), self.output_dim(), self.k_atoms())
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?;
let mut solve_options = plan
.solve_options_for_border_dim(sys.k)
.with_gpu_policy(self.gpu_policy);
if sys.k > 0
&& matches!(
solve_options.mode,
ArrowSolverMode::Direct | ArrowSolverMode::SqrtBA | ArrowSolverMode::InexactPCG
)
{
match self.closed_form_beta_gauge_directions() {
Ok(dirs) if !dirs.is_empty() => {
let quotient = ArrowBetaGaugeQuotient::new(dirs).map_err(|err| {
format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: invalid closed-form \
beta-gauge quotient: {err}"
)
})?;
sys.set_beta_gauge_quotient(quotient).map_err(|err| {
format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: closed-form \
beta-gauge quotient does not match the assembled border: {err}"
)
})?;
}
Ok(_) => {}
Err(err) => {
return Err(format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: closed-form gauge \
directions: {err}"
));
}
}
}
let existing_resident_frame = self.arrow_assembly_workspace.resident_frame.take();
let resident_frame = prepare_sae_resident_frame(
&sys,
&solve_options,
existing_resident_frame,
)
.map_err(|error| {
format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: resident GPU frame preparation failed: {error}"
)
})?;
solve_options.sae_resident_frame = resident_frame;
self.arrow_assembly_workspace.resident_frame = solve_options.sae_resident_frame.clone();
let (mut delta_ext_coord, mut delta_beta, _diag) =
solve_with_lm_escalation_inner(
&sys,
globalization.lm_ridge_t,
globalization.lm_ridge_b,
&solve_options,
)
.map_err(|err| format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}"))?;
let n_rows = sys.rows.len();
let parallel =
n_rows >= SAE_LOSS_PARALLEL_ROW_MIN && rayon::current_thread_index().is_none();
if parallel {
use rayon::prelude::*;
const CHUNK: usize = 64;
let row_offsets = &sys.row_offsets;
let dt_slice = delta_ext_coord
.as_slice_mut()
.expect("delta_ext_coord contiguous");
let n_chunks = n_rows.div_ceil(CHUNK);
let mut remaining = dt_slice;
let mut segments: Vec<(usize, &mut [f64])> = Vec::with_capacity(n_chunks);
let mut prev_end = 0usize;
for chunk in 0..n_chunks {
let start = chunk * CHUNK;
let end = (start + CHUNK).min(n_rows);
let seg_len = row_offsets[end] - row_offsets[start];
assert!(
prev_end == row_offsets[start],
"sae gauge-fix: non-contiguous row segment at chunk start {start} \
(prev_end={prev_end}, row_offset={})",
row_offsets[start]
);
let (seg, rest) = remaining.split_at_mut(seg_len);
remaining = rest;
segments.push((start, seg));
prev_end = row_offsets[end];
}
segments.into_par_iter().for_each(|(start, seg)| {
let end = (start + CHUNK).min(n_rows);
let mut local = 0usize;
for row_idx in start..end {
let di = sys.row_dims[row_idx];
let dirs = with_nested_parallel(|| {
row_sub_floor_null_directions(sys.rows[row_idx].htt.view())
});
for dir in dirs {
if dir.len() != di {
continue;
}
let mut dot = 0.0;
for a in 0..di {
dot += dir[a] * seg[local + a];
}
for a in 0..di {
seg[local + a] -= dot * dir[a];
}
}
local += di;
}
});
} else {
for row_idx in 0..n_rows {
let off = sys.row_offsets[row_idx];
let di = sys.row_dims[row_idx];
for dir in row_sub_floor_null_directions(sys.rows[row_idx].htt.view()) {
if dir.len() != di {
continue;
}
let mut dot = 0.0;
for a in 0..di {
dot += dir[a] * delta_ext_coord[off + a];
}
for a in 0..di {
delta_ext_coord[off + a] -= dot * dir[a];
}
}
}
}
let mut grad_norm_sq = 0.0;
for (row_idx, row) in sys.rows.iter().enumerate() {
let di = sys.row_dims[row_idx];
for axis in 0..di {
grad_norm_sq += row.gt[axis] * row.gt[axis];
}
}
for idx in 0..sys.k {
grad_norm_sq += sys.gb[idx] * sys.gb[idx];
}
let grad_norm = grad_norm_sq.sqrt();
let iterate_scale = self.inner_iterate_scale();
let grad_tolerance = SAE_MANIFOLD_INNER_GRAD_REL_TOL * iterate_scale;
let step_tolerance = SAE_MANIFOLD_INNER_STEP_REL_TOL * iterate_scale;
let lambda_smooth = rho.lambda_smooth_vec()?;
let quotient_grad_norm =
self.quotient_gradient_norm_from_system(&sys, grad_norm_sq, &lambda_smooth);
if allow_heuristic_termination
&& (grad_norm <= grad_tolerance || quotient_grad_norm <= grad_tolerance)
{
termination = JointFitTermination::Heuristic;
self.reclaim_arrow_assembly_workspace(&mut sys);
break;
}
let mut step_norm_sq = 0.0;
for &v in delta_ext_coord.iter() {
step_norm_sq += v * v;
}
for &v in delta_beta.iter() {
step_norm_sq += v * v;
}
let mut quotient_step_norm = step_norm_sq.sqrt();
if delta_ext_coord.len() == self.n_obs() * self.assignment.row_block_dim()
&& delta_beta.len() == self.factored_border_dim()
{
let quotient_step_norm_sq = self.quotient_newton_step_norm_sq(
delta_ext_coord.view(),
delta_beta.view(),
step_norm_sq,
&lambda_smooth,
)?;
quotient_step_norm = quotient_step_norm_sq.sqrt();
let trust_radius = solve_options.trust_region.radius.min(iterate_scale);
if quotient_step_norm > trust_radius
&& trust_radius.is_finite()
&& trust_radius > 0.0
{
let scale = trust_radius / quotient_step_norm;
delta_ext_coord.mapv_inplace(|v| v * scale);
delta_beta.mapv_inplace(|v| v * scale);
step_norm_sq *= scale * scale;
quotient_step_norm = trust_radius;
}
}
if quotient_step_norm <= step_tolerance {
log::debug!(
"SAE inner quotient step {:.3e} <= tol {:.3e} with non-stationary gradient \
raw={:.3e}, quotient={:.3e}; continuing after quotient trust-region gate",
quotient_step_norm,
step_tolerance,
grad_norm,
quotient_grad_norm
);
}
let directional_decrease = sae_manifold_newton_directional_decrease(
&sys,
delta_ext_coord.view(),
delta_beta.view(),
);
let directional_decrease_floor = SAE_MANIFOLD_DIRECTIONAL_DECREASE_REL_FLOOR
* grad_norm_sq.sqrt()
* step_norm_sq.sqrt();
let snapshot = self.snapshot_mutable_state();
let pre_step_total =
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)?;
if !pre_step_total.is_finite() {
self.restore_mutable_state(&snapshot)?;
self.reclaim_arrow_assembly_workspace(&mut sys);
if !allow_heuristic_termination {
return Err(
"SaeManifoldTerm::run_joint_fit_arrow_schur: evidence polish \
encountered a non-finite pre-step objective"
.to_string(),
);
}
termination = JointFitTermination::Heuristic;
break;
}
if allow_heuristic_termination && previous_full_iterate_objective.is_finite() {
let round_improvement = (previous_full_iterate_objective - pre_step_total).max(0.0);
let objective_scale = previous_full_iterate_objective
.abs()
.max(pre_step_total.abs())
+ 1.0;
let relative_decrease = round_improvement / objective_scale;
if relative_decrease < SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL {
consecutive_objective_stalls += 1;
if consecutive_objective_stalls >= SAE_MANIFOLD_INNER_OBJECTIVE_STALL_MIN_ROUNDS
{
let orbit = if gauge_block_armed {
gauge_block_armed = false;
self.descend_gauge_orbit(
target,
rho,
analytic_penalties,
&rho.lambda_smooth_vec()?,
max_iter.saturating_sub(outer_iteration).max(1),
)?
} else {
GaugeOrbitDescent::default()
};
if orbit.moved() {
state_moved = true;
moved_at.get_or_insert(StateMoveSite::GaugeOrbitDescent);
consecutive_objective_stalls = 0;
previous_full_iterate_objective = f64::NAN;
log::debug!(
"run_joint_fit_arrow_schur: gauge-orbit descent recovered \
{:.6e} over {} round(s) at the objective-stall shortcut, \
iteration {outer_iteration} (span dim {}, \
maxᵢ|gᵀvᵢ|={:.6e}, {} objective evaluations)",
orbit.objective_decrease,
orbit.rounds,
orbit.dimension,
orbit.max_directional_derivative,
orbit.evaluations,
);
self.reclaim_arrow_assembly_workspace(&mut sys);
continue;
}
termination = JointFitTermination::Heuristic;
self.reclaim_arrow_assembly_workspace(&mut sys);
break;
}
} else {
consecutive_objective_stalls = 0;
}
}
previous_full_iterate_objective = pre_step_total;
let descent_direction_ok = directional_decrease.is_finite()
&& directional_decrease > 0.0
&& directional_decrease > directional_decrease_floor;
let mut first_trial = true;
let accepted_step = if descent_direction_ok {
backtracking_line_search::<_, String>(
BacktrackConfig {
initial_step: globalization.warm_step,
max_steps: SAE_MANIFOLD_MAX_LINESEARCH_HALVINGS + 1,
..BacktrackConfig::default()
},
|trial_step_size| {
if !std::mem::take(&mut first_trial) {
self.restore_mutable_state(&snapshot)?;
}
Ok(self
.apply_newton_step(
delta_ext_coord.view(),
delta_beta.view(),
trial_step_size,
)
.and_then(|()| {
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)
})
.ok()
.map(|post_step_total| (post_step_total, ())))
},
|trial_step_size, post_step_total| {
let armijo_bound = pre_step_total
- SAE_MANIFOLD_ARMIJO_C1 * trial_step_size * directional_decrease;
let material_floor = if allow_heuristic_termination {
0.0
} else {
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL
* (1.0 + pre_step_total.abs())
};
post_step_total.is_finite()
&& post_step_total <= armijo_bound
&& pre_step_total - post_step_total >= material_floor
},
)?
} else {
None
};
let accepted = accepted_step.is_some();
log::debug!(
"[SAE/inner] it={outer_iteration} ‖g‖={grad_norm:.6e} \
‖Π⊥g‖={quotient_grad_norm:.6e} ‖Δ‖={:.6e} gᵀΔ={directional_decrease:.6e} \
alpha={} warm={:.4e} ridge_t={:.3e} ridge_b={:.3e} \
obj={pre_step_total:.9e}",
step_norm_sq.sqrt(),
match accepted_step.as_ref() {
Some(step) => format!("{:.4e}", step.step),
None => "rejected".to_string(),
},
globalization.warm_step,
globalization.lm_ridge_t,
globalization.lm_ridge_b,
);
if let Some(step) = accepted_step {
state_moved = true;
moved_at.get_or_insert(StateMoveSite::AcceptedNewtonStep);
gauge_block_armed = true;
let alpha = step.step;
let actual = pre_step_total - step.value;
let d_th_d = (directional_decrease
- globalization.lm_ridge_b * step_norm_sq)
.max(0.0);
let predicted =
(alpha * directional_decrease - 0.5 * alpha * alpha * d_th_d).max(0.0);
let gain_ratio = if predicted > 0.0 {
actual / predicted
} else {
1.0
};
globalization.record_accepted_step(
step.step,
warm_growth,
step_size.max(1.0),
gain_ratio,
ridge_ext_coord,
ridge_beta,
);
self.inner_globalization_hint = Some(globalization);
}
if !accepted {
globalization.reset(step_size, ridge_ext_coord, ridge_beta);
self.inner_globalization_hint = Some(globalization);
self.restore_mutable_state(&snapshot)?;
let correction = ArrowProximalCorrectionOptions {
initial_ridge: ridge_ext_coord
.max(ridge_beta)
.max(SAE_MANIFOLD_ROW_RIDGE_FLOOR),
armijo_c1: SAE_MANIFOLD_ARMIJO_C1,
..ArrowProximalCorrectionOptions::default()
};
let accepted_step = match solve_arrow_newton_step_with_proximal_correction(
&sys,
ridge_ext_coord,
ridge_beta,
pre_step_total,
&solve_options,
&correction,
|trial_delta_t, trial_delta_beta| {
if self.restore_mutable_state(&snapshot).is_err() {
return f64::INFINITY;
}
self.apply_newton_step(trial_delta_t, trial_delta_beta, 1.0)
.and_then(|()| {
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)
})
.unwrap_or(f64::INFINITY)
},
) {
Ok(step) => step,
Err(err) => {
log::debug!(
"run_joint_fit_arrow_schur: proximal correction errored at \
iteration {outer_iteration} (gᵀΔ={directional_decrease:.3e}, \
floor={directional_decrease_floor:.3e}, \
‖g‖={:.3e}): {err}",
grad_norm_sq.sqrt()
);
self.restore_mutable_state(&snapshot)?;
self.reclaim_arrow_assembly_workspace(&mut sys);
if !allow_heuristic_termination {
return Err(format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: evidence \
proximal correction failed before a no-descent certificate: {err}"
));
}
termination = JointFitTermination::Heuristic;
break;
}
};
let proximal_material_floor = if allow_heuristic_termination {
0.0
} else {
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL * (1.0 + pre_step_total.abs())
};
if !(accepted_step.trial_objective_value.is_finite()
&& pre_step_total - accepted_step.trial_objective_value
> proximal_material_floor)
{
log::debug!(
"run_joint_fit_arrow_schur: proximal correction made no decrease at \
iteration {outer_iteration} (trial={:.9e}, pre={pre_step_total:.9e}, \
‖g‖={:.3e})",
accepted_step.trial_objective_value,
grad_norm_sq.sqrt()
);
self.restore_mutable_state(&snapshot)?;
self.reclaim_arrow_assembly_workspace(&mut sys);
let orbit = self.descend_gauge_orbit(
target,
rho,
analytic_penalties,
&rho.lambda_smooth_vec()?,
max_iter.saturating_sub(outer_iteration).max(1),
)?;
if orbit.moved() {
state_moved = true;
moved_at.get_or_insert(StateMoveSite::GaugeOrbitDescent);
log::debug!(
"run_joint_fit_arrow_schur: gauge-orbit descent recovered \
{:.6e} over {} round(s) at iteration {outer_iteration} \
(span dim {}, maxᵢ|gᵀvᵢ|={:.6e}, {} objective evaluations) \
where both Newton movers found none",
orbit.objective_decrease,
orbit.rounds,
orbit.dimension,
orbit.max_directional_derivative,
orbit.evaluations,
);
continue;
}
termination = JointFitTermination::NoStrictDecrease;
break;
}
state_moved = true;
moved_at.get_or_insert(StateMoveSite::ProximalCorrectionStep);
gauge_block_armed = true;
}
self.run_objective_guarded_hook(target, rho, analytic_penalties, 0.0, |term| {
term.canonicalize_affine_gauge_after_accept(Some(rho))
})?;
let kkt_quiescent = grad_norm_sq.is_finite()
&& grad_norm_sq.sqrt()
<= 10.0 * SAE_MANIFOLD_INNER_GRAD_REL_TOL * self.inner_iterate_scale();
if !kkt_quiescent {
self.enforce_active_mass_guard(outer_iteration, Some(rho))?;
}
self.enforce_decoder_norm_guard(target, outer_iteration, rho, Some(&target_col_stats))?;
if !kkt_quiescent {
self.enforce_structural_coherence_guard(target, outer_iteration, rho)?;
}
if self.dictionary_cocollapse_reseeds
>= crate::assignment::SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET
{
let verdict = self.dictionary_collapse_verdict(target, rho, None)?;
if let Some(reason) = verdict.proof_unavailable_reason() {
return Err(format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: \
decoder-vanishing proof unavailable after co-collapse reseeds: {reason}"
));
}
let ev_now = verdict.explained_variance;
let atom_count = self.k_atoms();
if verdict.degenerate(atom_count) {
let decoder_norms: Vec<f64> = self
.atoms
.iter()
.map(|atom| {
atom.decoder_coefficients()
.iter()
.map(|value| value * value)
.sum::<f64>()
.sqrt()
})
.collect();
let arm = if verdict.all_decoders_vanished(atom_count) {
let max_signal = verdict
.decoder_vanishing
.max_signal_upper_bound()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted signal bounds"
.to_string()
})?;
let roundoff_floor = verdict
.decoder_vanishing
.residual_roundoff_floor()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its roundoff floor"
.to_string()
})?;
let residual_scale = verdict
.decoder_vanishing
.residual_scale_upper()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its residual scale"
.to_string()
})?;
let signal_boundary = verdict
.decoder_vanishing
.signal_vanish_boundary()
.ok_or_else(|| {
"certified decoder-vanishing verdict omitted its signal boundary"
.to_string()
})?;
format!(
"every atom's actual gated decoder signal vanished \
[max_signal_upper_bound={max_signal:.3e}, \
residual_scale_upper={residual_scale:.3e}, \
residual_roundoff_floor={roundoff_floor:.3e}, \
derived_signal_boundary={signal_boundary:.3e}]"
)
} else {
let evidence = verdict.structural_collapse.ok_or_else(|| {
"degenerate dictionary verdict omitted structural evidence"
.to_string()
})?;
format!(
"the decoder union output span collapsed structurally \
[decoder_span_rank={}, target_reach={:.4}, \
rank-matched random-subspace null={:.4}]",
evidence.decoder_span_rank,
evidence.target_reach,
evidence.random_subspace_null
)
};
return Err(format!(
"SaeManifoldTerm::run_joint_fit_arrow_schur: dictionary {} after {} \
reseed multi-starts: {arm}, so no residual \
structure could anchor K={} distinct charts for this input \
[EV telemetry={ev_now:.4}, decoder_norms={decoder_norms:.4?}]. \
Refusing to continue the degenerate \
fit. Try fewer atoms (a smaller K), a different atom_topology/assignment, \
more observations, or a different random_state.",
ProbeRefusalKind::total_co_collapse_marker(),
crate::assignment::SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET,
atom_count,
));
}
}
if self.run_objective_guarded_hook(target, rho, analytic_penalties, 0.0, |term| {
term.retract_unit_speed_charts_in_loop()?;
term.fix_decoder_scale_gauge()?;
if term.frames_active() {
term.refresh_active_frames_from_data(target)
.map_err(|err| {
format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}")
})?;
}
Ok(())
})? {
state_moved = true;
moved_at.get_or_insert(StateMoveSite::FrameRefresh);
}
let boundary_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
if boundary_obj.is_finite() && boundary_obj < warranty_obj {
warranty_obj = boundary_obj;
warranty_state = Some(self.snapshot_mutable_state());
}
if let Ok(ev) = self.dictionary_reconstruction_ev(target, rho) {
if self.structural_coherence_collapse_detected()?.is_none() {
let candidate_uniformity = self.coordinate_uniformity_aggregate();
let candidate_obj = boundary_obj;
if prefer_candidate_state(
candidate_obj,
ev,
candidate_uniformity,
best_reconstruction_obj,
best_reconstruction_ev,
best_reconstruction_uniformity,
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL,
SAE_FINAL_EV_DEGRADATION_TOL,
) {
best_reconstruction_ev = ev;
best_reconstruction_obj = candidate_obj;
best_reconstruction_uniformity = candidate_uniformity;
best_reconstruction_state = Some(self.snapshot_mutable_state());
}
}
}
self.reclaim_arrow_assembly_workspace(&mut sys);
}
let mut inner_incumbent_restored = false;
if let Some(best_state) = best_reconstruction_state.as_ref()
&& best_reconstruction_obj.is_finite()
{
let final_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
let obj_scale = SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL
* (1.0 + final_obj.abs().max(best_reconstruction_obj.abs()));
if !(final_obj <= best_reconstruction_obj + obj_scale) {
let final_ev = self
.dictionary_reconstruction_ev(target, rho)
.unwrap_or(f64::NAN);
log::warn!(
"[#1026] restoring inner-fit incumbent: final penalized objective \
{final_obj:.6e} degraded past banked {best_reconstruction_obj:.6e} \
(EV {final_ev:.4} vs banked {best_reconstruction_ev:.4}) — \
non-monotone boundary-hook damage, not line-search descent"
);
self.restore_mutable_state(best_state)?;
inner_incumbent_restored = true;
state_moved = true;
moved_at.get_or_insert(StateMoveSite::InnerIncumbentRestore);
if self.frames_active() {
self.run_objective_guarded_hook(
target,
rho,
analytic_penalties,
obj_scale,
|term| {
term.refresh_active_frames_from_data(target).map_err(|err| {
format!("SaeManifoldTerm::run_joint_fit_arrow_schur: {err}")
})?;
Ok(())
},
)?;
}
}
}
if max_iter > 0
&& self.sweep_blocks_to_objective_fixed_point(
target,
rho,
analytic_penalties,
allow_heuristic_termination,
)?
{
state_moved = true;
moved_at.get_or_insert(StateMoveSite::ExitBlockSweep);
}
if let Some(bank) = warranty_state.as_ref() {
let final_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
let warranty_tol = SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL
* (1.0 + final_obj.abs().max(warranty_obj.abs()));
if !(final_obj <= warranty_obj + warranty_tol) {
log::warn!(
"[#2228] exit warranty: final penalized objective {final_obj:.6e} degraded \
past the best accepted boundary {warranty_obj:.6e}; restoring the banked \
state (non-monotone boundary-mover damage leaked to the exit)"
);
self.restore_mutable_state(bank)?;
state_moved = true;
moved_at.get_or_insert(StateMoveSite::ExitWarrantyRestore);
self.best_fit_incumbent = None;
}
}
let exact_objective_incumbent_restored = inner_incumbent_restored
&& best_reconstruction_state
.as_ref()
.is_some_and(|state| self.matches_mutable_state(state));
if exact_objective_incumbent_restored {
let consecutive_inner_restores = self
.best_fit_incumbent
.as_ref()
.filter(|prior| self.matches_mutable_state(&prior.state))
.map_or(1, |prior| {
prior.consecutive_inner_restores.saturating_add(1)
});
self.best_fit_incumbent = Some(SaeFitIncumbent {
state: self.snapshot_mutable_state(),
consecutive_inner_restores,
});
} else {
self.best_fit_incumbent = None;
}
let loss = self.loss(target, rho)?;
drop(faer_sequential_inner_fit);
Ok(JointFitOutcome {
loss,
termination,
state_moved,
moved_at,
})
}
pub(crate) fn run_objective_guarded_hook<F>(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
allowed_increase: f64,
hook: F,
) -> Result<bool, String>
where
F: FnOnce(&mut Self) -> Result<(), String>,
{
let pre_obj = self.penalized_objective_total(target, rho, analytic_penalties, 1.0)?;
let pre_state = self.snapshot_mutable_state();
hook(self)?;
let post_obj = self
.penalized_objective_total(target, rho, analytic_penalties, 1.0)
.unwrap_or(f64::INFINITY);
if !(post_obj.is_finite() && post_obj <= pre_obj + allowed_increase) {
self.restore_mutable_state(&pre_state)?;
return Ok(false);
}
Ok(!self.matches_mutable_state(&pre_state))
}
pub(crate) fn sweep_blocks_to_objective_fixed_point(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
discovery_lane: bool,
) -> Result<bool, String> {
let mut best_objective =
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)?;
if !best_objective.is_finite() {
return Ok(false);
}
let floor_rel = if discovery_lane {
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL
} else {
SAE_MANIFOLD_INNER_OBJECTIVE_STALL_FRACTION
};
let frames = self.frames_active();
let mut moved = false;
loop {
let snapshot = self.snapshot_mutable_state();
let round = self
.refit_decoder_least_squares_at_current_state(target, Some(rho))
.and_then(|()| {
if frames {
self.refresh_active_frames_from_data(target)
.map_err(|err| format!("sweep frame re-polar: {err}"))?;
}
Ok(())
})
.and_then(|()| self.seed_coords_by_decoder_projection(target))
.and_then(|()| self.refit_decoder_least_squares_at_current_state(target, Some(rho)))
.and_then(|()| {
if frames {
self.refresh_active_frames_from_data(target)
.map_err(|err| format!("sweep frame re-polar: {err}"))?;
}
Ok(())
})
.and_then(|()| {
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)
});
let accept_floor = floor_rel * (1.0 + best_objective.abs());
match round {
Ok(value) if value.is_finite() && value < best_objective - accept_floor => {
best_objective = value;
moved = true;
}
_ => {
self.restore_mutable_state(&snapshot)?;
break;
}
}
}
for _ in 0..(if discovery_lane { 2 } else { 0 }) {
let snapshot = self.snapshot_mutable_state();
let round = self
.anchor_logits_to_residual_ownership(target)
.and_then(|()| self.refit_decoder_least_squares_at_current_state(target, Some(rho)))
.and_then(|()| {
if frames {
self.refresh_active_frames_from_data(target)
.map_err(|err| format!("sweep frame re-polar: {err}"))?;
}
Ok(())
})
.and_then(|()| self.seed_coords_by_decoder_projection(target))
.and_then(|()| self.refit_decoder_least_squares_at_current_state(target, Some(rho)))
.and_then(|()| {
self.penalized_objective_total(target, rho, analytic_penalties, 1.0)
});
let accept_floor = floor_rel * (1.0 + best_objective.abs());
match round {
Ok(value) if value.is_finite() && value < best_objective - accept_floor => {
best_objective = value;
moved = true;
}
_ => {
self.restore_mutable_state(&snapshot)?;
break;
}
}
}
Ok(moved)
}
pub(crate) fn empty_decoder_gram_accumulator(&self) -> Vec<Array2<f64>> {
self.atoms
.iter()
.map(|atom| {
let m = atom.basis_size();
Array2::<f64>::zeros((m, m))
})
.collect()
}
pub(crate) fn accumulate_decoder_gram(&self, grams: &mut [Array2<f64>]) -> Result<(), String> {
let n = self.n_obs();
let assignments = self.assignment.assignments();
let weights: Vec<Array1<f64>> = (0..self.atoms.len())
.map(|atom_idx| {
let col = assignments.column(atom_idx);
col.mapv(|a| a * a)
})
.collect();
let cpu_one = |atom_idx: usize, gram: &mut Array2<f64>| {
let atom = &self.atoms[atom_idx];
let m = atom.basis_size();
let assign_col = assignments.column(atom_idx);
let mut weighted = vec![0.0_f64; m];
for row in 0..n {
let a_k = assign_col[row];
if a_k == 0.0 {
continue;
}
for col in 0..m {
weighted[col] = a_k * atom.basis_values[[row, col]];
}
for i in 0..m {
let wi = weighted[i];
if wi == 0.0 {
continue;
}
for j in 0..m {
gram[[i, j]] += wi * weighted[j];
}
}
}
};
let max_atom_gram_flops: u128 = self
.atoms
.iter()
.map(|atom| {
let m = atom.basis_size() as u128;
2u128 * (n as u128) * m * m
})
.max()
.unwrap_or(0);
let rt = if max_atom_gram_flops < crate::gpu::GpuDispatchPolicy::MIN_CALIBRATABLE_GEMM_FLOPS
{
None
} else {
crate::gpu::device_runtime::GpuRuntime::resolve(self.gpu_policy)
.map_err(|error| format!("decoder-Gram CUDA admission failed: {error}"))?
};
match rt {
None => {
for atom_idx in 0..self.atoms.len() {
if self.atoms[atom_idx].basis_size() == 0 {
continue;
}
cpu_one(atom_idx, &mut grams[atom_idx]);
}
}
Some(rt) => {
let mut items: Vec<usize> = (0..self.atoms.len())
.filter(|&i| self.atoms[i].basis_size() > 0)
.collect();
let device_grams: std::sync::Mutex<Vec<(usize, Array2<f64>)>> =
std::sync::Mutex::new(Vec::with_capacity(items.len()));
let declined: std::sync::Mutex<Vec<usize>> = std::sync::Mutex::new(Vec::new());
let atoms_ref = &self.atoms;
let weights_ref = &weights;
let ok = crate::gpu::pool::scatter_batched(rt, &mut items, |_, slice| {
for &atom_idx in slice.iter() {
let phi = atoms_ref[atom_idx].basis_values.view();
let w = weights_ref[atom_idx].view();
match crate::gpu::linalg_dispatch::try_fast_xt_diag_x(phi, w) {
Some(g) => device_grams
.lock()
.expect("device_grams mutex poisoned")
.push((atom_idx, g)),
None => declined
.lock()
.expect("declined mutex poisoned")
.push(atom_idx),
}
}
Some(())
});
match ok {
Some(()) => {
for (atom_idx, g) in device_grams
.into_inner()
.expect("device_grams mutex poisoned")
{
grams[atom_idx] += &g;
}
let declined = declined.into_inner().expect("declined mutex poisoned");
if !declined.is_empty() {
return Err(format!(
"decoder-Gram device path declined admitted atoms {declined:?}"
));
}
}
None => {
return Err(
"decoder-Gram device scatter declined after CUDA admission".to_string()
);
}
}
}
}
Ok(())
}
pub(crate) fn finalize_decoder_identifiability_audit(
&self,
grams: &[Array2<f64>],
n_total: usize,
) -> Result<(), String> {
let mut any_identifiable = false;
let mut audited_atoms = 0usize;
for (atom_idx, atom) in self.atoms.iter().enumerate() {
let m = atom.basis_size();
if m == 0 {
continue;
}
audited_atoms += 1;
let rank = gam_identifiability::audit::rank_of_gram(&grams[atom_idx], n_total)
.map_err(|e| {
format!(
"SaeManifoldTerm: pre-fit decoder audit (atom '{}'): \
Gram eigendecomposition failed: {e}",
atom.name,
)
})?;
if rank > 0 {
any_identifiable = true;
}
if rank < m {
let dropped = m - rank;
log::info!(
"[SAE-AUDIT] decoder atom '{}' weighted design is rank-deficient \
(rank={rank}/{m}, {dropped} weakly-identified column(s), n={n_total}); the \
Arrow-Schur ridge will regularise the deficient directions{}",
atom.name,
if rank == 0 {
" (atom is fully unweighted — parked dead, β_k → 0)"
} else {
""
},
);
}
}
if audited_atoms > 0 && !any_identifiable {
return Err(format!(
"SaeManifoldTerm: pre-fit identifiability audit: ALL {audited_atoms} decoder \
atoms have rank-0 weighted design (n={n_total}); the entire dictionary is \
unidentifiable — every atom's assignment weights vanish or every basis is \
degenerate, so the joint Arrow-Schur Newton system is singular with no \
ridge-recoverable signal"
));
}
Ok(())
}
fn synthesize_monomial_patch_evaluator(
atom: &SaeManifoldAtom,
) -> Option<Arc<dyn SaeBasisEvaluator>> {
match atom.basis_kind() {
SaeAtomBasisKind::EuclideanPatch
| SaeAtomBasisKind::Linear
| SaeAtomBasisKind::Poincare => {}
_ => return None,
}
let latent_dim = atom.latent_dim();
let target = atom.basis_size();
for degree in 0..=target {
if gam_terms::basis::monomial_exponents(latent_dim, degree).len() == target {
return crate::basis::EuclideanPatchEvaluator::new(latent_dim, degree)
.ok()
.map(|ev| Arc::new(ev) as Arc<dyn SaeBasisEvaluator>);
}
}
None
}
pub(crate) fn chunk_frozen_logits(&self, start: usize, end: usize) -> Option<Array2<f64>> {
self.assignment
.frozen_logits
.as_ref()
.map(|f| f.slice(ndarray::s![start..end, ..]).to_owned())
}
pub fn materialize_chunk(
&self,
chunk_logits: Array2<f64>,
chunk_coords: Vec<Array2<f64>>,
chunk_frozen_logits: Option<Array2<f64>>,
) -> Result<SaeManifoldTerm, String> {
let k_atoms = self.k_atoms();
if chunk_logits.ncols() != k_atoms {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: chunk_logits has {} cols but K={k_atoms}",
chunk_logits.ncols()
));
}
if chunk_coords.len() != k_atoms {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: chunk_coords has {} atoms but K={k_atoms}",
chunk_coords.len()
));
}
let n_chunk = chunk_logits.nrows();
let mut atoms = Vec::with_capacity(k_atoms);
for (atom_idx, atom) in self.atoms.iter().enumerate() {
let coords = &chunk_coords[atom_idx];
if coords.nrows() != n_chunk || coords.ncols() != atom.latent_dim() {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: atom {atom_idx} coords shape {:?} != ({n_chunk}, {})",
coords.dim(),
atom.latent_dim()
));
}
let synthesized_evaluator = match atom.basis_evaluator.as_ref() {
Some(_) => None,
None => Self::synthesize_monomial_patch_evaluator(atom),
};
let evaluator = match atom
.basis_evaluator
.as_ref()
.or(synthesized_evaluator.as_ref())
{
Some(evaluator) => evaluator,
None => {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: atom '{}' has no basis evaluator; a \
streaming fit must re-evaluate Φ(t) at each chunk's coordinates",
atom.name
));
}
};
let (phi, jet) = evaluator.evaluate(coords.view())?;
let m = atom.basis_size();
if phi.dim() != (n_chunk, m) {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: atom '{}' evaluator returned Φ {:?}, expected ({n_chunk}, {m})",
atom.name,
phi.dim()
));
}
if jet.dim() != (n_chunk, m, atom.latent_dim()) {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: atom '{}' evaluator returned jet {:?}, expected ({n_chunk}, {m}, {})",
atom.name,
jet.dim(),
atom.latent_dim()
));
}
let mut chunk_atom = SaeManifoldAtom::new_with_provided_function_gram(
atom.name.clone(),
atom.basis_kind().clone(),
atom.latent_dim(),
phi,
jet,
atom.decoder_coefficients().clone(),
atom.smooth_penalty().clone(),
)?;
chunk_atom.basis_evaluator = atom
.basis_evaluator
.clone()
.or_else(|| synthesized_evaluator.clone());
chunk_atom.basis_second_jet = atom.basis_second_jet.clone();
chunk_atom.decoder_frame = atom.decoder_frame.clone();
atoms.push(chunk_atom);
}
let coord_values: Vec<LatentCoordValues> = chunk_coords
.iter()
.zip(self.assignment.coords.iter())
.map(|(c, src)| {
LatentCoordValues::from_matrix_with_manifold(
c.view(),
LatentIdMode::None,
src.manifold().clone(),
)
})
.collect();
let mut assignment =
SaeAssignment::with_mode(chunk_logits, coord_values, self.assignment.mode)?;
assignment.ungated = self.assignment.ungated.clone();
assignment.ordered_beta_bernoulli_alpha_override =
self.assignment.ordered_beta_bernoulli_alpha_override;
if let Some(frozen) = chunk_frozen_logits {
if frozen.dim() != (n_chunk, k_atoms) {
return Err(format!(
"SaeManifoldTerm::materialize_chunk: chunk_frozen_logits shape {:?} != ({n_chunk}, {k_atoms})",
frozen.dim()
));
}
assignment.frozen_logits = Some(frozen);
}
let mut term = SaeManifoldTerm::new(atoms, assignment)?;
term.host_available_bytes = self.host_available_bytes;
term.temperature_schedule = self.temperature_schedule.clone();
if self.streaming_gates_frozen {
term.decoder_repulsion_gate = self.decoder_repulsion_gate.clone();
term.barrier_coactivation_gate = self.barrier_coactivation_gate.clone();
term.amplitude_barrier_gate = self.amplitude_barrier_gate;
term.streaming_gates_frozen = true;
}
Ok(term)
}
}
fn leading_direction_above_noise_floor(energies: &[f64]) -> bool {
if energies.is_empty() {
return false;
}
let mut sorted: Vec<f64> = energies.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let Some(&peak) = sorted.last() else {
return false;
};
if !(peak > 0.0) {
return false;
}
let m = sorted.len();
let noise_scale = sorted[(((m - 1) as f64) * 0.25).round() as usize];
let floor = (peak * 1e-12).max(noise_scale * (m as f64).max(2.0).log2());
peak > floor
}
pub(crate) fn union_output_frame_rank(frames: &[Array2<f64>], p: usize) -> usize {
let total_cols: usize = frames.iter().map(|q| q.ncols()).sum();
if p == 0 || total_cols == 0 {
return 0;
}
let mut stacked = Array2::<f64>::zeros((p, total_cols));
let mut col = 0usize;
for q in frames {
let m = q.ncols();
if m == 0 {
continue;
}
if q.nrows() != p {
return p;
}
for qc in 0..m {
for row in 0..p {
stacked[[row, col + qc]] = q[[row, qc]];
}
}
col += m;
}
let sv = match stacked.svd(false, false) {
Ok((_, sv, _)) => sv,
Err(_) => return p,
};
let max_sv = sv.iter().copied().fold(0.0_f64, f64::max);
if !(max_sv > 0.0) {
return 0;
}
let tol = crate::frames::SAE_FRAME_RANK_CUTOFF * max_sv;
sv.iter().filter(|&&v| v > tol).count().min(p)
}
fn union_output_frame_basis(frames: &[Array2<f64>], p: usize) -> Array2<f64> {
let total_cols: usize = frames.iter().map(|q| q.ncols()).sum();
if p == 0 || total_cols == 0 {
return Array2::<f64>::zeros((p, 0));
}
let mut stacked = Array2::<f64>::zeros((p, total_cols));
let mut col = 0usize;
for q in frames {
let m = q.ncols();
if m == 0 {
continue;
}
if q.nrows() != p {
return Array2::<f64>::zeros((p, 0));
}
for qc in 0..m {
for row in 0..p {
stacked[[row, col + qc]] = q[[row, qc]];
}
}
col += m;
}
let (u_opt, sv, _) = match stacked.svd(true, false) {
Ok(factors) => factors,
Err(_) => return Array2::<f64>::zeros((p, 0)),
};
let Some(u) = u_opt else {
return Array2::<f64>::zeros((p, 0));
};
let max_sv = sv.iter().copied().fold(0.0_f64, f64::max);
if !(max_sv > 0.0) {
return Array2::<f64>::zeros((p, 0));
}
let tol = crate::frames::SAE_FRAME_RANK_CUTOFF * max_sv;
let rank = sv
.iter()
.filter(|&&v| v > tol)
.count()
.min(p)
.min(u.ncols());
let mut basis = Array2::<f64>::zeros((p, rank));
for c in 0..rank {
for row in 0..p {
basis[[row, c]] = u[[row, c]];
}
}
basis
}
#[cfg(test)]
mod projection_policy_tests {
use super::*;
use crate::basis::{AmbientSphereHarmonicEvaluator, SaeBasisEvaluator};
use ndarray::array;
use std::sync::Arc;
#[test]
fn multivariate_compact_projection_skips_without_mutation() {
let coordinates = array![[0.36, 0.48, 0.8]];
let evaluator = Arc::new(AmbientSphereHarmonicEvaluator::new(2).unwrap());
let (phi, jet) = evaluator.evaluate(coordinates.view()).unwrap();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"sphere",
SaeAtomBasisKind::Sphere,
3,
phi,
jet,
Array2::<f64>::zeros((9, 3)),
Array2::<f64>::eye(9),
)
.unwrap()
.with_basis_evaluator(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((1, 1)),
vec![coordinates.clone()],
vec![SaeAtomBasisKind::Sphere.latent_manifold(2)],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let before = term.assignment.coords[0].as_matrix();
term.seed_coords_by_decoder_projection(Array2::<f64>::zeros((1, 2)).view())
.expect("compact multivariate chart is skipped, not an error");
assert_eq!(term.assignment.coords[0].as_matrix(), before);
}
#[test]
fn leading_direction_above_noise_floor_separates_signal_from_noise_2132() {
let noise: Vec<f64> = (0..20)
.map(|i| 1.0 + 0.15 * ((i % 5) as f64 - 2.0))
.collect();
assert!(
!leading_direction_above_noise_floor(&noise),
"a flat (pure-noise) residual spectrum must read as NO uncovered signal"
);
let mut signal = noise.clone();
signal[7] = 100.0;
assert!(
leading_direction_above_noise_floor(&signal),
"a residual spectrum with a dominant direction must read as uncovered signal"
);
assert!(
leading_direction_above_noise_floor(&[7.68, 7.68, 48.0, 48.0]),
"two circles filling all p=4 directions must read as uncovered signal, not noise"
);
assert!(
leading_direction_above_noise_floor(&[1.0e-9, 1.0e-9, 7.68, 7.68]),
"the weaker uncovered circle must still clear the floor after the dominant peel"
);
assert!(!leading_direction_above_noise_floor(&[]));
assert!(!leading_direction_above_noise_floor(&[0.0, 0.0, 0.0]));
}
}