use ndarray::{Array1, Array2};
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use gam_math::probability::beta_quantile;
use gam_problem::topology_certificates::CertificateLedger;
use gam_problem::{EstimationError, MetricProvenance};
use gam_solve::inference::residual_factor::{ResidualFactorInput, StructuredResidualModel};
use gam_solve::rho_optimizer::{OuterProblem, OuterResult, audit_stationary_point};
use gam_solve::structure_search::{MoveBudget, StructureMove};
use gam_terms::analytic_penalties::AnalyticPenaltyRegistry;
use gam_terms::inference::structure_evidence::StructureLedger;
use crate::structure_harvest;
use crate::tiered::Tier0Mean;
use super::{
AmortizedEncoderConsistency, AssignmentMode, ChartDegeneracyReport,
ChartNondegeneracyCertificate, CoordinateFidelityCertificate, CrossFitConfig,
CrossFitReport, SaeManifoldFitDiagnostics, SaeManifoldLoss, SaeManifoldOuterObjective,
SaeInnerKktScaleError, SaeManifoldRho, SaeManifoldTerm, SaeOuterTermination,
SaeShapeUncertainty,
SaeTrustDiagnostics, TopologyPersistenceCertificate, VanishedAtoms,
cross_fit_reconstruction_ev,
};
pub const STRUCTURED_RESIDUAL_PASSES_MAX: usize = 4;
fn validate_structured_residual_passes(passes: usize) -> Result<(), SaeFitError> {
if passes > STRUCTURED_RESIDUAL_PASSES_MAX {
return Err(SaeFitError::InvalidRequest(format!(
"structured_residual_passes={passes} exceeds the hard maximum {STRUCTURED_RESIDUAL_PASSES_MAX}"
)));
}
Ok(())
}
pub(crate) const STRUCTURED_RESIDUAL_MIN_REL_ENERGY: f64 = 1.0e-10;
fn promotion_alignment_threshold(factor_rank: usize) -> f64 {
if factor_rank <= 1 {
return 1.0;
}
let rank = factor_rank as f64;
beta_quantile(0.95, 0.5, (rank - 1.0) / 2.0)
.sqrt()
.clamp(0.0, 1.0)
}
#[derive(Clone, Debug)]
pub struct StructuredResidualPassDiagnostic {
pub pass: usize,
pub gamma: f64,
pub factor_rank: usize,
pub log_evidence: f64,
pub factor_energy: f64,
pub diagonal_mean: f64,
pub dispersion_before: f64,
pub dispersion_after: f64,
pub log_lambda_smooth_before: Vec<f64>,
pub log_lambda_smooth_after: Vec<f64>,
}
pub fn metric_provenance_label(provenance: MetricProvenance) -> &'static str {
match provenance {
MetricProvenance::Euclidean => "Euclidean",
MetricProvenance::OutputFisher { .. } => "OutputFisher",
MetricProvenance::OutputFisherDownstream { .. } => "OutputFisherDownstream",
MetricProvenance::BehavioralFisher { .. } => "BehavioralFisher",
MetricProvenance::WhitenedStructured { .. } => "WhitenedStructured",
}
}
fn sae_structured_residual_model(
term: &SaeManifoldTerm,
target: ndarray::ArrayView2<'_, f64>,
) -> Result<Option<StructuredResidualModel>, String> {
let fitted = term.try_fitted_target_aware(target, None)?;
let (n, p) = fitted.dim();
if n == 0 || p <= 1 {
return Ok(None);
}
if target.dim() != (n, p) {
return Err(format!(
"sae_structured_residual_model: target must be ({n}, {p}); got {:?}",
target.dim()
));
}
let mut residuals = target.to_owned();
residuals -= &fitted;
let target_energy: f64 = target.iter().map(|v| v * v).sum();
let residual_energy: f64 = residuals.iter().map(|v| v * v).sum();
if residual_energy <= STRUCTURED_RESIDUAL_MIN_REL_ENERGY * target_energy {
return Ok(None);
}
let assignments = term.assignment.assignments();
let activity: ndarray::Array1<f64> = (0..n).map(|r| assignments.row(r).sum()).collect();
let max_factor_rank = p.saturating_sub(1);
match StructuredResidualModel::fit(ResidualFactorInput {
residuals: residuals.view(),
activity: activity.view(),
max_factor_rank,
}) {
Ok(m) => Ok(Some(m)),
Err(e) => Err(format!(
"sae_structured_residual_model: structured residual-covariance fit failed: {e}"
)),
}
}
pub struct SaeFitReport {
pub term: SaeManifoldTerm,
pub rho: SaeManifoldRho,
pub loss: SaeManifoldLoss,
pub penalized_quasi_laplace_criterion: f64,
pub assignments: Array2<f64>,
pub fitted: Array2<f64>,
pub active_mask: Vec<bool>,
pub reconstruction_r2: f64,
pub reconstruction_optimism_reference: Option<CrossFitReport>,
pub outer_termination: SaeOuterTermination,
pub shape_uncertainty: SaeShapeUncertainty,
pub metric_provenance: &'static str,
pub structured_residual_diagnostics: Vec<StructuredResidualPassDiagnostic>,
pub trust_diagnostics: SaeTrustDiagnostics,
pub fit_diagnostics: SaeManifoldFitDiagnostics,
pub amortized_encoder_consistency: AmortizedEncoderConsistency,
pub chart_degeneracy: ChartDegeneracyReport,
pub certificate_ledger: CertificateLedger,
pub structure_search_json: Option<String>,
pub structure_certificate_json: Option<String>,
pub reported_log_alpha: f64,
}
#[derive(Clone, Debug, PartialEq)]
pub enum SaeParameterSpaceKktAudit {
Resolved {
scaled_gradient_max: f64,
stationarity_bound: f64,
},
Unresolved(SaeInnerKktScaleError),
}
impl SaeParameterSpaceKktAudit {
pub fn certifies(&self) -> bool {
match self {
Self::Resolved {
scaled_gradient_max,
stationarity_bound,
} => {
scaled_gradient_max.is_finite()
&& stationarity_bound.is_finite()
&& *stationarity_bound >= 0.0
&& scaled_gradient_max <= stationarity_bound
}
Self::Unresolved(_) => false,
}
}
}
impl std::fmt::Display for SaeParameterSpaceKktAudit {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Resolved {
scaled_gradient_max,
stationarity_bound,
} => write!(
formatter,
"scaled_max={scaled_gradient_max:.6e}, bound={stationarity_bound:.6e}"
),
Self::Unresolved(reason) => write!(formatter, "unresolved ({reason})"),
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct SaeInstalledInnerKktAudit {
pub raw_gradient_norm: f64,
pub quotient_gradient_norm: f64,
pub stationarity_bound: f64,
pub parameter_space: SaeParameterSpaceKktAudit,
}
impl SaeInstalledInnerKktAudit {
pub fn certifies(&self) -> bool {
SaeManifoldTerm::quasi_laplace_kkt_stationary(
self.raw_gradient_norm,
self.quotient_gradient_norm,
self.stationarity_bound,
) || self.parameter_space.certifies()
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct SaeExternalEvaluationReport {
pub inner: SaeInstalledInnerKktAudit,
pub outer_raw_gradient_norm: Option<f64>,
pub outer_projected_gradient_norm: Option<f64>,
pub outer_stationarity_bound: Option<f64>,
pub optimization_iterations: usize,
pub reason: String,
}
pub enum SaeExternalCertificationOutcome {
Certified(SaeFitReport),
NonStationary(SaeExternalEvaluationReport),
}
fn installed_inner_kkt_audit(
term: &mut SaeManifoldTerm,
target: ndarray::ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
registry: &AnalyticPenaltyRegistry,
) -> Result<SaeInstalledInnerKktAudit, SaeFitError> {
let system = term
.assemble_arrow_schur(target, rho, Some(registry))
.map_err(SaeFitError::Fit)?;
let raw_gradient_norm_sq = SaeManifoldTerm::system_grad_norm_sq(&system);
let raw_gradient_norm = raw_gradient_norm_sq.sqrt();
let lambda_smooth = rho.lambda_smooth_vec().map_err(SaeFitError::Fit)?;
let quotient_gradient_norm =
term.quotient_gradient_norm_from_system(&system, raw_gradient_norm_sq, &lambda_smooth);
let parameter_space = match SaeManifoldTerm::system_scaled_grad_max(&system) {
Ok(scaled_gradient_max) => match term.inner_iterate_max() {
Ok(iterate_max) => SaeParameterSpaceKktAudit::Resolved {
scaled_gradient_max,
stationarity_bound: super::SAE_MANIFOLD_INNER_GRAD_REL_TOL * iterate_max,
},
Err(reason) => SaeParameterSpaceKktAudit::Unresolved(reason),
},
Err(reason) => SaeParameterSpaceKktAudit::Unresolved(reason),
};
Ok(SaeInstalledInnerKktAudit {
raw_gradient_norm,
quotient_gradient_norm,
stationarity_bound: super::SAE_MANIFOLD_INNER_GRAD_REL_TOL * term.inner_iterate_scale(),
parameter_space,
})
}
fn external_nonstationary_report(
inner: SaeInstalledInnerKktAudit,
outer: Option<&OuterResult>,
reason: String,
) -> SaeExternalCertificationOutcome {
let stationarity = outer
.and_then(|result| result.criterion_certificate.as_ref())
.map(|certificate| &certificate.stationarity);
SaeExternalCertificationOutcome::NonStationary(SaeExternalEvaluationReport {
inner,
outer_raw_gradient_norm: stationarity.map(|certificate| certificate.raw_norm()),
outer_projected_gradient_norm: stationarity.map(|certificate| certificate.projected_norm()),
outer_stationarity_bound: stationarity.map(|certificate| certificate.bound()),
optimization_iterations: outer.map_or(0, |result| result.iterations),
reason,
})
}
pub struct SaeNullFitReport {
pub tier0: Tier0Mean,
pub fitted: Array2<f64>,
pub residual_sum_squares: f64,
pub reconstruction_r2: f64,
pub metric_provenance: &'static str,
pub vanished_atoms: VanishedAtoms,
}
pub enum SaeFitOutcome {
Manifold(SaeFitReport),
Null(SaeNullFitReport),
}
impl SaeFitOutcome {
pub fn manifold_or_error(self) -> Result<SaeFitReport, String> {
match self {
Self::Manifold(report) => Ok(report),
Self::Null(report) => Err(format!(
"fit selected the exact Tier-0 null after {} atom(s) vanished",
report.vanished_atoms.len()
)),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SaeFitStage {
Primary,
StructuredResidual {
pass: usize,
total_passes: usize,
},
}
impl SaeFitStage {
fn checkpoint_tag(self) -> String {
match self {
Self::Primary => "primary".to_string(),
Self::StructuredResidual { pass, total_passes } => {
format!("structured-residual-{pass}-of-{total_passes}")
}
}
}
}
impl std::fmt::Display for SaeFitStage {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Primary => f.write_str("primary"),
Self::StructuredResidual { pass, total_passes } => {
write!(f, "structured-residual pass {pass}/{total_passes}")
}
}
}
}
#[derive(Debug)]
pub enum SaeFitError {
InvalidRequest(String),
Fit(String),
OuterRun {
stage: SaeFitStage,
source: EstimationError,
},
OuterDidNotConverge {
stage: SaeFitStage,
result: Box<OuterResult>,
},
DegenerateChart {
atoms: Vec<usize>,
evidence: String,
report: Box<ChartDegeneracyReport>,
},
}
impl From<String> for SaeFitError {
fn from(message: String) -> Self {
Self::Fit(message)
}
}
impl std::fmt::Display for SaeFitError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidRequest(message) | Self::Fit(message) => f.write_str(message),
Self::DegenerateChart {
atoms, evidence, ..
} => write!(
f,
"SAE manifold fit produced a DEGENERATE CHART on load-bearing atom(s) {atoms:?}: \
every chart axis of those atoms collapsed to one point of its own manifold, so \
they decode to a constant and carry no displacements; refusing to mint a fit \
[{evidence}]"
),
Self::OuterRun { stage, source } => {
write!(f, "SAE manifold {stage} outer search failed: {source}")
}
Self::OuterDidNotConverge { stage, result } => {
let grad = result
.final_grad_norm
.map(|value| format!("{value:.6e}"))
.unwrap_or_else(|| "unmeasured".to_string());
write!(
f,
"SAE manifold {stage} outer search stopped without a stationarity \
certificate (iterations={}, final_value={:.6e}, final_grad_norm={}, \
plan={}, stop_reason={:?}, rho_checkpoint={:?}); refusing to mint a fit",
result.iterations,
result.final_value,
grad,
result.plan_used,
result.operator_stop_reason,
result.rho,
)
}
}
}
}
impl std::error::Error for SaeFitError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::OuterRun { source, .. } => Some(source),
Self::InvalidRequest(_)
| Self::Fit(_)
| Self::OuterDidNotConverge { .. }
| Self::DegenerateChart { .. } => None,
}
}
}
pub(crate) fn scope_outer_checkpoint_to_stage(
objective: &mut SaeManifoldOuterObjective,
stage: SaeFitStage,
) {
let mut path =
super::checkpoint::SaeFitCheckpoint::default_store_path(&objective.checkpoint_fingerprint);
path.set_file_name(format!(
"{}.{}.json",
objective.checkpoint_fingerprint.content_hash,
stage.checkpoint_tag(),
));
objective.checkpoint_path = path;
}
pub(crate) fn certify_outer_stage(
objective: SaeManifoldOuterObjective,
stage: SaeFitStage,
run_result: Result<OuterResult, EstimationError>,
) -> Result<SaeManifoldOuterObjective, SaeFitError> {
match run_result {
Ok(result) if result.converged() => {
let mut objective = objective;
match objective.certify_outer_result(&result) {
Ok(()) => Ok(objective),
Err(_) => Err(SaeFitError::OuterDidNotConverge {
stage,
result: Box::new(result),
}),
}
}
Ok(result) => Err(SaeFitError::OuterDidNotConverge {
stage,
result: Box::new(result),
}),
Err(source) => Err(SaeFitError::OuterRun { stage, source }),
}
}
enum SaeStageFit {
Certified(SaeManifoldOuterObjective),
Null(SaeNullFitReport),
}
fn exact_null_report(
state: super::SaeVanishedStageState,
target: &Array2<f64>,
metric_provenance: &'static str,
) -> SaeNullFitReport {
let p = target.ncols();
let mean = state
.term
.tier0_mean()
.cloned()
.unwrap_or_else(|| Array1::<f64>::zeros(p));
let fitted = Array2::from_shape_fn(target.dim(), |(_, col)| mean[col]);
let target_mean = target
.mean_axis(ndarray::Axis(0))
.unwrap_or_else(|| Array1::<f64>::zeros(p));
let mut residual_sum_squares = 0.0_f64;
let mut total_sum_squares = 0.0_f64;
for row in 0..target.nrows() {
for col in 0..p {
let residual = target[[row, col]] - fitted[[row, col]];
let centered = target[[row, col]] - target_mean[col];
residual_sum_squares += residual * residual;
total_sum_squares += centered * centered;
}
}
let reconstruction_r2 = if total_sum_squares > 0.0 {
crate::tiered::explained_variance_from_sums(residual_sum_squares, total_sum_squares)
} else {
0.0
};
SaeNullFitReport {
tier0: Tier0Mean { mean },
fitted,
residual_sum_squares,
reconstruction_r2,
metric_provenance,
vanished_atoms: state.atoms,
}
}
enum SaeBoundaryDisposition {
Restart {
term: SaeManifoldTerm,
rho: SaeManifoldRho,
},
Null(SaeNullFitReport),
}
fn vanished_disposition(
mut state: super::SaeVanishedStageState,
target: &Array2<f64>,
metric_provenance: &'static str,
) -> Result<SaeBoundaryDisposition, SaeFitError> {
if state.atoms.len() == state.term.k_atoms() {
return Ok(SaeBoundaryDisposition::Null(exact_null_report(
state,
target,
metric_provenance,
)));
}
let remove = state.atoms.as_btree_set();
structure_harvest::remove_atoms(&mut state.term, &mut state.rho, &remove)
.map_err(SaeFitError::Fit)?;
Ok(SaeBoundaryDisposition::Restart {
term: state.term,
rho: state.rho,
})
}
fn fit_outer_stage_to_boundary(
mut term: SaeManifoldTerm,
target: &Array2<f64>,
registry: &AnalyticPenaltyRegistry,
mut rho: SaeManifoldRho,
max_iter: usize,
learning_rate: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
run_outer_rho_search: bool,
stage: SaeFitStage,
cancel_flag: &Arc<AtomicBool>,
metric_provenance: &'static str,
) -> Result<SaeStageFit, SaeFitError> {
loop {
let mut objective = SaeManifoldOuterObjective::new(
term,
target.clone(),
Some(registry.clone()),
rho,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
);
let rho_flat = objective.current_rho_flat();
scope_outer_checkpoint_to_stage(&mut objective, stage);
objective.set_cancel_flag(Arc::clone(cancel_flag));
let boundary = if run_outer_rho_search {
let search_init_rho = match objective.try_resume_from_checkpoint(rho_flat.len())? {
Some(banked) => ndarray::Array1::from(banked),
None => rho_flat,
};
let problem =
OuterProblem::new(search_init_rho.len()).with_initial_rho(search_init_rho);
match problem.run(&mut objective, "SAE manifold") {
Ok(result) if result.converged() => {
return certify_outer_stage(objective, stage, Ok(result))
.map(SaeStageFit::Certified);
}
Ok(result) => {
let terminal_rho = Array1::from(result.rho.clone());
match objective.vanished_stage_state_at(terminal_rho.view()) {
Ok(Some(state)) => Some(state),
Ok(None) => {
return Err(SaeFitError::OuterDidNotConverge {
stage,
result: Box::new(result),
});
}
Err(error) => {
log::debug!(
"SAE vanished-atom boundary probe refused at the terminal rho \
({error}); reporting the outer non-convergence it classifies"
);
return Err(SaeFitError::OuterDidNotConverge {
stage,
result: Box::new(result),
});
}
}
}
Err(source) => {
let terminal_rho = objective.current_rho_flat();
match objective.vanished_stage_state_at(terminal_rho.view()) {
Ok(Some(state)) => Some(state),
Ok(None) => {
return Err(SaeFitError::OuterRun { stage, source });
}
Err(error) => {
log::debug!(
"SAE vanished-atom boundary probe refused at the terminal rho \
({error}); reporting the outer-run failure it classifies"
);
return Err(SaeFitError::OuterRun { stage, source });
}
}
}
}
} else {
match objective.fit_at_fixed_rho(rho_flat.view()) {
Ok(()) => return Ok(SaeStageFit::Certified(objective)),
Err(original) => match objective.vanished_stage_state_at(rho_flat.view()) {
Ok(Some(state)) => Some(state),
Ok(None) => return Err(SaeFitError::Fit(original)),
Err(error) => {
log::debug!(
"SAE vanished-atom boundary probe refused at the fixed rho \
({error}); reporting the fit failure it classifies"
);
return Err(SaeFitError::Fit(original));
}
},
}
};
let state = boundary.expect("each non-returning branch installs a boundary state");
objective.remove_checkpoint();
match vanished_disposition(state, target, metric_provenance)? {
SaeBoundaryDisposition::Restart {
term: reduced_term,
rho: reduced_rho,
} => {
term = reduced_term;
rho = reduced_rho;
}
SaeBoundaryDisposition::Null(report) => return Ok(SaeStageFit::Null(report)),
}
}
}
pub struct SaeFitRequest {
pub reconstruction_optimism_folds: Option<usize>,
pub base_term: SaeManifoldTerm,
pub target: Array2<f64>,
pub registry: AnalyticPenaltyRegistry,
pub initial_rho: SaeManifoldRho,
pub max_iter: usize,
pub learning_rate: f64,
pub ridge_ext_coord: f64,
pub ridge_beta: f64,
pub alpha: f64,
pub isometry_pin_active: bool,
pub metric_provenance: &'static str,
pub promote_from_residual: bool,
pub run_structure_search: bool,
pub run_outer_rho_search: bool,
pub structured_residual_passes: usize,
pub cancel: Option<Arc<AtomicBool>>,
}
pub fn run_sae_manifold_fit(mut request: SaeFitRequest) -> Result<SaeFitOutcome, SaeFitError> {
validate_structured_residual_passes(request.structured_residual_passes)?;
if request.base_term.tier0_mean().is_some() {
return run_sae_manifold_fit_on_target(request);
}
let Some(mu) = request.target.mean_axis(ndarray::Axis(0)) else {
return run_sae_manifold_fit_on_target(request);
};
for mut row in request.target.rows_mut() {
row -= μ
}
let tier0_residual_sum_squares = request
.target
.iter()
.map(|value| value * value)
.sum::<f64>();
let standardizable = request.base_term.behavior.is_none()
&& request.base_term.crosscoder_layout.is_none()
&& request.target.nrows() > 0;
let sigma = if standardizable {
let n = request.target.nrows() as f64;
let mut sigma = Array1::<f64>::zeros(request.target.ncols());
for (col_idx, col) in request.target.columns().into_iter().enumerate() {
sigma[col_idx] = (col.iter().map(|v| v * v).sum::<f64>() / n).sqrt();
}
let sigma_max = sigma.iter().cloned().fold(0.0_f64, f64::max);
if sigma_max.is_finite() && sigma_max > 0.0 {
let floor = sigma_max * f64::EPSILON.sqrt();
for s in sigma.iter_mut() {
if !(*s > floor) {
*s = 1.0;
}
}
for mut row in request.target.rows_mut() {
row /= σ
}
for atom in &mut request.base_term.atoms {
for (col_idx, s) in sigma.iter().enumerate() {
for coeff in atom.decoder_coefficients_mut().column_mut(col_idx).iter_mut() {
*coeff /= *s;
}
}
}
Some(sigma)
} else {
None
}
} else {
None
};
let mut outcome = run_sae_manifold_fit_on_target(request)?;
match &mut outcome {
SaeFitOutcome::Manifold(report) => {
report
.term
.set_tier0_mean(mu.clone())
.map_err(SaeFitError::Fit)?;
if let Some(sigma) = sigma.as_ref() {
report
.term
.set_tier0_scale(sigma.clone())
.map_err(SaeFitError::Fit)?;
}
lift_tier0_rows(&mut report.fitted, &mu, sigma.as_ref());
}
SaeFitOutcome::Null(report) => {
report.tier0 = Tier0Mean { mean: mu.clone() };
lift_tier0_rows(&mut report.fitted, &mu, sigma.as_ref());
report.residual_sum_squares = tier0_residual_sum_squares;
report.reconstruction_r2 = 0.0;
}
}
Ok(outcome)
}
#[cfg(test)]
mod structured_pass_request_tests {
use super::*;
#[test]
fn explicit_structured_pass_count_above_hard_cap_is_rejected_2267() {
assert!(validate_structured_residual_passes(0).is_ok());
assert!(validate_structured_residual_passes(STRUCTURED_RESIDUAL_PASSES_MAX).is_ok());
assert!(matches!(
validate_structured_residual_passes(STRUCTURED_RESIDUAL_PASSES_MAX + 1),
Err(SaeFitError::InvalidRequest(_))
));
}
}
#[cfg(test)]
mod vanished_stage_tests {
use super::*;
use crate::basis::EuclideanPatchEvaluator;
use crate::manifold::{AssignmentMode, SaeAssignment, SaeAtomBasisKind, SaeManifoldAtom};
use gam_terms::latent::LatentManifold;
use ndarray::Array3;
fn fixed_boundary_term(k: usize, live_first: bool) -> (SaeManifoldTerm, SaeManifoldRho) {
let n = 8usize;
let p = 2usize;
let mut atoms = Vec::with_capacity(k);
for atom in 0..k {
let mut decoder = Array2::<f64>::zeros((1, p));
if atom == 0 && live_first {
decoder[[0, 0]] = 1.0;
}
let evaluator = Arc::new(
EuclideanPatchEvaluator::new(1, 0).expect("degree-zero Euclidean evaluator"),
);
atoms.push(
SaeManifoldAtom::new_with_provided_function_gram(
format!("atom{atom}"),
SaeAtomBasisKind::EuclideanPatch,
1,
Array2::<f64>::ones((n, 1)),
Array3::<f64>::zeros((n, 1, 1)),
decoder,
Array2::<f64>::eye(1),
)
.unwrap()
.with_basis_second_jet(evaluator),
);
}
let mut logits = Array2::<f64>::zeros((n, k));
if k > 1 {
logits.column_mut(1).fill(-40.0);
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![Array2::<f64>::zeros((n, 1)); k],
vec![LatentManifold::Euclidean; k],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); k]);
(term, rho)
}
#[test]
fn committed_k2_boundary_compacts_and_fixed_rho_restart_certifies_k1() {
let (term, rho) = fixed_boundary_term(2, true);
let mut target = Array2::<f64>::zeros((8, 2));
target.column_mut(0).fill(1.0);
let registry = AnalyticPenaltyRegistry::new();
let cancel = Arc::new(AtomicBool::new(false));
let stage = fit_outer_stage_to_boundary(
term,
&target,
®istry,
rho,
0,
1.0,
1.0e-6,
1.0e-6,
false,
SaeFitStage::Primary,
&cancel,
"Euclidean",
)
.expect("proper vanished subset must restart on the compacted stratum");
let SaeStageFit::Certified(objective) = stage else {
panic!("one live atom must not collapse to the Tier-0 null");
};
let fitted = objective
.into_fitted()
.expect("reduced fixed-rho state must carry an inner certificate");
assert_eq!(fitted.term.k_atoms(), 1);
assert_eq!(fitted.rho.log_lambda_smooth.len(), 1);
assert_eq!(fitted.rho.log_ard.len(), 1);
assert!(fitted.penalized_quasi_laplace_criterion.is_finite());
}
#[test]
fn committed_k1_boundary_returns_exact_tier0_null_not_manifold_fit() {
let (term, rho) = fixed_boundary_term(1, false);
let target = Array2::<f64>::ones((8, 2));
let registry = AnalyticPenaltyRegistry::new();
let cancel = Arc::new(AtomicBool::new(false));
let stage = fit_outer_stage_to_boundary(
term,
&target,
®istry,
rho,
0,
1.0,
1.0e-6,
1.0e-6,
false,
SaeFitStage::Primary,
&cancel,
"Euclidean",
)
.expect("all-vanished state must be an exact structural result");
let SaeStageFit::Null(report) = stage else {
panic!("K=1 vanished boundary must not mint a manifold fit");
};
assert_eq!(report.vanished_atoms.iter().collect::<Vec<_>>(), vec![0]);
assert_eq!(report.tier0.mean, Array1::<f64>::zeros(2));
assert!(report.residual_sum_squares.is_finite());
assert_eq!(report.fitted, Array2::<f64>::zeros((8, 2)));
}
}
fn lift_tier0_rows(recon: &mut Array2<f64>, mu: &Array1<f64>, sigma: Option<&Array1<f64>>) {
for mut row in recon.rows_mut() {
if let Some(sigma) = sigma {
row *= sigma;
}
row += mu;
}
}
struct SaeFinalizeRequest<'a> {
z: &'a Array2<f64>,
registry: &'a AnalyticPenaltyRegistry,
run_structure_search: bool,
shape_uncertainty_invalidated: bool,
carried_loss: Option<SaeManifoldLoss>,
structured_residual_diagnostics: Vec<StructuredResidualPassDiagnostic>,
outer_termination: SaeOuterTermination,
penalized_quasi_laplace_criterion: f64,
metric_provenance: &'static str,
alpha: f64,
isometry_pin_active: bool,
max_iter: usize,
learning_rate: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
entry_label: &'a str,
reconstruction_optimism_folds: Option<usize>,
}
fn finalize_sae_fit_report(
mut term: SaeManifoldTerm,
mut rho: SaeManifoldRho,
mut shape_uncertainty: SaeShapeUncertainty,
request: SaeFinalizeRequest<'_>,
) -> Result<SaeFitReport, SaeFitError> {
let SaeFinalizeRequest {
z,
registry,
run_structure_search,
shape_uncertainty_invalidated,
carried_loss,
structured_residual_diagnostics,
outer_termination,
penalized_quasi_laplace_criterion,
metric_provenance,
alpha,
isometry_pin_active,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
entry_label,
reconstruction_optimism_folds,
} = request;
let (n_obs, p_out) = z.dim();
term.record_fit_data_collapse_if_needed(z.view(), &rho, max_iter)?;
let mut structure_ledger = StructureLedger::new();
let mut structure_changed = false;
let structure_search_json = 'structure: {
if !run_structure_search {
break 'structure None;
}
let harvest_params = structure_harvest::HarvestParams {
max_fusions: 1,
max_fissions: 1,
max_births: 1,
};
let refit_params = structure_harvest::ProductionRefitParams {
inner_max_iter: max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
};
let budget = MoveBudget {
max_moves: 1,
alpha: 0.05,
};
let n_shards = n_obs.saturating_sub(n_obs / 2).max(1);
let config = structure_harvest::RoundDriverConfig {
n_shards,
budget,
harvest_params,
curl: None,
};
match structure_harvest::run_production_structure_search(
term,
rho,
z.view(),
config,
refit_params,
&mut structure_ledger,
) {
Ok(result) => {
structure_changed = result.structure_changed();
term = result.term;
rho = result.rho;
Some(structure_harvest::rounds_to_json(&result.rounds)?)
}
Err(e) => {
return Err(SaeFitError::Fit(format!(
"structure search around {entry_label} failed: {e}"
)));
}
}
};
term.clear_row_loss_weights();
let k_atoms = term.k_atoms();
if structure_changed || shape_uncertainty_invalidated {
shape_uncertainty = term.recompute_joint_shape_uncertainty(
z.view(),
&rho,
Some(registry),
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
)?;
}
term.set_certificate_dispersion(shape_uncertainty.dispersion)?;
term.set_atom_inner_fits(z.view(), shape_uncertainty.dispersion)?;
if shape_uncertainty.atoms.len() != k_atoms {
return Err(SaeFitError::Fit(
"final joint shape uncertainty does not match the final atom count".to_string(),
));
}
for (atom_idx, uncertainty) in shape_uncertainty.atoms.iter().enumerate() {
match (
&uncertainty.band_coords,
&uncertainty.band_mean,
&uncertainty.band_sd,
) {
(None, None, None) => {
if uncertainty.decoder_covariance.is_some() || uncertainty.band_sd_robust.is_some()
{
return Err(SaeFitError::Fit(format!(
"atom {atom_idx} has a partial unavailable shape-uncertainty payload"
)));
}
}
(Some(coords), Some(mean), Some(sd)) => {
if coords.nrows() != mean.nrows()
|| mean.dim() != sd.dim()
|| coords
.iter()
.chain(mean.iter())
.chain(sd.iter())
.any(|value| !value.is_finite())
{
return Err(SaeFitError::Fit(format!(
"atom {atom_idx} has inconsistent or non-finite joint shape uncertainty"
)));
}
if let Some(covariance) = &uncertainty.decoder_covariance
&& covariance.iter().any(|value| !value.is_finite())
{
return Err(SaeFitError::Fit(format!(
"atom {atom_idx} has non-finite decoder covariance"
)));
}
if let Some(robust) = &uncertainty.band_sd_robust
&& (robust.dim() != sd.dim() || robust.iter().any(|value| !value.is_finite()))
{
return Err(SaeFitError::Fit(format!(
"atom {atom_idx} has inconsistent robust shape uncertainty"
)));
}
}
_ => {
return Err(SaeFitError::Fit(format!(
"atom {atom_idx} has a partial joint shape-uncertainty band"
)));
}
}
}
term.assignment
.validate_rho_domain(&rho)
.map_err(SaeFitError::Fit)?;
let ard_variances: Vec<Option<Array1<f64>>> = term
.validated_ard_precisions(&rho)
.map_err(SaeFitError::Fit)?
.iter()
.map(|precision| {
if precision.is_empty() {
None
} else {
Some(precision.mapv(|alpha| alpha.recip()))
}
})
.collect();
let assignments = term.assignment.assignments();
let fitted = term.try_fitted_target_aware(z.view(), Some(&rho))?;
term.record_fit_data_collapse_if_needed(z.view(), &rho, max_iter)?;
let trust_diagnostics = term.trust_diagnostics_report(assignments.view())?;
let fit_diagnostics = term.fit_diagnostics_report(
Some(&ard_variances),
isometry_pin_active,
Some(shape_uncertainty.dispersion),
fitted.view(),
Some(assignments.view()),
)?;
let amortized_encoder_consistency = term.amortized_encoder_consistency(z.view(), &rho)?;
let chart_degeneracy = term.chart_degeneracy_report();
let collapsed_atoms = chart_degeneracy.atoms_without_a_chart();
let refused_atoms = if collapsed_atoms.is_empty() {
Vec::new()
} else {
let load_bearing = chart_degeneracy.chart_less_load_bearing_atoms(assignments.view());
if load_bearing.is_empty() && collapsed_atoms.len() == chart_degeneracy.atom_count {
collapsed_atoms
} else {
load_bearing
}
};
if !refused_atoms.is_empty() {
let evidence = chart_degeneracy.atom_evidence(&refused_atoms);
return Err(SaeFitError::DegenerateChart {
atoms: refused_atoms,
evidence,
report: Box::new(chart_degeneracy),
});
}
let mut certificate_ledger = CertificateLedger::new();
certificate_ledger.record(&ChartNondegeneracyCertificate::new(&chart_degeneracy));
certificate_ledger.record(&fit_diagnostics.residual_gauge);
certificate_ledger.record(&CoordinateFidelityCertificate::new(
&fit_diagnostics.coordinate_fidelity,
));
certificate_ledger.record(&TopologyPersistenceCertificate::new(
&fit_diagnostics.topology_persistence,
));
if let Some(report) = &fit_diagnostics.incoherence_report {
certificate_ledger.record(report);
}
let active_mask: Vec<bool> = (0..k_atoms)
.map(|atom_idx| assignments.column(atom_idx).sum() > 1.0e-8)
.collect();
let mut means = vec![0.0_f64; p_out];
for row in 0..n_obs {
for out_col in 0..p_out {
means[out_col] += z[[row, out_col]];
}
}
if n_obs > 0 {
let inv_n = 1.0 / n_obs as f64;
for mean in means.iter_mut() {
*mean *= inv_n;
}
}
let mut rss = 0.0_f64;
let mut tss = 0.0_f64;
for row in 0..n_obs {
for out_col in 0..p_out {
let residual = z[[row, out_col]] - fitted[[row, out_col]];
let centered = z[[row, out_col]] - means[out_col];
rss += residual * residual;
tss += centered * centered;
}
}
let reconstruction_r2 = if tss > 0.0 {
crate::tiered::explained_variance_from_sums(rss, tss)
} else {
0.0
};
let reconstruction_optimism_reference = reconstruction_optimism_folds.and_then(|k_folds| {
let q = term
.atoms
.iter()
.map(|atom| atom.latent_dim())
.sum::<usize>()
.min(p_out.saturating_sub(1));
if q == 0 || k_folds < 2 {
return None;
}
cross_fit_reconstruction_ev(z.view(), CrossFitConfig { k_folds, seed: 0 }, q).ok()
});
let reported_log_alpha = match term.assignment.mode {
AssignmentMode::OrderedBetaBernoulli { alpha, .. } => alpha.ln(),
_ => alpha.ln(),
};
let structure_certificate_json = structure_search_json
.as_ref()
.map(|_| {
structure_ledger
.certify(0.05)
.map_err(|error| error.to_string())
.and_then(|certificate| {
serde_json::to_string(&certificate).map_err(|error| error.to_string())
})
})
.transpose()?;
let loss = match carried_loss {
Some(loss) => loss,
None => term.loss(z.view(), &rho)?,
};
Ok(SaeFitReport {
term,
rho,
loss,
penalized_quasi_laplace_criterion,
assignments,
fitted,
active_mask,
reconstruction_r2,
reconstruction_optimism_reference,
outer_termination,
shape_uncertainty,
metric_provenance,
structured_residual_diagnostics,
trust_diagnostics,
fit_diagnostics,
amortized_encoder_consistency,
chart_degeneracy,
certificate_ledger,
structure_search_json,
structure_certificate_json,
reported_log_alpha,
})
}
fn run_sae_manifold_fit_on_target(request: SaeFitRequest) -> Result<SaeFitOutcome, SaeFitError> {
let SaeFitRequest {
base_term,
target: z,
registry,
initial_rho: init_rho,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
alpha,
isometry_pin_active,
metric_provenance: metric_provenance_initial,
promote_from_residual,
run_structure_search,
run_outer_rho_search,
structured_residual_passes,
cancel,
reconstruction_optimism_folds,
} = request;
let (n_obs, p_out) = z.dim();
let mut metric_provenance: &'static str = metric_provenance_initial;
let init_rho = init_rho.for_assignment(base_term.assignment.mode);
base_term
.assignment
.validate_rho_domain(&init_rho)
.map_err(SaeFitError::Fit)?;
let cancel_flag = cancel.unwrap_or_else(|| Arc::new(AtomicBool::new(false)));
let mut objective = match fit_outer_stage_to_boundary(
base_term,
&z,
®istry,
init_rho,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
run_outer_rho_search,
SaeFitStage::Primary,
&cancel_flag,
metric_provenance,
)? {
SaeStageFit::Certified(objective) => objective,
SaeStageFit::Null(report) => return Ok(SaeFitOutcome::Null(report)),
};
let mut shape_uncertainty = objective.decoder_shape_uncertainty()?;
objective.remove_checkpoint();
let fitted_result = objective.into_fitted().map_err(SaeFitError::Fit)?;
let mut finalization_invalidated_shape_uncertainty =
fitted_result.invalidates_pre_final_shape_uncertainty();
let mut outer_termination = fitted_result.termination;
let mut term = fitted_result.term;
let mut rho = fitted_result.rho;
let mut loss = fitted_result.loss;
let mut penalized_quasi_laplace_criterion = fitted_result.penalized_quasi_laplace_criterion;
let structured_passes = structured_residual_passes;
let mut structured_residual_diagnostics: Vec<StructuredResidualPassDiagnostic> = Vec::new();
if structured_passes > 0 && metric_provenance == "Euclidean" {
let mut prev_model: Option<StructuredResidualModel> = None;
const PROMOTION_ENERGY_FLOOR_MULT: f64 = 1.0;
const PROMOTION_NURSERY_MIN_PASSES: usize = 2;
let mut nursery: Vec<(Array1<f64>, usize)> = Vec::new();
let mut total_passes = structured_passes;
let mut pass = 0usize;
while pass < total_passes {
let Some(model) = sae_structured_residual_model(&term, z.view())? else {
break;
};
let gamma = (pass as f64 + 1.0) / (total_passes as f64 + 1.0);
let metric = model.row_metric_damped(n_obs, gamma, prev_model.as_ref())?;
let installed_label = metric_provenance_label(metric.provenance());
let factor_energy = model.factor().iter().map(|v| v * v).sum::<f64>();
let diagonal_mean = model.diagonal().iter().copied().sum::<f64>() / p_out as f64;
let dispersion_before = shape_uncertainty.dispersion;
let log_lambda_smooth_before = rho.log_lambda_smooth.clone();
term.set_row_metric(metric)?;
let stage = SaeFitStage::StructuredResidual {
pass: pass + 1,
total_passes,
};
let mut objective = match fit_outer_stage_to_boundary(
term,
&z,
®istry,
rho,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
run_outer_rho_search,
stage,
&cancel_flag,
installed_label,
)? {
SaeStageFit::Certified(objective) => objective,
SaeStageFit::Null(report) => return Ok(SaeFitOutcome::Null(report)),
};
shape_uncertainty = objective.decoder_shape_uncertainty()?;
objective.remove_checkpoint();
let fitted_result = objective.into_fitted().map_err(SaeFitError::Fit)?;
finalization_invalidated_shape_uncertainty =
fitted_result.invalidates_pre_final_shape_uncertainty();
outer_termination = fitted_result.termination;
term = fitted_result.term;
rho = fitted_result.rho;
loss = fitted_result.loss;
penalized_quasi_laplace_criterion = fitted_result.penalized_quasi_laplace_criterion;
structured_residual_diagnostics.push(StructuredResidualPassDiagnostic {
pass: pass + 1,
gamma,
factor_rank: model.factor_rank(),
log_evidence: model.log_evidence(),
factor_energy,
diagonal_mean,
dispersion_before,
dispersion_after: shape_uncertainty.dispersion,
log_lambda_smooth_before,
log_lambda_smooth_after: rho.log_lambda_smooth.clone(),
});
metric_provenance = installed_label;
let prev_for_promotion = if promote_from_residual {
prev_model.as_ref()
} else {
None
};
if let Some(prev) = prev_for_promotion {
let align_min = promotion_alignment_threshold(model.factor_rank());
let cands = model.promotion_candidates(
Some(prev),
align_min,
PROMOTION_ENERGY_FLOOR_MULT,
)?;
let mut seen = vec![false; nursery.len()];
for cand in &cands {
let hit = nursery
.iter()
.position(|(d, _)| cand.direction.dot(d).abs() >= align_min);
match hit {
Some(i) => {
nursery[i].0 = cand.direction.clone();
nursery[i].1 += 1;
seen[i] = true;
}
None => {
nursery.push((cand.direction.clone(), 1));
seen.push(true);
}
}
}
let mut keep = seen.into_iter();
nursery.retain(|_| keep.next().unwrap_or(false));
if !nursery.is_empty()
&& pass + 1 == total_passes
&& total_passes < STRUCTURED_RESIDUAL_PASSES_MAX
{
total_passes += 1;
}
let matured = if pass + 1 < total_passes {
nursery
.iter()
.find(|(_, count)| *count >= PROMOTION_NURSERY_MIN_PASSES)
.map(|(dir, _)| dir.clone())
} else {
None
};
if let Some(dir) = matured {
let m = term.atoms[0].basis_size();
let mut decoder = Array2::<f64>::zeros((m, p_out));
for out in 0..p_out {
decoder[[0, out]] = dir[out];
}
let (grown_term, grown_rho) = structure_harvest::apply_structure_move(
&term,
&rho,
&StructureMove::Birth { candidate: 0 },
std::slice::from_ref(&decoder),
)?;
term = grown_term;
rho = grown_rho;
nursery.retain(|(d, _)| d.dot(&dir).abs() < align_min);
}
}
prev_model = Some(model);
pass += 1;
}
}
let report = finalize_sae_fit_report(
term,
rho,
shape_uncertainty,
SaeFinalizeRequest {
z: &z,
registry: ®istry,
run_structure_search,
shape_uncertainty_invalidated: finalization_invalidated_shape_uncertainty,
carried_loss: Some(loss),
structured_residual_diagnostics,
outer_termination,
penalized_quasi_laplace_criterion,
metric_provenance,
alpha,
isometry_pin_active,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
reconstruction_optimism_folds,
entry_label: "SAE fit",
},
)?;
Ok(SaeFitOutcome::Manifold(report))
}
pub struct SaeCertifyRequest {
pub base_term: SaeManifoldTerm,
pub target: Array2<f64>,
pub registry: AnalyticPenaltyRegistry,
pub initial_rho: SaeManifoldRho,
pub max_iter: usize,
pub learning_rate: f64,
pub ridge_ext_coord: f64,
pub ridge_beta: f64,
pub alpha: f64,
pub isometry_pin_active: bool,
pub metric_provenance: &'static str,
pub run_structure_search: bool,
}
pub fn run_sae_manifold_certify(
request: SaeCertifyRequest,
) -> Result<SaeExternalCertificationOutcome, SaeFitError> {
let SaeCertifyRequest {
base_term,
target: z,
registry,
initial_rho,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
alpha,
isometry_pin_active,
metric_provenance,
run_structure_search,
} = request;
let mut term = base_term;
let rho = initial_rho.for_assignment(term.assignment.mode);
term.assignment
.validate_rho_domain(&rho)
.map_err(SaeFitError::Fit)?;
let inner_audit = installed_inner_kkt_audit(&mut term, z.view(), &rho, ®istry)?;
if !inner_audit.certifies() {
return Ok(external_nonstationary_report(
inner_audit.clone(),
None,
format!(
"installed external state failed inner KKT stationarity: raw={:.6e}, \
quotient={:.6e}, bound={:.6e}, parameter-space={}",
inner_audit.raw_gradient_norm,
inner_audit.quotient_gradient_norm,
inner_audit.stationarity_bound,
inner_audit.parameter_space,
),
));
}
let rho_flat = rho.to_flat();
let mut objective = SaeManifoldOuterObjective::new(
term,
z.clone(),
Some(registry.clone()),
rho,
0,
learning_rate,
ridge_ext_coord,
ridge_beta,
)
.for_installed_state_audit();
let outer_result = match audit_stationary_point(
&mut objective,
rho_flat,
"SAE external installed-state audit",
) {
Ok(result) => result,
Err(rejection) => {
return Ok(external_nonstationary_report(
inner_audit,
Some(&rejection.result),
rejection.source.to_string(),
));
}
};
objective
.certify_installed_state_audit(&outer_result)
.map_err(SaeFitError::Fit)?;
let shape_uncertainty = objective.decoder_shape_uncertainty()?;
let fitted_result = objective.into_fitted().map_err(SaeFitError::Fit)?;
let term = fitted_result.term;
let rho = fitted_result.rho;
let penalized_quasi_laplace_criterion = fitted_result.penalized_quasi_laplace_criterion;
let outer_termination = fitted_result.termination;
let report = finalize_sae_fit_report(
term,
rho,
shape_uncertainty,
SaeFinalizeRequest {
z: &z,
registry: ®istry,
run_structure_search,
shape_uncertainty_invalidated: false,
carried_loss: None,
structured_residual_diagnostics: Vec::new(),
outer_termination,
penalized_quasi_laplace_criterion,
metric_provenance,
alpha,
isometry_pin_active,
max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
reconstruction_optimism_folds: None,
entry_label: "SAE certify entry",
},
)?;
Ok(SaeExternalCertificationOutcome::Certified(report))
}
#[cfg(test)]
mod tests {
use super::promotion_alignment_threshold;
#[test]
fn promotion_alignment_threshold_is_core_owned_and_rank_aware() {
assert_eq!(promotion_alignment_threshold(0), 1.0);
assert_eq!(promotion_alignment_threshold(1), 1.0);
let rank_two = promotion_alignment_threshold(2);
let rank_four = promotion_alignment_threshold(4);
assert!(rank_two.is_finite() && (0.0..=1.0).contains(&rank_two));
assert!(rank_four.is_finite() && (0.0..=1.0).contains(&rank_four));
assert!(rank_four < rank_two);
}
}