use super::*;
use gam_math::special::bessel_i0_log_and_ratio;
use gam_solve::rho_optimizer::{
FixedPointCertificateEval, FixedPointCoordinateCertificate, OuterResult,
};
pub(crate) fn reconstruction_explained_variance(
target: ArrayView2<'_, f64>,
fitted: ArrayView2<'_, f64>,
) -> Option<f64> {
if target.dim() != fitted.dim() {
return None;
}
let (n, p) = target.dim();
if n == 0 || p == 0 {
return None;
}
let mut means = vec![0.0_f64; p];
for col in 0..p {
let mut acc = 0.0;
for row in 0..n {
acc += target[[row, col]];
}
means[col] = acc / n as f64;
}
let mut ssr = 0.0_f64;
let mut sst = 0.0_f64;
for row in 0..n {
for col in 0..p {
let residual = target[[row, col]] - fitted[[row, col]];
ssr += residual * residual;
let centered = target[[row, col]] - means[col];
sst += centered * centered;
}
}
if ssr.is_finite() && sst.is_finite() && sst > f64::MIN_POSITIVE {
Some(1.0 - ssr / sst)
} else {
None
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct AmortizedWarmStartTelemetry {
pub attempts: usize,
pub warm_started_evals: usize,
pub zero_certified_evals: usize,
pub failed_attempts: usize,
pub total_rows_warm_started: usize,
}
#[derive(Debug, Clone)]
pub(crate) struct OuterTerminationLedger {
evals: u64,
last_improvement_eval: u64,
best_cost: Option<f64>,
wall_start: std::time::Instant,
}
impl OuterTerminationLedger {
pub(crate) fn new() -> Self {
Self {
evals: 0,
last_improvement_eval: 0,
best_cost: None,
wall_start: std::time::Instant::now(),
}
}
pub(crate) fn record(&mut self, cost: f64, gradient_norm: Option<f64>) -> bool {
self.evals += 1;
let gradient_field = match gradient_norm {
Some(norm) => format!("{norm:.6e}"),
None => "na".to_string(),
};
if !cost.is_finite() {
log::info!(
"[SAE/outer] eval={} criterion={:.9e} grad={} best={:.9e} improved=false",
self.evals,
cost,
gradient_field,
self.best_cost.unwrap_or(f64::NAN),
);
return false;
}
let improved = match self.best_cost {
None => true,
Some(best) => cost < best - SAE_FINAL_EV_DEGRADATION_TOL * (1.0 + best.abs()),
};
if improved {
self.best_cost = Some(match self.best_cost {
Some(best) => best.min(cost),
None => cost,
});
self.last_improvement_eval = self.evals;
}
log::info!(
"[SAE/outer] eval={} criterion={:.9e} grad={} best={:.9e} improved={improved}",
self.evals,
cost,
gradient_field,
self.best_cost.unwrap_or(cost),
);
improved
}
pub(crate) fn seed_from_checkpoint(
&mut self,
evals: u64,
last_improvement_eval: u64,
best_cost: Option<f64>,
) {
self.evals = evals;
self.last_improvement_eval = last_improvement_eval.min(evals);
self.best_cost = best_cost.filter(|c| c.is_finite());
}
pub(crate) fn checkpoint_counters(&self) -> (u64, u64, Option<f64>) {
(self.evals, self.last_improvement_eval, self.best_cost)
}
pub(crate) fn reset_improvement_baseline(&mut self) {
self.last_improvement_eval = self.evals;
}
pub(crate) fn report(&self, verdict: SaeOuterVerdict) -> SaeOuterTermination {
SaeOuterTermination {
verdict,
evals: self.evals,
evals_since_improvement: self.evals.saturating_sub(self.last_improvement_eval),
wall: self.wall_start.elapsed(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum SaeOuterVerdict {
Search(OuterConvergedVia),
FixedRho,
Audited(OuterConvergedVia),
}
impl SaeOuterVerdict {
pub fn as_str(&self) -> &'static str {
match self {
Self::Search(via) => via.as_str(),
Self::FixedRho => "fixed_rho",
Self::Audited(_) => "audited_stationary",
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct SaeOuterTermination {
pub verdict: SaeOuterVerdict,
pub evals: u64,
pub evals_since_improvement: u64,
pub wall: std::time::Duration,
}
#[derive(Debug)]
pub struct SaeIntoFittedResult {
pub term: SaeManifoldTerm,
pub rho: SaeManifoldRho,
pub loss: SaeManifoldLoss,
pub penalized_quasi_laplace_criterion: f64,
pub charts_canonicalized: bool,
pub termination: SaeOuterTermination,
}
pub(crate) struct SaeVanishedStageState {
pub term: SaeManifoldTerm,
pub rho: SaeManifoldRho,
pub atoms: VanishedAtoms,
}
impl SaeIntoFittedResult {
pub fn invalidates_pre_final_shape_uncertainty(&self) -> bool {
self.charts_canonicalized
}
}
impl AmortizedWarmStartTelemetry {
pub(crate) fn record(&mut self, outcome: &Result<usize, String>) {
self.attempts += 1;
match outcome {
Ok(0) => self.zero_certified_evals += 1,
Ok(rows) => {
self.warm_started_evals += 1;
self.total_rows_warm_started += rows;
}
Err(_) => self.failed_attempts += 1,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum ProbeRefusalKind {
InnerNotConverged,
NonPdPerRow,
NonPdSchur,
AllZeroGatedDesign,
TotalCoCollapse,
}
impl ProbeRefusalKind {
pub(crate) const ALL: [Self; 5] = [
Self::InnerNotConverged,
Self::NonPdPerRow,
Self::NonPdSchur,
Self::AllZeroGatedDesign,
Self::TotalCoCollapse,
];
pub(crate) fn inner_not_converged_marker() -> &'static str {
"inner solve did not converge at fixed ρ"
}
pub(crate) fn non_pd_per_row_marker() -> &'static str {
"non-PD per-row H_tt block"
}
pub(crate) fn all_zero_gated_design_marker() -> &'static str {
"gated off at every row (all-zero gated design)"
}
pub(crate) fn total_co_collapse_marker() -> &'static str {
"did not escape total co-collapse"
}
pub(crate) fn classify(err: &str) -> Option<Self> {
if err.contains(Self::inner_not_converged_marker()) {
return Some(Self::InnerNotConverged);
}
if err.contains(Self::non_pd_per_row_marker()) {
return Some(Self::NonPdPerRow);
}
if ArrowSchurError::rendered_is_non_pd_schur_complement(err) {
return Some(Self::NonPdSchur);
}
if err.contains(Self::all_zero_gated_design_marker()) {
return Some(Self::AllZeroGatedDesign);
}
if err.contains(Self::total_co_collapse_marker()) {
return Some(Self::TotalCoCollapse);
}
None
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct OuterProbeTelemetry {
pub criterion_calls: usize,
pub infeasible_non_pd_per_row: usize,
pub infeasible_schur: usize,
pub infeasible_inner_not_converged: usize,
pub infeasible_all_zero_gated_design: usize,
pub infeasible_total_co_collapse: usize,
pub infeasible_criterion_evals: usize,
pub basin_envelope_evals: usize,
pub basin_admissions: usize,
pub basin_envelope_rescues: usize,
pub basin_max_members: usize,
pub basin_member_capacity: usize,
pub reactive_scalar_installs: usize,
pub reactive_target_restores: usize,
}
impl OuterProbeTelemetry {
fn record_refusal_kind(&mut self, err: &str) {
let Some(kind) = ProbeRefusalKind::classify(err) else {
return;
};
*self.counter_mut(kind) += 1;
}
fn counter_mut(&mut self, kind: ProbeRefusalKind) -> &mut usize {
match kind {
ProbeRefusalKind::InnerNotConverged => &mut self.infeasible_inner_not_converged,
ProbeRefusalKind::NonPdPerRow => &mut self.infeasible_non_pd_per_row,
ProbeRefusalKind::NonPdSchur => &mut self.infeasible_schur,
ProbeRefusalKind::AllZeroGatedDesign => &mut self.infeasible_all_zero_gated_design,
ProbeRefusalKind::TotalCoCollapse => &mut self.infeasible_total_co_collapse,
}
}
pub fn infeasible_total(&self) -> usize {
ProbeRefusalKind::ALL
.iter()
.map(|kind| self.counter_of(*kind))
.sum()
}
fn counter_of(&self, kind: ProbeRefusalKind) -> usize {
match kind {
ProbeRefusalKind::InnerNotConverged => self.infeasible_inner_not_converged,
ProbeRefusalKind::NonPdPerRow => self.infeasible_non_pd_per_row,
ProbeRefusalKind::NonPdSchur => self.infeasible_schur,
ProbeRefusalKind::AllZeroGatedDesign => self.infeasible_all_zero_gated_design,
ProbeRefusalKind::TotalCoCollapse => self.infeasible_total_co_collapse,
}
}
}
struct ProbeConvergedHandoff {
rho_flat: Array1<f64>,
term: SaeManifoldTerm,
}
#[derive(Clone)]
struct CrosscoderBlockPricing {
p_x: usize,
block_dims: Vec<usize>,
pristine_blocks: Array2<f64>,
last_log_lambda: Vec<f64>,
}
struct ReactiveWaypointCheckpoint {
term: SaeManifoldTerm,
target: Array2<f64>,
registry_isometry_weights: Vec<f64>,
current_rho: SaeManifoldRho,
last_loss: Option<SaeManifoldLoss>,
terminal_penalized_quasi_laplace_criterion: Option<f64>,
seeded_beta: Option<Array1<f64>>,
probe_converged_handoff: Option<ProbeConvergedHandoff>,
basin_bundle: BasinBundle<SaeManifoldTerm>,
termination: OuterTerminationLedger,
fit_verdict: Option<SaeOuterVerdict>,
crosscoder_blocks: Option<CrosscoderBlockPricing>,
}
struct MatrixFreeOuterArtifacts {
system: ArrowSchurSystem,
exact_a_cache: ArrowFactorCache,
logdet_derivative_bundle: RationalLogdetDerivativeBundle,
efs_inverse_probe_bundle: Option<(Vec<Array1<f64>>, Vec<Array1<f64>>)>,
}
pub(crate) struct OuterCriterionEvaluation {
pub(crate) cost: f64,
loss: SaeManifoldLoss,
cache: ArrowFactorCache,
matrix_free: Option<MatrixFreeOuterArtifacts>,
}
pub struct SaeManifoldOuterObjective {
pub(crate) term: SaeManifoldTerm,
pub(crate) baseline_term: SaeManifoldTerm,
pub(crate) target: Array2<f64>,
pub(crate) registry: Option<AnalyticPenaltyRegistry>,
baseline_isometry_weights: Vec<f64>,
pub(crate) current_rho: SaeManifoldRho,
pub(crate) baseline_rho: SaeManifoldRho,
pub(crate) inner_max_iter: usize,
pub(crate) learning_rate: f64,
pub(crate) ridge_ext_coord: f64,
pub(crate) ridge_beta: f64,
pub(crate) last_loss: Option<SaeManifoldLoss>,
pub(crate) terminal_penalized_quasi_laplace_criterion: Option<f64>,
pub(crate) seeded_beta: Option<Array1<f64>>,
pub(crate) warm_start_telemetry: AmortizedWarmStartTelemetry,
pub(crate) probe_telemetry: OuterProbeTelemetry,
pub(crate) cancel_flag: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
probe_converged_handoff: Option<ProbeConvergedHandoff>,
surrogate_lane: Option<SurrogateLaneState>,
basin_bundle: BasinBundle<SaeManifoldTerm>,
pub(crate) termination: OuterTerminationLedger,
fit_verdict: Option<SaeOuterVerdict>,
audit_installed_state: bool,
pub(crate) checkpoint_fingerprint: super::checkpoint::SaeCheckpointFingerprint,
pub(crate) checkpoint_path: std::path::PathBuf,
crosscoder_blocks: Option<CrosscoderBlockPricing>,
reactive_waypoint_checkpoint: Option<ReactiveWaypointCheckpoint>,
}
fn basin_bundle_member_capacity(term: &SaeManifoldTerm) -> usize {
let host_available = term.host_available_bytes;
let host_budget = super::sae_host_in_core_budget_from_available(host_available);
let total_basis: usize = term.atoms.iter().map(SaeManifoldAtom::basis_size).sum();
let d_max = term
.atoms
.iter()
.map(SaeManifoldAtom::latent_dim)
.max()
.unwrap_or(0);
let border_dim = if term.any_frame_active() {
term.factored_border_dim()
} else {
term.beta_dim()
};
let plan = super::sae_streaming_plan_from_budget(
term.n_obs(),
total_basis,
term.k_atoms(),
d_max,
border_dim,
host_budget,
super::SAE_CPU_L2_CACHE_BYTES * super::SAE_CHUNK_CACHE_MULTIPLE,
host_available,
);
if !plan.direct_logdet_admitted() {
return 0;
}
let bytes_per_saved_state = plan
.estimated_direct_peak_bytes
.max(plan.estimated_full_batch_bytes)
.max(std::mem::size_of::<SaeManifoldTerm>());
host_budget.saturating_sub(plan.estimated_direct_peak_bytes) / bytes_per_saved_state
}
pub(crate) fn sae_outer_gradient_capability() -> Derivative {
Derivative::Analytic
}
pub(crate) fn assignment_strength_gradient_coordinate(rho: &SaeManifoldRho) -> Option<usize> {
rho.sparse_flat_index()
}
const SAE_SURROGATE_LANE_QUADRATURE_REL_TOL: f64 = 1.0e-8;
const SAE_SURROGATE_LANE_POWER_ITERS: usize = 40;
const SAE_SURROGATE_LANE_CG_REL_TOL: f64 = 1.0e-8;
const SAE_SURROGATE_LANE_CG_MAX_ITERS: usize = 20_000;
const SAE_SURROGATE_LANE_DEFLATION_MAX_RANK: usize = 128;
const SAE_SURROGATE_LANE_DEFLATION_SUBSPACE_ITERS: usize = 4;
pub(crate) fn sae_surrogate_lane_config() -> SurrogateLaneConfig {
SurrogateLaneConfig {
num_probes: SCHUR_SLQ_LOGDET_PROBES,
seed: SCHUR_SLQ_LOGDET_SEED,
rel_tol: SAE_SURROGATE_LANE_QUADRATURE_REL_TOL,
power_iters: SAE_SURROGATE_LANE_POWER_ITERS,
cg_rel_tol: SAE_SURROGATE_LANE_CG_REL_TOL,
cg_max_iters: SAE_SURROGATE_LANE_CG_MAX_ITERS,
deflation_max_rank: SAE_SURROGATE_LANE_DEFLATION_MAX_RANK,
deflation_subspace_iters: SAE_SURROGATE_LANE_DEFLATION_SUBSPACE_ITERS,
deflation_target_std_err_rel: 0.1 * SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL,
}
}
impl SaeManifoldOuterObjective {
fn curvature_seed(term: &SaeManifoldTerm) -> Vec<(usize, f64)> {
term.atoms
.iter()
.enumerate()
.filter_map(|(atom, value)| {
value
.geometry_plan()
.and_then(SaeAtomGeometryPlan::constant_curvature)
.map(|kappa| (atom, kappa))
})
.collect()
}
fn apply_curvature_state(&mut self, rho: &SaeManifoldRho) -> Result<(), String> {
let mut prepared = Vec::with_capacity(rho.kappa.len());
for (&atom_index, &kappa) in rho.kappa_atoms.iter().zip(rho.kappa.iter()) {
let atom = self.term.atoms.get(atom_index).ok_or_else(|| {
format!(
"curvature coordinate names atom {atom_index}, outside K={}",
self.term.atoms.len()
)
})?;
let already_installed = atom
.geometry_plan()
.and_then(SaeAtomGeometryPlan::constant_curvature)
.is_some_and(|current| current.to_bits() == kappa.to_bits())
&& atom.smooth_penalty_kappa_derivative().is_some();
if !already_installed {
prepared.push((atom_index, atom.prepare_constant_curvature(kappa)?));
}
}
for (atom_index, state) in prepared {
self.term.atoms[atom_index].commit_prepared_constant_curvature(state);
}
Ok(())
}
fn curvature_domain_bounds(&self) -> Result<Vec<(usize, f64, f64)>, EstimationError> {
let mut out = Vec::with_capacity(self.baseline_rho.kappa.len());
for &atom_index in &self.baseline_rho.kappa_atoms {
let flat = self
.baseline_rho
.kappa_flat_index(atom_index)
.ok_or_else(|| {
EstimationError::InvalidInput(format!(
"curvature atom {atom_index} has no flat outer coordinate"
))
})?;
let atom = self.baseline_term.atoms.get(atom_index).ok_or_else(|| {
EstimationError::InvalidInput(format!(
"curvature coordinate names atom {atom_index}, outside K={}",
self.baseline_term.atoms.len()
))
})?;
let (lower, upper) = atom
.geometry_plan()
.ok_or_else(|| {
EstimationError::InvalidInput(format!(
"curvature atom {atom_index} has no typed geometry plan"
))
})?
.constant_curvature_domain()
.map_err(EstimationError::InvalidInput)?
.ok_or_else(|| {
EstimationError::InvalidInput(format!(
"atom {atom_index} owns a curvature coordinate but its metric is not constant-curvature"
))
})?;
out.push((flat, lower, upper));
}
Ok(out)
}
pub(crate) fn current_rho_flat(&self) -> Array1<f64> {
self.current_rho.to_flat()
}
pub(crate) fn vanished_stage_state_at(
&self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<Option<SaeVanishedStageState>, String> {
let rho = self.baseline_rho.from_flat(rho_flat)?;
let mut term = self.term.clone();
let evaluated = if term.streaming_plan()?.direct_logdet_admitted() {
term.penalized_quasi_laplace_criterion_with_cache(
self.target.view(),
&rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
)
} else {
term.penalized_quasi_laplace_criterion_streaming_exact_with_cache(
self.target.view(),
&rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
)
};
let atoms = match evaluated {
Ok(_) => return Ok(None),
Err(err @ SaeCriterionError::IndefiniteObservedInformation { .. }) => {
return Err(err.to_string());
}
Err(SaeCriterionError::Numerical(message)) => return Err(message),
Err(SaeCriterionError::VanishedAtoms(atoms)) => atoms,
};
let vanished = atoms.as_btree_set();
for atlas in term.chart_atlases() {
let removed = atlas
.charts()
.iter()
.filter(|chart| vanished.contains(chart))
.count();
if removed > 0 && removed < atlas.charts().len() {
return Err(format!(
"vanished-atom boundary would partially delete live atlas {:?}; \
chart-atlas disappearance must be adjudicated at semantic-atlas granularity",
atlas.charts()
));
}
}
Ok(Some(SaeVanishedStageState { term, rho, atoms }))
}
pub fn new(
mut term: SaeManifoldTerm,
target: Array2<f64>,
registry: Option<AnalyticPenaltyRegistry>,
init_rho: SaeManifoldRho,
inner_max_iter: usize,
learning_rate: f64,
ridge_ext_coord: f64,
ridge_beta: f64,
) -> Self {
let init_rho = init_rho
.for_assignment(term.assignment.mode)
.with_curvature(Self::curvature_seed(&term));
term.expected_criterion_gauge_deflated_directions = None;
term.criterion_gauge_deflation_reanchors = 0;
term.criterion_gauge_deflation_last_delta_sign = 0;
term.dictionary_cocollapse_reseeds = 0;
term.best_cocollapse_incumbent = None;
term.structural_cocollapse_reseeds = 0;
let baseline_term = term.clone();
let baseline_rho = init_rho.clone();
let baseline_isometry_weights = registry
.as_ref()
.map(AnalyticPenaltyRegistry::isometry_scalar_weights)
.unwrap_or_default();
let term_k_atoms = term.k_atoms();
let basin_member_capacity = basin_bundle_member_capacity(&term);
let checkpoint_fingerprint =
super::checkpoint::SaeCheckpointFingerprint::of_target(target.view(), term_k_atoms);
let checkpoint_path =
super::checkpoint::SaeFitCheckpoint::default_store_path(&checkpoint_fingerprint);
Self {
term,
baseline_term,
target,
registry,
baseline_isometry_weights,
current_rho: init_rho,
baseline_rho,
inner_max_iter,
learning_rate,
ridge_ext_coord,
ridge_beta,
last_loss: None,
terminal_penalized_quasi_laplace_criterion: None,
seeded_beta: None,
warm_start_telemetry: AmortizedWarmStartTelemetry::default(),
probe_telemetry: OuterProbeTelemetry::default(),
cancel_flag: None,
probe_converged_handoff: None,
surrogate_lane: Some(SurrogateLaneState::new(sae_surrogate_lane_config())),
basin_bundle: BasinBundle::new(basin_member_capacity),
termination: OuterTerminationLedger::new(),
fit_verdict: None,
audit_installed_state: false,
checkpoint_fingerprint,
checkpoint_path,
crosscoder_blocks: None,
reactive_waypoint_checkpoint: None,
}
}
pub(crate) fn evaluate_outer_criterion_route(
&mut self,
rho: &SaeManifoldRho,
direct_logdet_admitted: bool,
need_efs_inverse_probes: bool,
) -> Result<OuterCriterionEvaluation, SaeCriterionError> {
if direct_logdet_admitted {
let (cost, loss, cache) = self.term.penalized_quasi_laplace_criterion_with_cache(
self.target.view(),
rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
)?;
return Ok(OuterCriterionEvaluation {
cost,
loss,
cache,
matrix_free: None,
});
}
let lane = self.surrogate_lane.as_mut().ok_or_else(|| {
SaeCriterionError::Numerical(
"streaming outer evaluation requires the frozen rational-logdet surrogate lane"
.to_string(),
)
})?;
let evaluated = self
.term
.penalized_quasi_laplace_streaming_outer_evaluation(
self.target.view(),
rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
lane,
need_efs_inverse_probes,
)?;
Ok(OuterCriterionEvaluation {
cost: evaluated.cost,
loss: evaluated.loss,
cache: evaluated.cache,
matrix_free: Some(MatrixFreeOuterArtifacts {
system: evaluated.system,
exact_a_cache: evaluated.exact_a_cache,
logdet_derivative_bundle: evaluated.logdet_derivative_bundle,
efs_inverse_probe_bundle: evaluated.efs_inverse_probe_bundle,
}),
})
}
pub(crate) fn analytic_gradient_for_outer_evaluation(
&self,
rho: &SaeManifoldRho,
evaluation: &OuterCriterionEvaluation,
) -> Result<Array1<f64>, OuterGradientError> {
let components = if let Some(matrix_free) = evaluation.matrix_free.as_ref() {
let derivative_vectors = &matrix_free.logdet_derivative_bundle.vectors;
let solver = DeflatedArrowSolver::plain(&evaluation.cache);
self.term
.analytic_outer_rho_gradient_components_with_bundle(
self.target.view(),
rho,
&evaluation.loss,
&evaluation.cache,
&solver,
Some(BundleEvidenceGeometry {
operator: EvidenceOperator::ExactObservedInformation,
cache: &matrix_free.exact_a_cache,
probes: derivative_vectors,
sinv: derivative_vectors,
}),
Some(&matrix_free.system),
)?
} else {
let lambda_smooth = rho
.lambda_smooth_vec()
.map_err(OuterGradientError::internal)?;
let solver = self
.term
.outer_gradient_arrow_solver(&evaluation.cache, &lambda_smooth)?;
self.term
.analytic_outer_rho_gradient_components_with_bundle(
self.target.view(),
rho,
&evaluation.loss,
&evaluation.cache,
&solver,
None,
None,
)?
};
let mut gradient = components.gradient();
if let Some(block_grad) = self
.block_log_lambda_gradient(rho)
.map_err(OuterGradientError::internal)?
{
let trailing = rho.kappa.len();
let tail = gradient.len() - trailing - block_grad.len();
for (block, value) in block_grad.into_iter().enumerate() {
gradient[tail + block] += value;
}
}
Ok(gradient)
}
pub fn with_crosscoder_blocks(
mut self,
p_x: usize,
block_dims: Vec<usize>,
) -> Result<Self, String> {
if p_x == 0 {
return Err("with_crosscoder_blocks: anchor width p_x must be non-zero".to_string());
}
if block_dims.is_empty() {
return Err(
"with_crosscoder_blocks: block_dims is empty — a plain SAE must not install \
crosscoder pricing (leave crosscoder_blocks = None)"
.to_string(),
);
}
let block_total: usize = block_dims.iter().sum();
let p_tot = self.target.ncols();
if p_x + block_total != p_tot {
return Err(format!(
"with_crosscoder_blocks: p_x ({p_x}) + Σ block_dims ({block_total}) = {} \
must equal the stacked target width p̃ = {p_tot}",
p_x + block_total
));
}
let template_blocks = self.baseline_rho.log_lambda_block.len();
if block_dims.len() != template_blocks {
return Err(format!(
"with_crosscoder_blocks: block_dims length ({}) must match the ρ template's \
log_lambda_block count ({template_blocks})",
block_dims.len()
));
}
if self.term.row_loss_weights.is_some() {
return Err(
"with_crosscoder_blocks: the outer row-subsample (row_loss_weights, #991) is \
engaged; block pricing needs the pristine copy restricted to the sampled rows \
and the Jacobian n set to the effective sample size — deferred (stage 1)"
.to_string(),
);
}
let (budget_bytes, _) = super::sae_host_in_core_budget_bytes();
crate::front_door::admit_crosscoder_border(
self.term.factored_border_dim(),
self.term.beta_dim(),
budget_bytes,
)?;
let pristine_blocks = self.target.slice(s![.., p_x..]).to_owned();
self.term.crosscoder_pricing_spans = Some((p_x, block_dims.clone()));
self.crosscoder_blocks = Some(CrosscoderBlockPricing {
p_x,
last_log_lambda: vec![0.0; block_dims.len()],
block_dims,
pristine_blocks,
});
Ok(self)
}
fn apply_block_scaling(&mut self, rho: &SaeManifoldRho) -> Result<(), String> {
self.term.assignment.validate_rho_domain(rho)?;
self.apply_curvature_state(rho)?;
let Self {
target,
crosscoder_blocks: Some(blocks),
..
} = self
else {
return Ok(());
};
if rho.log_lambda_block.len() != blocks.block_dims.len() {
return Err(format!(
"crosscoder block log-strength count {} != pricing block count {}",
rho.log_lambda_block.len(),
blocks.block_dims.len()
));
}
let mut moved: Vec<(usize, usize, f64)> = Vec::new(); let mut pristine_off = 0usize;
for l in 0..blocks.block_dims.len() {
let p_l = blocks.block_dims[l];
let new_ll = rho.log_lambda_block[l];
if new_ll != blocks.last_log_lambda[l] {
moved.push((pristine_off, p_l, (0.5 * new_ll).exp()));
blocks.last_log_lambda[l] = new_ll;
}
pristine_off += p_l;
}
if moved.is_empty() {
return Ok(());
}
let p_x = blocks.p_x;
let pristine = &blocks.pristine_blocks;
use rayon::prelude::*;
target
.axis_iter_mut(ndarray::Axis(0))
.into_par_iter()
.zip(pristine.axis_iter(ndarray::Axis(0)).into_par_iter())
.for_each(|(mut dst_row, src_row)| {
let src = src_row
.to_slice()
.expect("pristine block rows are contiguous");
let dst = dst_row
.as_slice_mut()
.expect("stacked target rows are contiguous");
for &(off, p_l, sqrt_lambda) in &moved {
let dst_span = &mut dst[p_x + off..p_x + off + p_l];
let src_span = &src[off..off + p_l];
for (d, &s) in dst_span.iter_mut().zip(src_span) {
*d = s * sqrt_lambda;
}
}
});
Ok(())
}
fn block_jacobian(&self, rho: &SaeManifoldRho) -> f64 {
let Some(blocks) = self.crosscoder_blocks.as_ref() else {
return 0.0;
};
let n = self.target.nrows() as f64;
blocks
.block_dims
.iter()
.zip(rho.log_lambda_block.iter())
.map(|(&p_l, &log_lambda)| -(n * p_l as f64 / 2.0) * log_lambda)
.sum()
}
fn block_scaled_rss(&self, rho: &SaeManifoldRho) -> Result<Option<Vec<f64>>, String> {
let Some(blocks) = self.crosscoder_blocks.as_ref() else {
return Ok(None);
};
let residual = self.term.reconstruction_residual(self.target.view(), rho)?;
let mut out = Vec::with_capacity(blocks.block_dims.len());
let mut off = blocks.p_x;
for &p_l in &blocks.block_dims {
let mut rss = 0.0_f64;
for row in residual.rows() {
for j in off..off + p_l {
let r = row[j];
rss += r * r;
}
}
out.push(rss);
off += p_l;
}
Ok(Some(out))
}
fn block_log_lambda_gradient(&self, rho: &SaeManifoldRho) -> Result<Option<Vec<f64>>, String> {
let Some(scaled_rss) = self.block_scaled_rss(rho)? else {
return Ok(None);
};
let blocks = self
.crosscoder_blocks
.as_ref()
.expect("block_scaled_rss returned Some ⇒ crosscoder pricing is installed");
let n = self.target.nrows() as f64;
Ok(Some(
blocks
.block_dims
.iter()
.zip(scaled_rss.iter())
.map(|(&p_l, &r_tilde)| 0.5 * r_tilde - 0.5 * n * p_l as f64)
.collect(),
))
}
pub(crate) fn bank_checkpoint(&self, rho_flat: &Array1<f64>) {
let (evals, last_improvement_eval, best_cost) = self.termination.checkpoint_counters();
let rho_owned = rho_flat.to_vec();
let incumbent_ev = self
.term
.dictionary_reconstruction_ev(self.target.view(), &self.current_rho)
.ok()
.filter(|ev| ev.is_finite())
.unwrap_or(-1.0);
let ckpt = super::checkpoint::SaeFitCheckpoint::capture(
&self.term,
&self.checkpoint_fingerprint,
&rho_owned,
super::checkpoint::SaeCheckpointLedger {
evals,
last_improvement_eval,
best_cost,
},
incumbent_ev,
);
if let Some(dir) = self.checkpoint_path.parent()
&& let Err(e) = std::fs::create_dir_all(dir)
{
log::warn!("SAE fit checkpoint: create dir {}: {e}", dir.display());
return;
}
if let Err(e) = ckpt.save_atomic(&self.checkpoint_path) {
log::warn!("SAE fit checkpoint: {e}");
}
}
pub fn try_resume_from_checkpoint(
&mut self,
expected_rho_len: usize,
) -> Result<Option<Vec<f64>>, String> {
self.fit_verdict = None;
if !self.checkpoint_path.exists() {
return Ok(None);
}
let ckpt = match super::checkpoint::SaeFitCheckpoint::load(&self.checkpoint_path) {
Ok(c) => c,
Err(e) => {
log::warn!("SAE fit checkpoint resume: {e}; fitting cold");
return Ok(None);
}
};
if let Err(e) = ckpt.verify_compatible(&self.checkpoint_fingerprint, expected_rho_len) {
log::warn!("SAE fit checkpoint resume: {e}; fitting cold");
return Ok(None);
}
if let Err(e) = self
.baseline_rho
.from_flat(ArrayView1::from(ckpt.rho_flat.as_slice()))
{
return Err(format!(
"SAE fit checkpoint resume refused invalid rho payload: {e}"
));
}
let install_result = ckpt.install_into(&mut self.term);
if install_result.is_ok()
&& let Err(e) = ckpt.install_into(&mut self.baseline_term)
{
log::warn!("SAE fit checkpoint resume (baseline): {e}");
}
if let Err(e) = install_result {
log::warn!("SAE fit checkpoint resume: {e}; fitting cold");
return Ok(None);
}
self.termination.seed_from_checkpoint(
ckpt.ledger.evals,
ckpt.ledger.last_improvement_eval,
ckpt.ledger.best_cost,
);
log::warn!(
"SAE fit checkpoint resume: installed banked incumbent from {} \
(evals {}, best cost {:?}); the resumed search must still converge on its own",
self.checkpoint_path.display(),
ckpt.ledger.evals,
ckpt.ledger.best_cost,
);
Ok(Some(ckpt.rho_flat))
}
pub fn remove_checkpoint(&self) {
if self.checkpoint_path.exists()
&& let Err(e) = std::fs::remove_file(&self.checkpoint_path)
{
log::warn!(
"SAE fit checkpoint: remove {}: {e}",
self.checkpoint_path.display()
);
}
}
pub fn set_cancel_flag(&mut self, flag: std::sync::Arc<std::sync::atomic::AtomicBool>) {
self.cancel_flag = Some(flag);
}
fn check_cancelled(&self) -> Result<(), EstimationError> {
if let Some(flag) = &self.cancel_flag {
if flag.load(std::sync::atomic::Ordering::Relaxed) {
return Err(EstimationError::RemlOptimizationFailed(
"SAE fit cancelled by host (Python interrupt)".to_string(),
));
}
}
Ok(())
}
pub fn probe_telemetry(&self) -> OuterProbeTelemetry {
self.probe_telemetry
}
fn record_warm_start(&mut self, outcome: Result<usize, String>) -> Result<(), String> {
self.warm_start_telemetry.record(&outcome);
outcome?;
Ok(())
}
pub fn certify_outer_result(&mut self, result: &OuterResult) -> Result<(), String> {
self.fit_verdict = None;
self.terminal_penalized_quasi_laplace_criterion = None;
if !result.converged() {
return Err("outer result is not converged".to_string());
}
let via = result
.converged_via()
.ok_or_else(|| "converged outer result is missing converged_via".to_string())?;
let certificate = result.criterion_certificate.as_ref().ok_or_else(|| {
"converged outer result is missing its analytic criterion certificate".to_string()
})?;
if !certificate.certifies() {
return Err(format!(
"outer criterion certificate does not certify the installed state: {}",
certificate.summary()
));
}
if self.last_loss.is_none() {
return Err("outer result has no installed converged inner loss".to_string());
}
let installed_rho = self.current_rho.to_flat();
let rho_matches = installed_rho.len() == result.rho.len()
&& installed_rho
.iter()
.zip(result.rho.iter())
.all(|(installed, certified)| installed.to_bits() == certified.to_bits());
if !rho_matches {
return Err(format!(
"outer result rho does not match the installed state (certified={:?}, installed={:?})",
result.rho, installed_rho
));
}
if !result.final_value.is_finite() {
return Err("converged outer result has a non-finite final criterion value".into());
}
self.terminal_penalized_quasi_laplace_criterion = Some(result.final_value);
self.fit_verdict = Some(SaeOuterVerdict::Search(via));
Ok(())
}
pub(crate) fn for_installed_state_audit(mut self) -> Self {
self.audit_installed_state = true;
self.inner_max_iter = 0;
self
}
fn record_search_criterion(&mut self, cost: f64, gradient_norm: Option<f64>) -> bool {
!self.audit_installed_state && self.termination.record(cost, gradient_norm)
}
pub(crate) fn certify_installed_state_audit(
&mut self,
result: &OuterResult,
) -> Result<(), String> {
self.fit_verdict = None;
self.terminal_penalized_quasi_laplace_criterion = None;
if !self.audit_installed_state {
return Err("installed-state audit was not enabled on this objective".to_string());
}
if result.iterations != 0 || !result.converged() {
return Err("installed-state audit result is not a zero-step convergence".to_string());
}
let via = result
.converged_via()
.ok_or_else(|| "installed-state audit is missing converged_via".to_string())?;
let certificate = result.criterion_certificate.as_ref().ok_or_else(|| {
"installed-state audit is missing its analytic criterion certificate".to_string()
})?;
if !certificate.certifies() {
return Err(format!(
"installed-state outer certificate does not certify: {}",
certificate.summary()
));
}
if self.last_loss.is_none() {
return Err("installed-state audit has no evaluated inner loss".to_string());
}
let installed_rho = self.current_rho.to_flat();
if installed_rho.len() != result.rho.len()
|| installed_rho
.iter()
.zip(result.rho.iter())
.any(|(installed, certified)| installed.to_bits() != certified.to_bits())
{
return Err("installed-state audit rho does not match the evaluated state".to_string());
}
if !result.final_value.is_finite() {
return Err("installed-state audit produced a non-finite criterion".to_string());
}
self.terminal_penalized_quasi_laplace_criterion = Some(result.final_value);
self.fit_verdict = Some(SaeOuterVerdict::Audited(via));
Ok(())
}
pub fn into_fitted(self) -> Result<SaeIntoFittedResult, String> {
let verdict = self.fit_verdict.ok_or_else(|| {
"SaeManifoldOuterObjective::into_fitted: installed state is not explicitly certified; \
run fit_at_fixed_rho or certify a converged OuterResult before minting a fit"
.to_string()
})?;
let termination_report = self.termination.report(verdict);
let Self {
term,
target,
registry,
current_rho,
last_loss,
terminal_penalized_quasi_laplace_criterion,
..
} = self;
let mut fitted_rho = current_rho;
let mut fitted = term;
if last_loss.is_none() {
return Err(
"SaeManifoldOuterObjective::into_fitted: certified state has no converged inner loss"
.to_string(),
);
}
let penalized_quasi_laplace_criterion = terminal_penalized_quasi_laplace_criterion.ok_or_else(|| {
"SaeManifoldOuterObjective::into_fitted: terminal state has no penalized quasi-Laplace criterion value"
.to_string()
})?;
let pre_canonical_flags = fitted
.atoms
.iter()
.map(|atom| atom.chart_canonicalized)
.collect::<Vec<_>>();
if let Err(err) =
fitted.canonicalize_charts_post_fit(target.view(), &fitted_rho, registry.as_ref())
{
log::debug!("into_fitted: chart canonicalization refused: {err}");
}
let charts_canonicalized = fitted
.atoms
.iter()
.zip(pre_canonical_flags.iter())
.any(|(atom, before)| atom.chart_canonicalized != *before);
if fitted
.assignment
.persist_resolved_ordered_beta_bernoulli_alpha(&fitted_rho)
{
fitted_rho.log_lambda_sparse = 0.0;
}
let fitted_loss = fitted.loss(target.view(), &fitted_rho)?;
let termination = termination_report;
log::warn!(
"[#2235] outer search concluded: {} evals ({} since last improvement, wall {:.1?})",
termination.evals,
termination.evals_since_improvement,
termination.wall
);
Ok(SaeIntoFittedResult {
term: fitted,
rho: fitted_rho,
loss: fitted_loss,
penalized_quasi_laplace_criterion,
charts_canonicalized,
termination,
})
}
pub fn curvature_walk_report(&self) -> Option<&CurvatureWalkReport> {
self.term.curvature_walk_report()
}
pub fn decoder_shape_uncertainty(&mut self) -> Result<SaeShapeUncertainty, String> {
self.probe_converged_handoff = None;
self.basin_bundle.clear();
let rho = self.current_rho.clone();
let plan = self.term.streaming_plan()?.admitted_or_error(
self.term.n_obs(),
self.term.output_dim(),
self.term.k_atoms(),
)?;
if !plan.direct_logdet_admitted() {
let loss = self.term.loss(self.target.view(), &rho)?;
let n_scalar = (self.term.n_obs().saturating_mul(self.term.output_dim())).max(1) as f64;
let dispersion = (2.0 * loss.data_fit / n_scalar).max(f64::MIN_POSITIVE);
return Ok(self.term.unavailable_shape_uncertainty(dispersion));
}
let saved_term = self.term.clone();
let evaluated = self.term.penalized_quasi_laplace_criterion_with_cache(
self.target.view(),
&rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
);
let (_cost, loss, cache) = match evaluated {
Ok(evaluated) => evaluated,
Err(err) => {
self.term = saved_term;
return Err(err.to_string());
}
};
let residual = self
.term
.reconstruction_residual(self.target.view(), &rho)?;
let dispersion =
self.term
.reconstruction_dispersion(&loss, &cache, &rho, Some(residual.view()))?;
self.term.assemble_shape_uncertainty(&cache, dispersion)
}
fn record_fit_data_collapse_verdict(&mut self, rho: &SaeManifoldRho) -> Result<(), String> {
self.term.record_fit_data_collapse_if_needed(
self.target.view(),
rho,
self.inner_max_iter,
)?;
Ok(())
}
fn probe_value_is_infeasible(value: f64) -> bool {
!value.is_finite()
}
pub(crate) fn is_recoverable_value_probe_refusal(err: &str) -> bool {
ProbeRefusalKind::classify(err).is_some()
}
fn take_probe_converged_handoff(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Option<SaeManifoldTerm> {
let handoff = self.probe_converged_handoff.take()?;
let matches = handoff.rho_flat.len() == rho_flat.len()
&& handoff
.rho_flat
.iter()
.zip(rho_flat.iter())
.all(|(a, b)| a.to_bits() == b.to_bits());
if matches { Some(handoff.term) } else { None }
}
pub(crate) fn evaluate_authoritative_criterion(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<(f64, Array1<f64>), String> {
self.evaluate_authoritative_inner(rho_flat, false)
}
fn evaluate_authoritative_inner(
&mut self,
rho_flat: ArrayView1<'_, f64>,
basin_installed: bool,
) -> Result<(f64, Array1<f64>), String> {
self.fit_verdict = None;
let rho = self.baseline_rho.from_flat(rho_flat)?;
self.apply_block_scaling(&rho)?;
let probe_handoff_installed = if basin_installed {
true
} else if let Some(converged) = self.take_probe_converged_handoff(rho_flat) {
self.term = converged;
self.seeded_beta = None;
true
} else {
false
};
if let Some(beta) = self.seeded_beta.take() {
if beta.len() != self.term.beta_dim() {
return Err(format!(
"seeded decoder has length {}; expected {}",
beta.len(),
self.term.beta_dim()
));
}
self.term.set_flat_beta(beta.view())?;
}
if !probe_handoff_installed && !self.audit_installed_state {
let warm_start_outcome = self
.term
.warm_start_latents_from_amortized_encoder(self.target.view(), &rho);
self.record_warm_start(warm_start_outcome)?;
}
let criterion = self
.term
.penalized_quasi_laplace_criterion_with_refine_policy_and_lane(
self.target.view(),
&rho,
self.registry.as_ref(),
self.inner_max_iter,
self.learning_rate,
self.ridge_ext_coord,
self.ridge_beta,
true,
self.surrogate_lane.as_mut(),
);
let (penalized_quasi_laplace_cost, loss) = match criterion {
Ok(evaluated) => evaluated,
Err(SaeCriterionError::VanishedAtoms(atoms)) => {
log::debug!(
"SAE criterion reached fixed-K structural boundary at rho={:?}: {atoms}",
rho.to_flat()
);
let loss = self.term.loss(self.target.view(), &rho)?;
let beta_hat = self.term.flatten_beta();
self.current_rho = rho;
self.last_loss = Some(loss);
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok((f64::INFINITY, beta_hat));
}
Err(err @ SaeCriterionError::IndefiniteObservedInformation { .. }) => {
self.probe_telemetry.record_refusal_kind(&err.to_string());
log::debug!("SAE criterion mapped indefinite-A refusal to +inf: {err}");
let loss = self.term.loss(self.target.view(), &rho)?;
let beta_hat = self.term.flatten_beta();
self.current_rho = rho;
self.last_loss = Some(loss);
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok((f64::INFINITY, beta_hat));
}
Err(SaeCriterionError::Numerical(message)) => return Err(message),
};
let beta_hat = self.term.flatten_beta();
self.record_fit_data_collapse_verdict(&rho)?;
let cost = if penalized_quasi_laplace_cost.is_finite() {
penalized_quasi_laplace_cost
} else {
log::debug!(
"SAE criterion assembled a NON-FINITE value {penalized_quasi_laplace_cost:.6e} \
(loss total {:.6e}) at the converged inner state — mapping to +inf",
loss.total()
);
self.probe_telemetry.infeasible_criterion_evals += 1;
f64::INFINITY
};
self.current_rho = rho;
self.last_loss = Some(loss);
Ok((cost, beta_hat))
}
pub fn fit_at_fixed_rho(&mut self, rho_flat: ArrayView1<'_, f64>) -> Result<(), String> {
self.fit_verdict = None;
self.terminal_penalized_quasi_laplace_criterion = None;
let rho_state = self.baseline_rho.from_flat(rho_flat.clone())?;
let (criterion, _) = self.evaluate_authoritative_criterion(rho_flat)?;
let jacobian = self.block_jacobian(&rho_state);
let cost = criterion + jacobian;
if !cost.is_finite() {
return Err(format!(
"SaeManifoldOuterObjective::fit_at_fixed_rho: penalized quasi-Laplace criterion \
is infeasible at the requested rho (criterion={criterion:.6e}, \
block_jacobian={jacobian:.6e})"
));
}
self.terminal_penalized_quasi_laplace_criterion = Some(cost);
self.fit_verdict = Some(SaeOuterVerdict::FixedRho);
Ok(())
}
fn evaluate_authoritative_value_probe(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<(f64, Array1<f64>), String> {
let saved_term = self.term.clone();
let saved_rho = self.current_rho.clone();
let saved_loss = self.last_loss.clone();
let saved_seeded_beta = self.seeded_beta.clone();
let result = self.evaluate_authoritative_inner(rho_flat, false);
match &result {
Ok((cost, _beta)) if !Self::probe_value_is_infeasible(*cost) => {
let converged = std::mem::replace(&mut self.term, saved_term);
self.probe_converged_handoff = Some(ProbeConvergedHandoff {
rho_flat: rho_flat.to_owned(),
term: converged,
});
}
_ => {
self.term = saved_term;
}
}
self.current_rho = saved_rho;
self.last_loss = saved_loss;
self.seeded_beta = saved_seeded_beta;
result
}
fn authoritative_envelope_value_probe(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<(f64, Array1<f64>), String> {
if self.inner_max_iter == 0 || !self.term.streaming_plan()?.direct_logdet_admitted() {
return self.evaluate_authoritative_value_probe(rho_flat);
}
if self.basin_bundle.is_empty() {
self.basin_bundle
.admit_distinct(self.term.clone(), f64::INFINITY)
.map_err(|error| format!("SAE basin-envelope seed admission refused: {error}"))?;
}
let discovery = self.evaluate_authoritative_value_probe(rho_flat);
let discovery_cost = match &discovery {
Ok((cost, _)) if !Self::probe_value_is_infeasible(*cost) => Some(*cost),
_ => None,
};
let discovery_term = self.take_probe_converged_handoff(rho_flat);
let mut bundle = std::mem::replace(&mut self.basin_bundle, BasinBundle::new(0));
let member_eval = bundle.evaluate(|state: &SaeManifoldTerm| {
let (res, converged) = self.converge_member_criterion(rho_flat, state);
res.map(|value| (converged, value))
});
let rho_state = self.baseline_rho.from_flat(rho_flat)?;
let ss_tot =
super::fit_drivers::TargetCenteredColStats::compute(self.target.view()).ss_tot();
let len_before = bundle.len();
if let (Some(term), Some(cost)) = (discovery_term, discovery_cost) {
let admission = bundle.admit(term, cost, |a, b| {
Self::same_basin_at_rho(a, b, &rho_state, ss_tot)
});
if let Err(error) = admission {
self.basin_bundle = bundle;
return Err(format!(
"SAE exact basin-envelope discovery admission refused: {error}"
));
}
}
let grew = bundle.len() > len_before;
let bundle_len = bundle.len();
let envelope = bundle
.argmin()
.filter(|m| m.last_value.is_finite())
.map(|m| (m.last_value, m.state.flatten_beta(), m.state.clone()));
self.basin_bundle = bundle;
self.probe_telemetry.basin_envelope_evals += 1;
if grew {
self.probe_telemetry.basin_admissions += 1;
}
self.probe_telemetry.basin_max_members =
self.probe_telemetry.basin_max_members.max(bundle_len);
self.probe_telemetry.basin_member_capacity = self.basin_bundle.member_capacity();
match envelope {
Some((env_value, env_beta, env_term)) => {
if let Some(dcost) = discovery_cost {
let stall = SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL * dcost.abs().max(1.0);
if dcost - env_value > stall {
self.probe_telemetry.basin_envelope_rescues += 1;
}
}
if !Self::probe_value_is_infeasible(env_value) {
self.probe_converged_handoff = Some(ProbeConvergedHandoff {
rho_flat: rho_flat.to_owned(),
term: env_term,
});
}
Ok((env_value, env_beta))
}
None => {
drop(member_eval);
discovery
}
}
}
fn install_authoritative_envelope_basin(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<bool, String> {
if let Some(converged) = self.take_probe_converged_handoff(rho_flat) {
self.term = converged;
self.seeded_beta = None;
return Ok(true);
}
let (cost, _beta) = self.authoritative_envelope_value_probe(rho_flat)?;
if Self::probe_value_is_infeasible(cost) {
return Ok(false);
}
let converged = self
.take_probe_converged_handoff(rho_flat)
.ok_or_else(|| {
"SAE basin-envelope protocol violated: a finite probe at the requested rho did not install its exact-rho converged-state handoff"
.to_string()
})?;
self.term = converged;
self.seeded_beta = None;
Ok(true)
}
fn converge_member_criterion(
&mut self,
rho_flat: ArrayView1<'_, f64>,
member: &SaeManifoldTerm,
) -> (Result<f64, String>, SaeManifoldTerm) {
let saved_term = std::mem::replace(&mut self.term, member.clone());
let saved_rho = self.current_rho.clone();
let saved_loss = self.last_loss.clone();
let saved_seeded_beta = self.seeded_beta.take();
let res = self
.evaluate_authoritative_inner(rho_flat, true)
.map(|(cost, _beta)| cost);
let converged = std::mem::replace(&mut self.term, saved_term);
self.current_rho = saved_rho;
self.last_loss = saved_loss;
self.seeded_beta = saved_seeded_beta;
(res, converged)
}
fn same_basin_at_rho(
a: &SaeManifoldTerm,
b: &SaeManifoldTerm,
rho: &SaeManifoldRho,
ss_tot: f64,
) -> bool {
if !(ss_tot > 0.0) {
return false;
}
let (Ok(fa), Ok(fb)) = (a.try_fitted_for_rho(rho), b.try_fitted_for_rho(rho)) else {
return false;
};
if fa.dim() != fb.dim() {
return false;
}
let mut diff_sq = 0.0_f64;
for (x, y) in fa.iter().zip(fb.iter()) {
let d = x - y;
diff_sq += d * d;
}
(diff_sq / ss_tot) < SAE_FINAL_EV_DEGRADATION_TOL
}
pub(crate) fn efs_step(&mut self, rho_flat: ArrayView1<'_, f64>) -> Result<EfsEval, String> {
self.efs_step_with_certificate(rho_flat)
.map(|(evaluation, _)| evaluation)
}
fn efs_step_with_certificate(
&mut self,
rho_flat: ArrayView1<'_, f64>,
) -> Result<(EfsEval, Vec<FixedPointCoordinateCertificate>), String> {
self.fit_verdict = None;
self.probe_telemetry.criterion_calls += 1;
let rho = self.baseline_rho.from_flat(rho_flat)?;
let n_params = rho.to_flat().len();
self.apply_block_scaling(&rho)?;
let direct_logdet_admitted = self.term.streaming_plan()?.direct_logdet_admitted();
let infeasible_evaluation = |reason: &str| {
(
EfsEval {
cost: f64::INFINITY,
steps: vec![0.0_f64; n_params],
beta: None,
psi_gradient: None,
psi_indices: None,
inner_hessian_scale: None,
logdet_enclosure_gap: None,
consecutive_restored_incumbents: None,
},
(0..n_params)
.map(|index| {
FixedPointCoordinateCertificate::uncovered(format!(
"coordinate {index}: fixed-point evidence unavailable: {reason}"
))
})
.collect(),
)
};
if direct_logdet_admitted && self.inner_max_iter != 0 {
match self.install_authoritative_envelope_basin(rho_flat) {
Ok(true) => {}
Ok(false) => {
self.probe_converged_handoff = None;
self.basin_bundle.clear();
self.current_rho = rho;
return Ok(infeasible_evaluation(
"the authoritative basin envelope is infeasible at this rho",
));
}
Err(err) if Self::is_recoverable_value_probe_refusal(&err) => {
self.probe_telemetry.record_refusal_kind(&err);
self.probe_telemetry.infeasible_criterion_evals += 1;
self.probe_converged_handoff = None;
self.basin_bundle.clear();
self.current_rho = rho;
return Ok(infeasible_evaluation(
"infeasible penalized quasi-Laplace basin envelope",
));
}
Err(err) => return Err(err),
}
}
self.probe_converged_handoff = None;
self.basin_bundle.clear();
if let Some(beta) = self.seeded_beta.take()
&& beta.len() == self.term.beta_dim()
{
self.term.set_flat_beta(beta.view())?;
}
let criterion = self.evaluate_outer_criterion_route(&rho, direct_logdet_admitted, true);
let evaluation = match criterion {
Ok(evaluated) => evaluated,
Err(SaeCriterionError::VanishedAtoms(atoms)) => {
log::debug!("SAE EFS probe reached fixed-K structural boundary: {atoms}");
self.probe_telemetry.infeasible_criterion_evals += 1;
self.current_rho = rho;
return Ok(infeasible_evaluation("vanished-atom structural boundary"));
}
Err(SaeCriterionError::Numerical(err))
if Self::is_recoverable_value_probe_refusal(&err) =>
{
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
self.current_rho = rho;
return Ok(infeasible_evaluation(
"infeasible penalized quasi-Laplace score",
));
}
Err(err @ SaeCriterionError::IndefiniteObservedInformation { .. }) => {
self.probe_telemetry.record_refusal_kind(&err.to_string());
log::debug!("SAE criterion mapped indefinite-A refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
self.current_rho = rho;
return Ok(infeasible_evaluation(
"infeasible penalized quasi-Laplace score (indefinite exact A)",
));
}
Err(SaeCriterionError::Numerical(err)) => return Err(err),
};
let cost = evaluation.cost;
self.record_fit_data_collapse_verdict(&rho)?;
self.current_rho = rho.clone();
if !cost.is_finite() {
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(infeasible_evaluation(
"the penalized quasi-Laplace criterion is non-finite",
));
}
let n_eff = self.term.n_obs() as f64;
let sumsq = self.term.ard_coord_sumsq();
let complete_gradient = if rho.sparse_flat_index().is_some() || !rho.kappa.is_empty() {
Some(
self.analytic_gradient_for_outer_evaluation(&rho, &evaluation)
.map_err(|error| error.to_string())?,
)
} else {
None
};
let cache = &evaluation.cache;
let inverse_probe_bundle = evaluation
.matrix_free
.as_ref()
.and_then(|artifacts| artifacts.efs_inverse_probe_bundle.as_ref());
let traces = if let Some((probes, sinv)) = inverse_probe_bundle.as_ref() {
self.term
.ard_inverse_traces_from_probes(cache, probes, sinv)
.map_err(|e| {
format!("SaeManifoldOuterObjective::efs_step: ARD traces (matrix-free): {e}")
})?
} else {
self.term
.ard_inverse_traces(cache)
.map_err(|e| format!("SaeManifoldOuterObjective::efs_step: ARD traces: {e}"))?
};
let mut steps = vec![0.0_f64; n_params];
let mut fixed_point_coordinates = (0..n_params)
.map(|index| {
FixedPointCoordinateCertificate::uncovered(format!(
"coordinate {index}: no root-equivalent fixed-point equation was evaluated"
))
})
.collect::<Vec<_>>();
let mut psi_gradient = Vec::new();
let mut psi_indices = Vec::new();
if let Some(sparse_index) = rho.sparse_flat_index() {
assert_eq!(
assignment_strength_gradient_coordinate(&rho),
Some(sparse_index)
);
let gradient = complete_gradient
.as_ref()
.expect("sparse rho coordinate requested its complete analytic gradient")
[sparse_index];
let gradient_scale = gradient.abs().max(1.0);
let step = -gradient / gradient_scale;
steps[sparse_index] = step;
fixed_point_coordinates[sparse_index] =
FixedPointCoordinateCertificate::covered(step, 1.0);
psi_gradient.push(gradient);
psi_indices.push(sparse_index);
}
for &atom in &rho.kappa_atoms {
let coordinate = rho.kappa_flat_index(atom).ok_or_else(|| {
format!(
"SaeManifoldOuterObjective::efs_step: atom {atom} has curvature state but no flat coordinate"
)
})?;
let gradient = complete_gradient
.as_ref()
.expect("curvature coordinate requested its complete analytic gradient")
[coordinate];
let step = -gradient / gradient.abs().max(1.0);
steps[coordinate] = step;
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::covered(step, 1.0);
psi_gradient.push(gradient);
psi_indices.push(coordinate);
}
let k_smooth = rho.log_lambda_smooth.len();
let lambda_smooth_vec = rho.lambda_smooth_vec()?;
let quad_per_atom = self.term.decoder_smoothness_quadratic_form_per_atom()?;
let eff_dof_per_atom = if let Some((probes, sinv)) = inverse_probe_bundle.as_ref() {
self.term
.decoder_smoothness_effective_dof_per_atom_from_probes(
probes,
sinv,
&lambda_smooth_vec,
)
.map_err(|e| {
format!("SaeManifoldOuterObjective::efs_step: smooth dof (matrix-free): {e}")
})?
} else {
self.term
.decoder_smoothness_effective_dof_per_atom(&cache, &lambda_smooth_vec)
.map_err(|e| format!("SaeManifoldOuterObjective::efs_step: smooth dof: {e}"))?
};
for atom_idx in 0..k_smooth {
let coordinate = rho.smooth_flat_index(atom_idx);
let lambda_k = lambda_smooth_vec[atom_idx];
let rank_k = (self.term.atoms[atom_idx].border_frame_rank() as f64)
* (SaeManifoldTerm::symmetric_rank(self.term.atoms[atom_idx].smooth_penalty())?
as f64);
let quad_k = quad_per_atom[atom_idx];
let eff_dof_k = eff_dof_per_atom[atom_idx];
if !(quad_k > 0.0) {
fixed_point_coordinates[coordinate] = FixedPointCoordinateCertificate::uncovered(
format!("atom {atom_idx} smoothness energy is not positive"),
);
} else if !(rank_k - eff_dof_k > 0.0) {
fixed_point_coordinates[coordinate] = FixedPointCoordinateCertificate::uncovered(
format!("atom {atom_idx} smoothness rank-minus-edf numerator is not positive"),
);
} else if !(lambda_k > 0.0 && lambda_k.is_finite()) {
fixed_point_coordinates[coordinate] = FixedPointCoordinateCertificate::uncovered(
format!("atom {atom_idx} smoothness precision is not finite and positive"),
);
} else {
let lambda_new = (rank_k - eff_dof_k) / quad_k;
if lambda_new.is_finite() && lambda_new > 0.0 {
let step = lambda_new.ln() - rho.log_lambda_smooth[atom_idx];
steps[coordinate] = step;
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::covered(step, 1.0);
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"atom {atom_idx} smoothness equation proposed a non-finite precision"
));
}
}
}
let ard_periods: Vec<Vec<Option<f64>>> = self
.term
.assignment
.coords
.iter()
.map(|c| c.effective_axis_periods())
.collect();
match rho.ard_sharing() {
ArdSharing::PerAtom => {
for (k, axis_logard) in rho.log_ard.iter().enumerate() {
for (j, &logard_kj) in axis_logard.iter().enumerate() {
let coordinate = rho.ard_flat_index(k, j);
let denom = sumsq[k][j] + traces[k][j];
if denom > 0.0 {
let alpha_gauss = n_eff / denom;
let alpha_new = match ard_periods[k].get(j).copied().flatten() {
Some(period) => von_mises_ard_precision(
alpha_gauss,
std::f64::consts::TAU / period,
),
None => alpha_gauss,
};
if alpha_new.is_finite() && alpha_new > 0.0 {
let step = alpha_new.ln() - logard_kj;
steps[coordinate] = step;
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::covered(step, 1.0);
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"atom {k} ARD axis {j} equation proposed a non-finite precision"
));
}
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"atom {k} ARD axis {j} posterior second moment is not positive"
));
}
}
}
}
ArdSharing::Shared => {
let max_d = rho.max_ard_axes();
for axis in 0..max_d {
let mut denom = 0.0_f64;
let mut count = 0usize;
let mut shared_logard = 0.0_f64;
let mut shared_period: Option<f64> = None;
for (k, axis_logard) in rho.log_ard.iter().enumerate() {
if axis < axis_logard.len() {
denom += sumsq[k][axis] + traces[k][axis];
shared_logard = axis_logard[axis];
if shared_period.is_none() {
shared_period = ard_periods[k].get(axis).copied().flatten();
}
count += 1;
}
}
let coordinate = rho.ard_flat_index(0, axis);
if count == 0 {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"shared ARD axis {axis} has no owning atom"
));
} else if !(denom > 0.0) {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"shared ARD axis {axis} posterior second moment is not positive"
));
} else {
let alpha_gauss = n_eff * (count as f64) / denom;
let alpha_new = match shared_period {
Some(period) => {
von_mises_ard_precision(alpha_gauss, std::f64::consts::TAU / period)
}
None => alpha_gauss,
};
if alpha_new.is_finite() && alpha_new > 0.0 {
let step = alpha_new.ln() - shared_logard;
steps[coordinate] = step;
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::covered(step, 1.0);
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"shared ARD axis {axis} equation proposed a non-finite precision"
));
}
}
}
}
}
if let Some(scaled_rss) = self.block_scaled_rss(&rho)? {
let n = self.term.n_obs() as f64;
let blocks = self
.crosscoder_blocks
.as_ref()
.expect("block_scaled_rss returned Some ⇒ crosscoder pricing is installed");
let tail = n_params - rho.kappa.len() - blocks.block_dims.len();
for (l, (&p_l, &r_tilde)) in blocks.block_dims.iter().zip(scaled_rss.iter()).enumerate()
{
let coordinate = tail + l;
if r_tilde > 0.0 {
let step = (n * p_l as f64 / r_tilde).ln();
if step.is_finite() {
steps[coordinate] = step;
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"crosscoder block {l} EFS proposal omits the logdet IFT adjoint and is not a complete stationarity equation"
));
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"crosscoder block {l} equation proposed a non-finite update"
));
}
} else {
fixed_point_coordinates[coordinate] =
FixedPointCoordinateCertificate::uncovered(format!(
"crosscoder block {l} scaled residual energy is not positive"
));
}
}
}
let beta_hat = self.term.flatten_beta();
self.last_loss = Some(evaluation.loss);
let consecutive_restored_incumbents = self
.term
.best_fit_incumbent
.as_ref()
.map(|incumbent| incumbent.consecutive_inner_restores);
Ok((
EfsEval {
cost,
steps,
beta: Some(beta_hat),
psi_gradient: (!psi_gradient.is_empty()).then(|| Array1::from_vec(psi_gradient)),
psi_indices: (!psi_indices.is_empty()).then_some(psi_indices),
inner_hessian_scale: None,
logdet_enclosure_gap: None,
consecutive_restored_incumbents,
},
fixed_point_coordinates,
))
}
}
fn von_mises_ard_precision(alpha_gauss: f64, kappa: f64) -> f64 {
if !(alpha_gauss.is_finite() && alpha_gauss > 0.0 && kappa.is_finite() && kappa > 0.0) {
return alpha_gauss;
}
let kappa2 = kappa * kappa;
let eta_gauss = alpha_gauss / kappa2;
let a_target = 1.0 - 0.5 / eta_gauss;
if !(a_target > 0.0 && a_target < 1.0) {
return alpha_gauss;
}
let a_of = |eta: f64| bessel_i0_log_and_ratio(eta).1;
let mut lo = eta_gauss;
let mut hi = eta_gauss;
let mut guard = 0;
while lo > f64::MIN_POSITIVE && a_of(lo) > a_target && guard < 256 {
lo *= 0.5;
guard += 1;
}
guard = 0;
while hi.is_finite() && a_of(hi) < a_target && guard < 256 {
hi *= 2.0;
guard += 1;
}
if !(lo.is_finite() && hi.is_finite() && lo > 0.0 && hi > lo) {
return alpha_gauss;
}
for _ in 0..80 {
let mid = 0.5 * (lo + hi);
if a_of(mid) < a_target {
lo = mid;
} else {
hi = mid;
}
}
let alpha = kappa2 * 0.5 * (lo + hi);
if alpha.is_finite() && alpha > 0.0 {
alpha
} else {
alpha_gauss
}
}
fn reactive_smooth_curvature_scale(
term: &SaeManifoldTerm,
assignments: &Array2<f64>,
atom_idx: usize,
) -> Result<Option<f64>, String> {
let atom = &term.atoms[atom_idx];
let m = atom.basis_values.ncols();
if atom.smooth_penalty().dim() != (m, m) {
return Err(format!(
"reactive rho domain: atom {atom_idx} smooth penalty shape {:?} != ({m}, {m})",
atom.smooth_penalty().dim()
));
}
let penalty_geometry =
gam_linalg::utils::rank_certified_psd_pseudoinverse(atom.smooth_penalty(), 1.0e-10)
.map_err(|error| format!("reactive rho domain penalty spectrum failed: {error}"))?;
let rank = penalty_geometry.rank();
let penalty_pinv = penalty_geometry.into_pseudoinverse();
if rank == 0 {
return Ok(None);
}
let whitens = term
.row_metric
.as_ref()
.is_some_and(gam_problem::RowMetric::whitens_likelihood);
let mut data_gram = Array2::<f64>::zeros((m, m));
for row in 0..term.n_obs() {
let honesty_weight = term
.row_loss_weights
.as_ref()
.map_or(1.0, |weights| weights[row]);
let metric_norm_bound = match term.row_metric.as_ref() {
Some(metric) if whitens => metric.row_traces()[row],
_ => 1.0,
};
let gate = assignments[[row, atom_idx]];
let weight = honesty_weight * metric_norm_bound * gate * gate;
if !(weight.is_finite() && weight >= 0.0) {
return Err(format!(
"reactive rho domain: atom {atom_idx} row {row} has invalid data-curvature weight {weight}"
));
}
for left in 0..m {
let weighted_left = weight * atom.basis_values[[row, left]];
for right in 0..m {
data_gram[[left, right]] += weighted_left * atom.basis_values[[row, right]];
}
}
}
let (pinv_eigenvalues, pinv_eigenvectors) =
gam_linalg::faer_ndarray::strict_symmetric_eigh(&penalty_pinv, Side::Lower)
.map_err(|error| format!("reactive rho domain P⁺ spectrum failed: {error}"))?;
if !pinv_eigenvalues.iter().all(|value| value.is_finite()) {
return Err(format!(
"reactive rho domain: atom {atom_idx} P⁺ spectrum is non-finite"
));
}
let mut order: Vec<usize> = (0..m).collect();
order.sort_by(|&left, &right| {
pinv_eigenvalues[right]
.partial_cmp(&pinv_eigenvalues[left])
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut scaled_vectors = Array2::<f64>::zeros((m, m));
for &col in order.iter().take(rank) {
let eigenvalue = pinv_eigenvalues[col];
if !(eigenvalue.is_finite() && eigenvalue > 0.0) {
return Err(format!(
"reactive rho domain: atom {atom_idx} retained P⁺ eigenvalue is invalid ({eigenvalue})"
));
}
let scale = eigenvalue.sqrt();
for row in 0..m {
scaled_vectors[[row, col]] = pinv_eigenvectors[[row, col]] * scale;
}
}
let pinv_sqrt = scaled_vectors.dot(&pinv_eigenvectors.t());
let mut standardized_curvature = pinv_sqrt.dot(&data_gram).dot(&pinv_sqrt);
for row in 0..m {
for col in 0..row {
let symmetric =
0.5 * (standardized_curvature[[row, col]] + standardized_curvature[[col, row]]);
standardized_curvature[[row, col]] = symmetric;
standardized_curvature[[col, row]] = symmetric;
}
}
let (generalized_eigenvalues, _) = standardized_curvature
.eigh(Side::Lower)
.map_err(|error| format!("reactive rho domain generalized spectrum failed: {error}"))?;
if !generalized_eigenvalues
.iter()
.all(|value| value.is_finite())
{
return Err(format!(
"reactive rho domain: atom {atom_idx} generalized decoder spectrum is non-finite"
));
}
let largest = generalized_eigenvalues
.iter()
.copied()
.fold(0.0_f64, f64::max);
Ok((largest > 0.0).then_some(largest))
}
fn reactive_ard_curvature_scale(
term: &SaeManifoldTerm,
assignments: &Array2<f64>,
atom_idx: usize,
axis: usize,
) -> Result<Option<f64>, String> {
let periods = term.assignment.coords[atom_idx].effective_axis_periods();
if periods.get(axis).copied().flatten().is_some() {
return Ok(None);
}
observed_ard_curvature_scale(term, assignments, atom_idx, axis).map(Some)
}
fn observed_ard_curvature_scale(
term: &SaeManifoldTerm,
assignments: &Array2<f64>,
atom_idx: usize,
axis: usize,
) -> Result<f64, String> {
let atom = &term.atoms[atom_idx];
let p = atom.decoder_coefficients().ncols();
let m = atom.decoder_coefficients().nrows();
if atom.basis_jacobian.dim().1 != m || axis >= atom.basis_jacobian.dim().2 {
return Err(format!(
"reactive rho domain: atom {atom_idx} axis {axis} is incompatible with basis Jacobian {:?} and decoder {:?}",
atom.basis_jacobian.dim(),
atom.decoder_coefficients().dim()
));
}
let whitens = term
.row_metric
.as_ref()
.is_some_and(gam_problem::RowMetric::whitens_likelihood);
let mut tangent = vec![0.0_f64; p];
let mut maximum = 0.0_f64;
for row in 0..term.n_obs() {
tangent.fill(0.0);
let gate = assignments[[row, atom_idx]];
for basis in 0..m {
let coefficient = gate * atom.basis_jacobian[[row, basis, axis]];
for out in 0..p {
tangent[out] += coefficient * atom.decoder_coefficients()[[basis, out]];
}
}
let tangent_norm_sq = match term.row_metric.as_ref() {
Some(metric) if whitens => metric
.whiten_residual_row(row, ArrayView1::from(tangent.as_slice()))
.into_iter()
.map(|value| value * value)
.sum::<f64>(),
_ => tangent.iter().map(|value| value * value).sum(),
};
let honesty_weight = term
.row_loss_weights
.as_ref()
.map_or(1.0, |weights| weights[row]);
let curvature = honesty_weight * tangent_norm_sq;
if !(curvature.is_finite() && curvature >= 0.0) {
return Err(format!(
"reactive rho domain: atom {atom_idx} axis {axis} row {row} has invalid latent data curvature {curvature}"
));
}
maximum = maximum.max(curvature);
}
Ok(maximum)
}
fn reactive_rho_domain_upper(
term: &SaeManifoldTerm,
rho: &SaeManifoldRho,
entry_temperature: f64,
) -> Result<Array1<f64>, String> {
term.assignment.validate_rho_domain(rho)?;
let mut entry_term = term.clone();
entry_term
.assignment
.mode
.set_temperature(entry_temperature)?;
entry_term.temperature_schedule = None;
let assignments = entry_term.assignment.try_assignments()?;
let target = rho.to_flat();
let mut upper = target.clone();
let mut largest_native_scale = 0.0_f64;
for atom_idx in 0..rho.k_atoms() {
if let Some(scale) = reactive_smooth_curvature_scale(&entry_term, &assignments, atom_idx)? {
largest_native_scale = largest_native_scale.max(scale);
let index = rho.smooth_flat_index(atom_idx);
let target_strength = target[index].exp();
upper[index] = target_strength.max(scale).ln();
}
for axis in 0..rho.log_ard[atom_idx].len() {
if let Some(scale) =
reactive_ard_curvature_scale(&entry_term, &assignments, atom_idx, axis)?
{
largest_native_scale = largest_native_scale.max(scale);
let index = rho.ard_flat_index(atom_idx, axis);
let target_strength = target[index].exp();
upper[index] = upper[index].max(target_strength.max(scale).ln());
}
}
}
if let Some(index) = rho.sparse_flat_index()
&& !matches!(
entry_term.assignment.mode,
AssignmentMode::OrderedBetaBernoulli {
learnable_alpha: false,
..
}
)
&& largest_native_scale > 0.0
{
let target_strength = target[index].exp();
upper[index] = target_strength.max(largest_native_scale).ln();
}
if upper.iter().all(|value| value.is_finite()) {
Ok(upper)
} else {
Err(format!(
"reactive rho domain produced a non-finite upper face: {upper:?}"
))
}
}
pub(crate) fn periodic_ard_domain_upper(
term: &SaeManifoldTerm,
rho: &SaeManifoldRho,
assignments: &Array2<f64>,
) -> Result<Vec<(usize, f64)>, String> {
let n_rows = term.n_obs();
let mut faces = Vec::new();
if n_rows < 2 {
return Ok(faces);
}
for atom_idx in 0..rho.k_atoms() {
if atom_idx >= term.assignment.coords.len() {
return Err(format!(
"periodic ARD chart-resolution domain: atom {atom_idx} has no assignment \
coordinates ({} present)",
term.assignment.coords.len()
));
}
let periods = term.assignment.coords[atom_idx].effective_axis_periods();
for axis in 0..rho.log_ard[atom_idx].len() {
let Some(period) = periods.get(axis).copied().flatten() else {
continue;
};
if !(period.is_finite() && period > 0.0) {
return Err(format!(
"periodic ARD chart-resolution domain: atom {atom_idx} axis {axis} has \
non-positive period {period}"
));
}
let resolution_face = 2.0 * ((2.0 * n_rows as f64) / period).ln();
if !resolution_face.is_finite() {
return Err(format!(
"periodic ARD chart-resolution domain: atom {atom_idx} axis {axis} produced a \
non-finite face from n={n_rows} period={period}"
));
}
let curvature = observed_ard_curvature_scale(term, assignments, atom_idx, axis)?;
let face = if curvature > 0.0 {
resolution_face.min(curvature.ln())
} else {
resolution_face
};
if !face.is_finite() {
return Err(format!(
"periodic ARD domain: atom {atom_idx} axis {axis} produced a non-finite face \
from resolution={resolution_face} curvature={curvature}"
));
}
faces.push((rho.ard_flat_index(atom_idx, axis), face));
}
}
Ok(faces)
}
impl OuterObjective for SaeManifoldOuterObjective {
fn capability(&self) -> OuterCapability {
let gradient = sae_outer_gradient_capability();
let exact_gradient_certificate = matches!(gradient, Derivative::Analytic);
let psi_gradient_dim =
usize::from(assignment_strength_gradient_coordinate(&self.baseline_rho).is_some())
+ self.baseline_rho.kappa.len();
OuterCapability {
gradient,
hessian: DeclaredHessianForm::Unavailable,
n_params: self.baseline_rho.to_flat().len(),
psi_dim: if exact_gradient_certificate {
psi_gradient_dim
} else {
0
},
fixed_point_available: exact_gradient_certificate,
barrier_config: None,
prefer_gradient_only: false,
disable_fixed_point: !exact_gradient_certificate,
}
}
fn eval_cost(&mut self, rho: &Array1<f64>) -> Result<f64, EstimationError> {
self.check_cancelled()?;
self.probe_telemetry.criterion_calls += 1;
match self.authoritative_envelope_value_probe(rho.view()) {
Ok((cost, _beta)) => {
let rho_state = self
.baseline_rho
.from_flat(rho.view())
.map_err(EstimationError::InvalidInput)?;
let cost = cost + self.block_jacobian(&rho_state);
if !cost.is_finite() {
return Ok(f64::INFINITY);
}
if self.reactive_waypoint_checkpoint.is_none()
&& self.record_search_criterion(cost, None)
{
self.bank_checkpoint(rho);
}
Ok(cost)
}
Err(err) if Self::is_recoverable_value_probe_refusal(&err) => {
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
Ok(f64::INFINITY)
}
Err(err) => Err(EstimationError::RemlOptimizationFailed(err)),
}
}
fn eval(&mut self, rho: &Array1<f64>) -> Result<OuterEval, EstimationError> {
self.check_cancelled()?;
self.probe_telemetry.criterion_calls += 1;
let rho_state = self
.baseline_rho
.from_flat(rho.view())
.map_err(EstimationError::InvalidInput)?;
self.apply_block_scaling(&rho_state)
.map_err(EstimationError::InvalidInput)?;
if !self.audit_installed_state
&& !self
.term
.streaming_plan()
.map_err(EstimationError::RemlOptimizationFailed)?
.direct_logdet_admitted()
{
let (cost, _beta_hat) = match self.evaluate_authoritative_criterion(rho.view()) {
Ok(evaluated) => evaluated,
Err(err) if Self::is_recoverable_value_probe_refusal(&err) => {
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
Err(err) => return Err(EstimationError::RemlOptimizationFailed(err)),
};
let cost = cost + self.block_jacobian(&rho_state);
if !cost.is_finite() {
return Ok(OuterEval::infeasible(rho.len()));
}
if self.record_search_criterion(cost, None) {
self.bank_checkpoint(rho);
}
return Ok(OuterEval {
cost,
gradient: Array1::zeros(rho.len()),
hessian: HessianValue::Unavailable,
inner_beta_hint: None,
});
}
match self.install_authoritative_envelope_basin(rho.view()) {
Ok(true) => {}
Ok(false) => return Ok(OuterEval::infeasible(rho.len())),
Err(err) if Self::is_recoverable_value_probe_refusal(&err) => {
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
Err(err) => return Err(EstimationError::RemlOptimizationFailed(err)),
}
let direct_logdet_admitted = self
.term
.streaming_plan()
.map_err(EstimationError::RemlOptimizationFailed)?
.direct_logdet_admitted();
let evaluation =
match self.evaluate_outer_criterion_route(&rho_state, direct_logdet_admitted, false) {
Ok(evaluated) => evaluated,
Err(SaeCriterionError::VanishedAtoms(atoms)) => {
log::debug!(
"SAE analytic evaluation reached fixed-K structural boundary: {atoms}"
);
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
Err(SaeCriterionError::Numerical(err))
if Self::is_recoverable_value_probe_refusal(&err) =>
{
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
Err(err @ SaeCriterionError::IndefiniteObservedInformation { .. }) => {
self.probe_telemetry.record_refusal_kind(&err.to_string());
log::debug!("SAE criterion mapped indefinite-A refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
Err(SaeCriterionError::Numerical(err)) => {
return Err(EstimationError::RemlOptimizationFailed(err));
}
};
let cost = evaluation.cost;
self.record_fit_data_collapse_verdict(&rho_state)
.map_err(EstimationError::RemlOptimizationFailed)?;
if !cost.is_finite() {
self.probe_telemetry.infeasible_criterion_evals += 1;
return Ok(OuterEval::infeasible(rho.len()));
}
let gradient = self
.analytic_gradient_for_outer_evaluation(&rho_state, &evaluation)
.map_err(EstimationError::from)?;
let beta_hat = self.term.flatten_beta();
let hessian = match self.term.exact_fixed_stratum_outer_hessian(
self.target.view(),
&rho_state,
&evaluation.loss,
&evaluation.cache,
) {
Ok(dense) => HessianValue::Dense(dense),
Err(_incomplete) => HessianValue::Unavailable,
};
let cost = cost + self.block_jacobian(&rho_state);
self.current_rho = rho_state;
self.last_loss = Some(evaluation.loss);
if self.record_search_criterion(cost, Some(gradient.dot(&gradient).sqrt())) {
self.bank_checkpoint(rho);
}
Ok(OuterEval {
cost,
gradient,
hessian,
inner_beta_hint: Some(beta_hat),
})
}
fn eval_with_order(
&mut self,
rho: &Array1<f64>,
order: OuterEvalOrder,
) -> Result<OuterEval, EstimationError> {
self.check_cancelled()?;
match order {
OuterEvalOrder::Value => {
let (cost, beta_hat) = match self.authoritative_envelope_value_probe(rho.view()) {
Ok(evaluated) => evaluated,
Err(err) if Self::is_recoverable_value_probe_refusal(&err) => {
self.probe_telemetry.record_refusal_kind(&err);
log::debug!("SAE criterion eval mapped refusal to +inf: {err}");
self.probe_telemetry.infeasible_criterion_evals += 1;
if self.reactive_waypoint_checkpoint.is_some() {
return Err(EstimationError::RemlOptimizationFailed(format!(
"reactive coupled waypoint has undefined penalized quasi-Laplace score: {err}"
)));
}
return Ok(OuterEval::infeasible(rho.len()));
}
Err(err) => return Err(EstimationError::RemlOptimizationFailed(err)),
};
let rho_state = self
.baseline_rho
.from_flat(rho.view())
.map_err(EstimationError::InvalidInput)?;
let cost = cost + self.block_jacobian(&rho_state);
if !cost.is_finite() {
return Ok(OuterEval::infeasible(rho.len()));
}
if self.reactive_waypoint_checkpoint.is_none()
&& self.record_search_criterion(cost, None)
{
self.bank_checkpoint(rho);
}
Ok(OuterEval {
cost,
gradient: Array1::zeros(rho.len()),
hessian: HessianValue::Unavailable,
inner_beta_hint: Some(beta_hat),
})
}
OuterEvalOrder::ValueAndGradient => self.eval(rho),
OuterEvalOrder::ValueGradientHessian => self.eval(rho),
}
}
fn eval_efs(&mut self, rho: &Array1<f64>) -> Result<EfsEval, EstimationError> {
self.check_cancelled()?;
let mut eval = self
.efs_step(rho.view())
.map_err(EstimationError::RemlOptimizationFailed)?;
let rho_state = self
.baseline_rho
.from_flat(rho.view())
.map_err(EstimationError::InvalidInput)?;
eval.cost += self.block_jacobian(&rho_state);
if self.record_search_criterion(eval.cost, None) {
self.bank_checkpoint(rho);
}
Ok(eval)
}
fn eval_fixed_point_certificate(
&mut self,
rho: &Array1<f64>,
) -> Result<FixedPointCertificateEval, EstimationError> {
self.check_cancelled()?;
let evaluation = self.eval(rho)?;
let coordinates = evaluation
.gradient
.iter()
.map(|&gradient| FixedPointCoordinateCertificate::covered(-gradient, 1.0))
.collect();
Ok(FixedPointCertificateEval {
cost: evaluation.cost,
coordinates,
})
}
fn reset(&mut self) {
self.reactive_waypoint_checkpoint = None;
self.fit_verdict = None;
self.term = self.baseline_term.clone();
if let Some(registry) = self.registry.as_mut() {
registry.set_isometry_scalar_weights(&self.baseline_isometry_weights);
}
self.current_rho = self.baseline_rho.clone();
self.last_loss = None;
self.terminal_penalized_quasi_laplace_criterion = None;
self.seeded_beta = None;
self.probe_converged_handoff = None;
self.basin_bundle.clear();
self.termination.reset_improvement_baseline();
}
fn seed_inner_state(&mut self, beta: &Array1<f64>) -> Result<SeedOutcome, EstimationError> {
self.fit_verdict = None;
if beta.is_empty() {
return Ok(SeedOutcome::NoSlot);
}
if beta.len() != self.term.beta_dim() {
return Err(EstimationError::RemlOptimizationFailed(format!(
"SaeManifoldOuterObjective::seed_inner_state: β length {} != decoder dim {}",
beta.len(),
self.term.beta_dim()
)));
}
self.seeded_beta = Some(beta.clone());
self.probe_converged_handoff = None;
self.basin_bundle.clear();
Ok(SeedOutcome::Installed)
}
fn outer_domain_upper_bound(&self) -> Result<Option<Array1<f64>>, EstimationError> {
self.baseline_term
.assignment
.validate_rho_domain(&self.baseline_rho)
.map_err(EstimationError::InvalidInput)?;
let mut log_strength_upper = self.baseline_rho.flat_domain_upper_bound();
if let Some((_, alpha_upper)) = self
.baseline_term
.assignment
.learnable_alpha_rho_domain()
.map_err(EstimationError::InvalidInput)?
&& let (Some(bounds), Some(index)) = (
log_strength_upper.as_mut(),
self.baseline_rho.sparse_flat_index(),
)
{
bounds[index] = bounds[index].min(alpha_upper);
}
let curvature_bounds = self.curvature_domain_bounds()?;
let baseline_assignments = self
.baseline_term
.assignment
.try_assignments()
.map_err(EstimationError::RemlOptimizationFailed)?;
let chart_faces = periodic_ard_domain_upper(
&self.baseline_term,
&self.baseline_rho,
&baseline_assignments,
)
.map_err(EstimationError::RemlOptimizationFailed)?;
let target = self.baseline_rho.to_flat();
let Some(contract) = self.reactive_domain_scalar_contract()? else {
if let Some(bounds) = log_strength_upper.as_mut() {
for &(index, face) in &chart_faces {
bounds[index] = bounds[index].min(face.max(target[index]));
}
for &(index, _, upper) in &curvature_bounds {
bounds[index] = upper;
}
}
return Ok(log_strength_upper);
};
let mut entry_term = self.baseline_term.clone();
entry_term
.assignment
.mode
.set_temperature(contract.entry().assignment_temperature)
.map_err(EstimationError::RemlOptimizationFailed)?;
entry_term.temperature_schedule = None;
entry_term
.place_reactive_entry_disjoint_charts(self.target.view())
.map_err(|error| {
EstimationError::RemlOptimizationFailed(format!(
"reactive rho domain could not construct its separated entry geometry: {error}"
))
})?;
let mut reactive_upper = reactive_rho_domain_upper(
&entry_term,
&self.baseline_rho,
contract.entry().assignment_temperature,
)
.map_err(EstimationError::RemlOptimizationFailed)?;
if let Some(log_strength_upper) = log_strength_upper {
for index in 0..reactive_upper.len() {
reactive_upper[index] = reactive_upper[index].min(log_strength_upper[index]);
}
}
for &(index, face) in &chart_faces {
reactive_upper[index] = reactive_upper[index].min(face.max(target[index]));
}
for &(index, _, upper) in &curvature_bounds {
reactive_upper[index] = upper;
}
Ok(Some(reactive_upper))
}
fn outer_domain_lower_bound(&self) -> Result<Option<Array1<f64>>, EstimationError> {
self.baseline_term
.assignment
.validate_rho_domain(&self.baseline_rho)
.map_err(EstimationError::InvalidInput)?;
let mut lower = self.baseline_rho.flat_domain_lower_bound();
if let Some((alpha_lower, _)) = self
.baseline_term
.assignment
.learnable_alpha_rho_domain()
.map_err(EstimationError::InvalidInput)?
&& let (Some(bounds), Some(index)) =
(lower.as_mut(), self.baseline_rho.sparse_flat_index())
{
bounds[index] = bounds[index].max(alpha_lower);
}
if let Some(bounds) = lower.as_mut() {
for (index, curvature_lower, _) in self.curvature_domain_bounds()? {
bounds[index] = curvature_lower;
}
}
Ok(lower)
}
fn reactive_domain_scalar_contract(
&self,
) -> Result<Option<gam_solve::continuation_path::ContinuationScalarContract>, EstimationError>
{
if self.baseline_term.k_atoms() < 2
|| !self
.baseline_term
.streaming_plan()
.map_err(EstimationError::RemlOptimizationFailed)?
.direct_logdet_admitted()
{
return Ok(None);
}
let target_temperature = self.baseline_term.assignment.mode.temperature();
let routing_logits = self
.baseline_term
.assignment
.frozen_logits
.as_ref()
.unwrap_or(&self.baseline_term.assignment.logits);
let threshold = match self.baseline_term.assignment.mode {
AssignmentMode::ThresholdGate { threshold, .. } => threshold,
_ => 0.0,
};
let mut routing_scale = 0.0_f64;
for &logit in routing_logits {
let centered = logit - threshold;
if !centered.is_finite() {
return Err(EstimationError::RemlOptimizationFailed(
"reactive scalar continuation found a non-finite literal routing logit"
.to_string(),
));
}
routing_scale = routing_scale.max(centered.abs());
}
let entry = gam_solve::continuation_path::ContinuationScalarState::new(
target_temperature.max(routing_scale),
vec![0.0; self.baseline_isometry_weights.len()],
)
.map_err(EstimationError::RemlOptimizationFailed)?;
let target = gam_solve::continuation_path::ContinuationScalarState::new(
target_temperature,
self.baseline_isometry_weights.clone(),
)
.map_err(EstimationError::RemlOptimizationFailed)?;
gam_solve::continuation_path::ContinuationScalarContract::new(entry, target)
.map(Some)
.map_err(EstimationError::RemlOptimizationFailed)
}
fn install_reactive_domain_scalar_state(
&mut self,
state: &gam_solve::continuation_path::ContinuationScalarState,
) -> Result<(), EstimationError> {
let contract = self
.reactive_domain_scalar_contract()?
.ok_or_else(|| {
EstimationError::RemlOptimizationFailed(
"reactive scalar waypoint requested from an objective without a dense K>=2 contract"
.to_string(),
)
})?;
if state.isometry_weights.len() != self.baseline_isometry_weights.len() {
return Err(EstimationError::RemlOptimizationFailed(format!(
"reactive scalar waypoint isometry dimension {} != literal target dimension {}",
state.isometry_weights.len(),
self.baseline_isometry_weights.len(),
)));
}
self.fit_verdict = None;
self.term
.assignment
.mode
.set_temperature(state.assignment_temperature)
.map_err(EstimationError::RemlOptimizationFailed)?;
let installing_entry = state.bitwise_eq(contract.entry());
let restoring_target = state.bitwise_eq(contract.target());
self.term.temperature_schedule = None;
if let Some(registry) = self.registry.as_mut() {
registry.set_isometry_scalar_weights(&state.isometry_weights);
}
if installing_entry {
if self.reactive_waypoint_checkpoint.is_none() {
return Err(EstimationError::RemlOptimizationFailed(
"reactive scalar entry placement requires an active full-state waypoint transaction"
.to_string(),
));
}
self.term
.place_reactive_entry_disjoint_charts(self.target.view())
.map_err(|err| {
EstimationError::RemlOptimizationFailed(format!(
"reactive scalar entry could not install its separated legal basin: {err}"
))
})?;
let entry_rho_flat = reactive_rho_domain_upper(
&self.term,
&self.baseline_rho,
state.assignment_temperature,
)
.map_err(EstimationError::RemlOptimizationFailed)?;
let entry_rho = self
.baseline_rho
.from_flat(entry_rho_flat.view())
.map_err(EstimationError::InvalidInput)?;
self.term
.refit_reactive_entry_decoders_at_smooth_face(
self.target.view(),
&entry_rho,
)
.map_err(|err| {
EstimationError::RemlOptimizationFailed(format!(
"reactive scalar entry could not fit its separated decoders at the legal smooth face: {err}"
))
})?;
}
self.probe_converged_handoff = None;
self.basin_bundle.clear();
self.probe_telemetry.reactive_scalar_installs += 1;
if restoring_target {
self.probe_telemetry.reactive_target_restores += 1;
}
Ok(())
}
fn begin_reactive_domain_waypoint(&mut self) -> Result<(), EstimationError> {
if self.reactive_waypoint_checkpoint.is_some() {
return Err(EstimationError::RemlOptimizationFailed(
"reactive coupled waypoint began while another waypoint transaction was active"
.to_string(),
));
}
let bundle_capacity = self.basin_bundle.member_capacity();
let basin_bundle =
std::mem::replace(&mut self.basin_bundle, BasinBundle::new(bundle_capacity));
let registry_isometry_weights = self
.registry
.as_ref()
.map(AnalyticPenaltyRegistry::isometry_scalar_weights)
.unwrap_or_default();
self.reactive_waypoint_checkpoint = Some(ReactiveWaypointCheckpoint {
term: self.term.clone(),
target: self.target.clone(),
registry_isometry_weights,
current_rho: self.current_rho.clone(),
last_loss: self.last_loss.clone(),
terminal_penalized_quasi_laplace_criterion: self
.terminal_penalized_quasi_laplace_criterion,
seeded_beta: self.seeded_beta.clone(),
probe_converged_handoff: self.probe_converged_handoff.take(),
basin_bundle,
termination: self.termination.clone(),
fit_verdict: self.fit_verdict,
crosscoder_blocks: self.crosscoder_blocks.clone(),
});
Ok(())
}
fn commit_reactive_domain_waypoint(
&mut self,
rho: &Array1<f64>,
) -> Result<(), EstimationError> {
if self.reactive_waypoint_checkpoint.is_none() {
return Err(EstimationError::RemlOptimizationFailed(
"reactive coupled waypoint commit had no active transaction".to_string(),
));
}
let converged_term = self
.take_probe_converged_handoff(rho.view())
.ok_or_else(|| {
EstimationError::RemlOptimizationFailed(
"reactive coupled waypoint produced no exact-rho converged full-state handoff"
.to_string(),
)
})?;
let rho_state = self
.baseline_rho
.from_flat(rho.view())
.map_err(EstimationError::InvalidInput)?;
let target_contract = self.reactive_domain_scalar_contract()?.ok_or_else(|| {
EstimationError::RemlOptimizationFailed(
"active reactive waypoint lost its scalar contract before commit".to_string(),
)
})?;
let committed_isometry_weights = self
.registry
.as_ref()
.map(AnalyticPenaltyRegistry::isometry_scalar_weights)
.unwrap_or_default();
let committed_scalar = gam_solve::continuation_path::ContinuationScalarState::new(
converged_term.assignment.mode.temperature(),
committed_isometry_weights,
)
.map_err(EstimationError::RemlOptimizationFailed)?;
let committed_literal_target = committed_scalar.bitwise_eq(target_contract.target());
let loss = converged_term
.loss(self.target.view(), &rho_state)
.map_err(EstimationError::RemlOptimizationFailed)?;
self.term = converged_term;
if committed_literal_target {
self.term.temperature_schedule = self.baseline_term.temperature_schedule.clone();
self.term
.assignment
.mode
.set_temperature(target_contract.target().assignment_temperature)
.map_err(EstimationError::RemlOptimizationFailed)?;
}
self.current_rho = rho_state;
self.last_loss = Some(loss);
self.seeded_beta = None;
self.fit_verdict = None;
self.terminal_penalized_quasi_laplace_criterion = None;
self.reactive_waypoint_checkpoint = None;
Ok(())
}
fn rollback_reactive_domain_waypoint(&mut self) -> Result<(), EstimationError> {
let checkpoint = self.reactive_waypoint_checkpoint.take().ok_or_else(|| {
EstimationError::RemlOptimizationFailed(
"reactive coupled waypoint rollback had no active transaction".to_string(),
)
})?;
self.term = checkpoint.term;
self.target = checkpoint.target;
if let Some(registry) = self.registry.as_mut() {
registry.set_isometry_scalar_weights(&checkpoint.registry_isometry_weights);
}
self.current_rho = checkpoint.current_rho;
self.last_loss = checkpoint.last_loss;
self.terminal_penalized_quasi_laplace_criterion =
checkpoint.terminal_penalized_quasi_laplace_criterion;
self.seeded_beta = checkpoint.seeded_beta;
self.probe_converged_handoff = checkpoint.probe_converged_handoff;
self.basin_bundle = checkpoint.basin_bundle;
self.termination = checkpoint.termination;
self.fit_verdict = checkpoint.fit_verdict;
self.crosscoder_blocks = checkpoint.crosscoder_blocks;
Ok(())
}
}
pub(crate) fn sae_manifold_newton_directional_decrease(
sys: &ArrowSchurSystem,
delta_ext_coord: ArrayView1<'_, f64>,
delta_beta: ArrayView1<'_, f64>,
) -> f64 {
assert_eq!(delta_ext_coord.len(), sys.row_offsets[sys.rows.len()]);
assert_eq!(delta_beta.len(), sys.k);
let mut gradient_dot_step = 0.0;
for (row_idx, row) in sys.rows.iter().enumerate() {
let row_base = sys.row_offsets[row_idx];
let di = sys.row_dims[row_idx];
for axis in 0..di {
gradient_dot_step += row.gt[axis] * delta_ext_coord[row_base + axis];
}
}
for idx in 0..sys.k {
gradient_dot_step += sys.gb[idx] * delta_beta[idx];
}
-gradient_dot_step
}
pub(crate) fn batched_smooth_sb(
sb_inputs: &[(ArrayView2<'_, f64>, ArrayView2<'_, f64>)],
symmetrize: bool,
gpu_policy: gam_gpu::GpuPolicy,
) -> Result<Vec<Array2<f64>>, String> {
let n_atoms = sb_inputs.len();
let s_mats: Vec<Array2<f64>> = sb_inputs
.iter()
.map(|(s, _)| {
if symmetrize {
let m = s.nrows();
let mut sym = Array2::<f64>::zeros((m, m));
for i in 0..m {
for j in 0..m {
sym[[i, j]] = 0.5 * (s[[i, j]] + s[[j, i]]);
}
}
sym
} else {
s.to_owned()
}
})
.collect();
let cpu_one = |idx: usize| -> Array2<f64> { s_mats[idx].dot(&sb_inputs[idx].1) };
let mut groups: std::collections::BTreeMap<(usize, usize), Vec<usize>> =
std::collections::BTreeMap::new();
for (idx, (_, b)) in sb_inputs.iter().enumerate() {
let m = s_mats[idx].nrows();
let p = b.ncols();
groups.entry((m, p)).or_default().push(idx);
}
let group_op =
|atoms: usize, m: usize, p: usize| crate::gpu::linalg_dispatch::DispatchOp::BatchedGemm {
batch: atoms,
m,
n: p,
k: m,
};
if !groups
.iter()
.any(|(&(m, p), members)| group_op(members.len(), m, p).admissible_under_any_policy())
{
return Ok((0..n_atoms).map(cpu_one).collect());
}
let rt = match crate::gpu::device_runtime::GpuRuntime::resolve(gpu_policy)
.map_err(|error| format!("decoder-smoothness CUDA admission failed: {error}"))?
{
Some(rt) => rt,
None => return Ok((0..n_atoms).map(cpu_one).collect()),
};
let mut out: Vec<Option<Array2<f64>>> = (0..n_atoms).map(|_| None).collect();
for ((m, p), members) in groups {
if members.len() < 2
|| m == 0
|| p == 0
|| crate::gpu::linalg_dispatch::route_through_gpu_with_policy(
group_op(members.len(), m, p),
gpu_policy,
)
.is_none()
{
for &idx in &members {
out[idx] = Some(cpu_one(idx));
}
continue;
}
let mut items: Vec<usize> = members.clone();
let s_ref = &s_mats;
let tile_results: std::sync::Mutex<Vec<(usize, Array2<f64>)>> =
std::sync::Mutex::new(Vec::with_capacity(members.len()));
let ok = crate::gpu::pool::scatter_batched(rt, &mut items, |_, slice| {
if slice.is_empty() {
return Some(());
}
let batch = slice.len();
let mut a = Array3::<f64>::zeros((batch, m, m));
let mut bt = Array3::<f64>::zeros((batch, p, m));
for (t, &idx) in slice.iter().enumerate() {
let s = &s_ref[idx];
let b = &sb_inputs[idx].1;
for i in 0..m {
for j in 0..m {
a[[t, i, j]] = s[[i, j]];
}
}
for i in 0..p {
for j in 0..m {
bt[[t, i, j]] = b[[j, i]];
}
}
}
let prod = crate::gpu::try_fast_abt_strided_batched_with_policy(
a.view(),
bt.view(),
gpu_policy,
)?;
let mut sink = tile_results.lock().expect("tile_results mutex poisoned");
for (t, &idx) in slice.iter().enumerate() {
sink.push((idx, prod.slice(s![t, .., ..]).to_owned()));
}
Some(())
});
match ok {
Some(()) => {
let sink = tile_results
.into_inner()
.expect("tile_results mutex poisoned");
for (idx, mat) in sink {
out[idx] = Some(mat);
}
for &idx in &members {
if out[idx].is_none() {
return Err(format!(
"decoder-smoothness device scatter omitted atom {idx}"
));
}
}
}
None => {
return Err(format!(
"decoder-smoothness device scatter declined admitted group m={m}, p={p}, atoms={}",
members.len()
));
}
}
}
out.into_iter()
.enumerate()
.map(|(idx, slot)| {
slot.ok_or_else(|| format!("decoder-smoothness result missing atom {idx}"))
})
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CurvatureBifurcation {
pub eta: f64,
pub min_pivot: f64,
}
#[derive(Debug, Clone)]
pub struct CurvatureWalkReport {
pub arrived: bool,
pub anchor_residual_norm_sq: f64,
pub bifurcation: Option<CurvatureBifurcation>,
pub eta_steps: usize,
pub step_halvings: usize,
pub collapse_events: usize,
pub reseeds: usize,
}
#[derive(Debug, Clone)]
pub struct LinearSpanAtomAnchor {
pub gate_weight: f64,
pub frame: GrassmannFrame,
pub decoder_coordinates: Array2<f64>,
pub singular_values: Array1<f64>,
}
#[derive(Debug, Clone)]
pub struct LinearSpanAnchor {
pub atoms: Vec<LinearSpanAtomAnchor>,
pub reconstruction: Array2<f64>,
pub residual_norm_sq: f64,
}
pub(crate) fn sae_cholesky_solve_neg_gradient(
h: ArrayView2<'_, f64>,
g: ArrayView1<'_, f64>,
) -> Result<Array1<f64>, String> {
let n = h.nrows();
if h.ncols() != n || g.len() != n {
return Err(format!(
"sae_cholesky_solve_neg_gradient: shape mismatch H={:?}, g={}",
h.dim(),
g.len()
));
}
let mut l = Array2::<f64>::zeros((n, n));
for i in 0..n {
for j in 0..=i {
let mut sum = h[[i, j]];
for k in 0..j {
sum -= l[[i, k]] * l[[j, k]];
}
if i == j {
if !(sum.is_finite() && sum > 0.0) {
return Err(format!("non-positive Cholesky pivot at {i}: {sum}"));
}
l[[i, j]] = sum.sqrt();
} else {
l[[i, j]] = sum / l[[j, j]];
}
}
}
let mut y = Array1::<f64>::zeros(n);
for i in 0..n {
let mut sum = -g[i];
for k in 0..i {
sum -= l[[i, k]] * y[k];
}
y[i] = sum / l[[i, i]];
}
let mut x = Array1::<f64>::zeros(n);
for ii in 0..n {
let i = n - 1 - ii;
let mut sum = y[i];
for k in i + 1..n {
sum -= l[[k, i]] * x[k];
}
x[i] = sum / l[[i, i]];
}
if !x.iter().all(|v| v.is_finite()) {
return Err("sae_cholesky_solve_neg_gradient: non-finite solution".into());
}
Ok(x)
}
pub(crate) fn solve_basis_transport(
new_phi: ArrayView2<'_, f64>,
old_phi: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
solve_design_least_squares(new_phi, old_phi)
}
pub(crate) fn transport_smooth_penalty_for_decoder(
decoder_transport: ArrayView2<'_, f64>,
old_smooth_penalty: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
let m = decoder_transport.nrows();
if decoder_transport.ncols() != m {
return Err(format!(
"transport_smooth_penalty_for_decoder: decoder transport must be square; got {:?}",
decoder_transport.dim()
));
}
if old_smooth_penalty.dim() != (m, m) {
return Err(format!(
"transport_smooth_penalty_for_decoder: smooth penalty shape {:?} != ({m}, {m})",
old_smooth_penalty.dim()
));
}
let transport_inverse =
solve_design_least_squares(decoder_transport, Array2::<f64>::eye(m).view())?;
Ok(fast_atb(
&transport_inverse,
&fast_ab(&old_smooth_penalty.to_owned(), &transport_inverse),
))
}
pub(crate) fn solve_design_least_squares(
design: ArrayView2<'_, f64>,
rhs: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
if design.nrows() != rhs.nrows() {
return Err(format!(
"solve_design_least_squares: row mismatch design={} rhs={}",
design.nrows(),
rhs.nrows()
));
}
let (u_opt, sigma, vt_opt) = design
.to_owned()
.svd(true, true)
.map_err(|err| format!("solve_design_least_squares: SVD failed: {err}"))?;
let u = u_opt.ok_or_else(|| "solve_design_least_squares: SVD omitted U".to_string())?;
let vt = vt_opt.ok_or_else(|| "solve_design_least_squares: SVD omitted Vt".to_string())?;
let smax = sigma.iter().fold(0.0_f64, |acc, &v| acc.max(v));
if !(smax.is_finite() && smax > 0.0) {
return Err("solve_design_least_squares: design has zero numerical rank".to_string());
}
let cutoff = smax * f64::EPSILON * (design.nrows().max(design.ncols()) as f64);
let coeffs = u.t().dot(&rhs);
let mut scaled = Array2::<f64>::zeros(coeffs.dim());
for row in 0..sigma.len() {
if sigma[row] > cutoff {
let inv = 1.0 / sigma[row];
for col in 0..rhs.ncols() {
scaled[[row, col]] = inv * coeffs[[row, col]];
}
}
}
Ok(vt.t().dot(&scaled))
}
#[cfg(test)]
mod linear_parity_anchor_1026_tests {
use super::*;
fn ls_projection_ev(phi: ArrayView2<'_, f64>, target: ArrayView2<'_, f64>) -> f64 {
let m = phi.ncols();
let gram = phi.t().dot(&phi) + Array2::<f64>::eye(m) * 1.0e-10;
let rhs = phi.t().dot(&target);
let coeffs = gam_linalg::faer_ndarray::FaerCholesky::cholesky(&gram, faer::Side::Lower)
.map(|c| c.solve_mat(&rhs))
.expect("design Gram must be SPD");
let fitted = phi.dot(&coeffs);
reconstruction_explained_variance(target, fitted.view()).expect("projection EV finite")
}
#[test]
fn hybrid_curved_plus_linear_beats_either_alone_1026() {
let n = 80usize;
let p = 5usize;
let zf: Vec<f64> = (0..n).map(|i| ((i as f64 + 1.0) * 0.21).sin()).collect();
let theta: Vec<f64> = (0..n).map(|i| ((i as f64) * 0.6180339887) % 1.0).collect();
let a0 = Array1::from_shape_fn(p, |c| 0.5 + 0.3 * (c as f64));
let a1 = Array1::from_shape_fn(p, |c| (((c + 1) % 4) as f64 - 1.5) * 0.8);
let bs = Array1::from_shape_fn(p, |c| (((c * 2 + 1) % 5) as f64 - 2.0) * 0.9);
let bc = Array1::from_shape_fn(p, |c| (((c * 3 + 2) % 5) as f64 - 2.0) * 0.7);
let two_pi = std::f64::consts::TAU;
let target = Array2::from_shape_fn((n, p), |(i, c)| {
a0[c]
+ a1[c] * zf[i]
+ bs[c] * (two_pi * theta[i]).sin()
+ bc[c] * (two_pi * theta[i]).cos()
});
let mut phi_lin = Array2::<f64>::ones((n, 2));
for i in 0..n {
phi_lin[[i, 1]] = zf[i];
}
let eval = PeriodicHarmonicEvaluator::new(3).unwrap();
let theta_coords = Array2::from_shape_fn((n, 1), |(i, _)| theta[i]);
let (phi_curved, _jet) = eval.evaluate(theta_coords.view()).unwrap();
let mut phi_hybrid = Array2::<f64>::ones((n, 4));
for i in 0..n {
phi_hybrid[[i, 1]] = zf[i];
phi_hybrid[[i, 2]] = phi_curved[[i, 1]]; phi_hybrid[[i, 3]] = phi_curved[[i, 2]]; }
let ev_lin = ls_projection_ev(phi_lin.view(), target.view());
let ev_curved = ls_projection_ev(phi_curved.view(), target.view());
let ev_hybrid = ls_projection_ev(phi_hybrid.view(), target.view());
println!(
"[#1026] hybrid dominance: linear-only EV={ev_lin:.6} curved-only EV={ev_curved:.6} \
hybrid EV={ev_hybrid:.6} hybrid−max(either)={:.6}",
ev_hybrid - ev_lin.max(ev_curved)
);
assert!(
ev_lin.is_finite() && ev_curved.is_finite() && ev_hybrid.is_finite(),
"all three projection EVs must be finite: lin={ev_lin}, curved={ev_curved}, \
hybrid={ev_hybrid}"
);
assert!(
ev_hybrid > ev_lin + 0.05,
"#1026 hybrid: union basis EV {ev_hybrid:.6} must STRICTLY beat linear-only \
{ev_lin:.6} (the curved atom captures the periodic wave a line cannot)"
);
assert!(
ev_hybrid > ev_curved + 0.05,
"#1026 hybrid: union basis EV {ev_hybrid:.6} must STRICTLY beat curved-only \
{ev_curved:.6} (the linear tier captures the z-ramp the periodic atom cannot)"
);
assert!(
ev_hybrid > 0.999,
"#1026 hybrid: the union basis is the exact generating model, so its \
projection EV must be ~1; got {ev_hybrid:.6}"
);
}
}
#[cfg(test)]
mod decoder_smoothness_dispatch_2393_tests {
use super::batched_smooth_sb;
use ndarray::Array2;
#[test]
fn llm_decoder_group_routes_instead_of_refusing() {
const M: usize = 6;
const P: usize = 2048;
const ATOMS: usize = 8;
let s_mats: Vec<Array2<f64>> = (0..ATOMS)
.map(|atom| {
Array2::from_shape_fn((M, M), |(i, j)| {
if i == j {
1.0 + 0.1 * (atom as f64)
} else {
0.01 * ((i + j + atom) as f64).sin()
}
})
})
.collect();
let b_mats: Vec<Array2<f64>> = (0..ATOMS)
.map(|atom| {
Array2::from_shape_fn((M, P), |(i, j)| {
0.05 * (((i + 1) * (j + 3) + atom) as f64 * 0.0037).cos()
})
})
.collect();
let expected: Vec<Array2<f64>> = (0..ATOMS)
.map(|atom| s_mats[atom].dot(&b_mats[atom]))
.collect();
for policy in [crate::gpu::GpuPolicy::Auto, crate::gpu::GpuPolicy::Off] {
let inputs: Vec<_> = (0..ATOMS)
.map(|atom| (s_mats[atom].view(), b_mats[atom].view()))
.collect();
let got = batched_smooth_sb(&inputs, false, policy).unwrap_or_else(|error| {
panic!(
"#2393: the LLM decoder group must route to a product under \
{policy}, not refuse; got error: {error}"
)
});
assert_eq!(got.len(), ATOMS);
for (atom, product) in got.iter().enumerate() {
assert_eq!(
product, &expected[atom],
"#2393: atom {atom} product differs from the exact S·B under {policy}"
);
}
}
}
}
#[cfg(test)]
mod probe_refusal_classification_2593_tests {
use super::{ArrowSchurError, OuterProbeTelemetry, ProbeRefusalKind, SaeManifoldOuterObjective};
fn representative(kind: ProbeRefusalKind) -> String {
match kind {
ProbeRefusalKind::InnerNotConverged => {
"SaeManifoldTerm::penalized_quasi_laplace_criterion: inner solve did not \
converge at fixed ρ; refusing to rank an off-optimum state"
.to_string()
}
ProbeRefusalKind::NonPdPerRow => {
"SaeManifoldTerm::penalized_quasi_laplace_criterion: undamped evidence \
factorization hit a non-PD per-row H_tt block before KKT stationarity \
at an infeasible-ρ probe"
.to_string()
}
ProbeRefusalKind::NonPdSchur => ArrowSchurError::SchurFactorFailed {
reason: "leading minor is not positive definite".to_string(),
}
.to_string(),
ProbeRefusalKind::AllZeroGatedDesign => {
"run_joint_fit_arrow_schur: atom 2 is gated off at every row (all-zero \
gated design)"
.to_string()
}
ProbeRefusalKind::TotalCoCollapse => {
"run_joint_fit_arrow_schur: reseed budget spent and the fit did not \
escape total co-collapse"
.to_string()
}
}
}
#[test]
fn every_refusal_kind_is_counted_exactly_once() {
for kind in ProbeRefusalKind::ALL {
let message = representative(kind);
assert_eq!(
ProbeRefusalKind::classify(&message),
Some(kind),
"representative message must classify as its own kind: {message}"
);
assert!(
SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(&message),
"a classified refusal is ρ-local by construction: {message}"
);
let mut telemetry = OuterProbeTelemetry::default();
telemetry.record_refusal_kind(&message);
assert_eq!(
telemetry.infeasible_total(),
1,
"{kind:?} must increment exactly one counter that infeasible_total sums"
);
}
}
#[test]
fn an_unclassified_defect_is_fatal_and_uncounted() {
let defect = "SaeManifoldTerm::penalized_quasi_laplace_criterion: \
arrow_log_det_from_cache returned None (undamped joint Hessian \
log-det unavailable for the Laplace normaliser)";
assert_eq!(ProbeRefusalKind::classify(defect), None);
assert!(!SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
defect
));
let mut telemetry = OuterProbeTelemetry::default();
telemetry.record_refusal_kind(defect);
assert_eq!(telemetry.infeasible_total(), 0);
}
#[test]
fn every_per_row_producer_rendering_classifies_as_non_pd_per_row() {
let renderings = [
"SaeManifoldTerm::penalized_quasi_laplace_criterion: stationary undamped \
criterion factorization has a non-PD per-row H_tt block that spectral \
unit-stiffness deflation could not condition",
"SaeManifoldTerm::penalized_quasi_laplace_criterion: undamped evidence \
factorization hit a non-PD per-row H_tt block before KKT stationarity \
at an infeasible-ρ probe; returning the typed infeasible refusal \
without grinding the probe refinement budget",
"SaeManifoldTerm::penalized_quasi_laplace_criterion: undamped evidence \
factorization hit a non-PD per-row H_tt block before KKT stationarity \
and the refinement budget was exhausted",
];
for rendering in renderings {
assert_eq!(
ProbeRefusalKind::classify(rendering),
Some(ProbeRefusalKind::NonPdPerRow),
"a per-row refusal the crate actually renders must classify as \
NonPdPerRow: {rendering}"
);
assert!(
SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(rendering),
"the outer optimizer must read this ρ as +∞ and steer, not abort: \
{rendering}"
);
let mut telemetry = OuterProbeTelemetry::default();
telemetry.record_refusal_kind(rendering);
assert_eq!(
telemetry.infeasible_non_pd_per_row, 1,
"infeasible_non_pd_per_row counted zero of these before #2598: {rendering}"
);
}
}
#[test]
fn a_schur_failure_that_is_not_a_non_pd_pivot_stays_fatal() {
let defect = "arrow-Schur: Schur complement Cholesky failed: non-finite entry";
assert_eq!(ProbeRefusalKind::classify(defect), None);
assert!(!SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
defect
));
}
}