use ndarray::{Array1, Array2, ArrayView2};
use gam_solve::inference::residual_factor::{ResidualFactorInput, StructuredResidualModel};
use gam_solve::structure_search::StructureMove;
use crate::structure_harvest::apply_structure_move;
use super::*;
#[derive(Clone, Copy, Debug)]
pub struct StagewiseConfig {
pub inner_max_iter: usize,
pub learning_rate: f64,
pub ridge_ext_coord: f64,
pub ridge_beta: f64,
pub max_births: usize,
pub max_backfit_sweeps: usize,
pub min_effect_ev: f64,
pub max_factor_rank: usize,
pub structured_whitening: bool,
}
impl Default for StagewiseConfig {
fn default() -> Self {
Self {
inner_max_iter: 64,
learning_rate: 1.0,
ridge_ext_coord: 1e-6,
ridge_beta: 1e-6,
max_births: 32,
max_backfit_sweeps: 4,
min_effect_ev: 0.0,
max_factor_rank: 4,
structured_whitening: true,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BirthKind {
NewAtom,
ChartExtension,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StagewiseStop {
TwoConsecutiveRejections,
MaxBirths,
NoResidualStructure,
Cancelled,
}
#[derive(Clone, Debug, PartialEq)]
pub struct BirthRecord {
pub joint_penalized_quasi_laplace_before: f64,
pub min_effect_ev: f64,
pub candidates: Vec<BirthCandidateRecord>,
}
impl BirthRecord {
pub fn accepted(&self) -> bool {
self.candidates
.iter()
.any(|candidate| matches!(&candidate.decision, BirthCandidateDecision::Accepted))
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct BirthCandidateRecord {
pub kind: BirthKind,
pub delta_ev: Option<f64>,
pub factor_energy: f64,
pub joint_penalized_quasi_laplace: Option<f64>,
pub decision: BirthCandidateDecision,
}
#[derive(Clone, Debug, PartialEq)]
pub enum BirthCandidateDecision {
Accepted,
GateRejected(BirthRejection),
Outranked,
DeferredByBatchSelection,
FitFailed(String),
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum BirthRejection {
NonFiniteCriterion,
NonFiniteEv,
EvidenceNotImproved {
criterion: f64,
must_be_below: f64,
},
EffectBelowFloor {
delta_ev: f64,
floor: f64,
},
}
fn classify_birth_candidate(
criterion: f64,
ev: f64,
cur_ev: f64,
current_criterion: f64,
min_effect_ev: f64,
) -> Option<BirthRejection> {
if !criterion.is_finite() {
return Some(BirthRejection::NonFiniteCriterion);
}
if !(criterion < current_criterion) {
return Some(BirthRejection::EvidenceNotImproved {
criterion,
must_be_below: current_criterion,
});
}
if !ev.is_finite() {
return Some(BirthRejection::NonFiniteEv);
}
let delta_ev = ev - cur_ev;
if !(delta_ev >= min_effect_ev) {
return Some(BirthRejection::EffectBelowFloor {
delta_ev,
floor: min_effect_ev,
});
}
None
}
enum BirthCandidateAttempt {
Measured {
kind: BirthKind,
factor_energy: f64,
criterion: f64,
ev: f64,
},
FitFailed {
kind: BirthKind,
factor_energy: f64,
error: String,
},
}
#[derive(Clone, Copy)]
enum PassingCandidateDisposition {
Outranked,
DeferredByBatchSelection,
}
fn best_passing_birth_candidate(
attempts: &[BirthCandidateAttempt],
cur_ev: f64,
current_criterion: f64,
min_effect_ev: f64,
) -> Option<usize> {
attempts
.iter()
.enumerate()
.filter_map(|(index, attempt)| match attempt {
BirthCandidateAttempt::Measured { criterion, ev, .. }
if classify_birth_candidate(
*criterion,
*ev,
cur_ev,
current_criterion,
min_effect_ev,
)
.is_none() =>
{
Some((index, *criterion))
}
_ => None,
})
.min_by(|(left_index, left), (right_index, right)| {
left.partial_cmp(right)
.unwrap_or(std::cmp::Ordering::Equal)
.then(left_index.cmp(right_index))
})
.map(|(index, _)| index)
}
fn record_birth_candidates(
attempts: &[BirthCandidateAttempt],
selected_indices: &[usize],
passing_disposition: PassingCandidateDisposition,
cur_ev: f64,
current_criterion: f64,
min_effect_ev: f64,
) -> Result<Vec<BirthCandidateRecord>, String> {
let mut selected = vec![false; attempts.len()];
for &index in selected_indices {
let Some(slot) = selected.get_mut(index) else {
return Err(format!(
"birth candidate selection index {index} is outside {} attempts",
attempts.len()
));
};
if *slot {
return Err(format!(
"birth candidate selection index {index} was selected twice"
));
}
*slot = true;
}
attempts
.iter()
.enumerate()
.map(|(index, attempt)| match attempt {
BirthCandidateAttempt::FitFailed {
kind,
factor_energy,
error,
} => {
if selected[index] {
return Err(format!(
"birth candidate selection index {index} names a failed fit"
));
}
Ok(BirthCandidateRecord {
kind: *kind,
delta_ev: None,
factor_energy: *factor_energy,
joint_penalized_quasi_laplace: None,
decision: BirthCandidateDecision::FitFailed(error.clone()),
})
}
BirthCandidateAttempt::Measured {
kind,
factor_energy,
criterion,
ev,
} => {
let rejection = classify_birth_candidate(
*criterion,
*ev,
cur_ev,
current_criterion,
min_effect_ev,
);
if selected[index] && rejection.is_some() {
return Err(format!(
"birth candidate selection index {index} did not clear the birth gate"
));
}
let decision = match rejection {
Some(reason) => BirthCandidateDecision::GateRejected(reason),
None if selected[index] => BirthCandidateDecision::Accepted,
None => match passing_disposition {
PassingCandidateDisposition::Outranked => {
BirthCandidateDecision::Outranked
}
PassingCandidateDisposition::DeferredByBatchSelection => {
BirthCandidateDecision::DeferredByBatchSelection
}
},
};
Ok(BirthCandidateRecord {
kind: *kind,
delta_ev: Some(*ev - cur_ev),
factor_energy: *factor_energy,
joint_penalized_quasi_laplace: Some(*criterion),
decision,
})
}
})
.collect()
}
#[derive(Clone, Debug)]
pub struct StagewiseReport {
pub births_accepted: usize,
pub births_rejected: usize,
pub birth_records: Vec<BirthRecord>,
pub ev_trace: Vec<f64>,
pub backfit_ev_trace: Vec<f64>,
pub stopped_reason: StagewiseStop,
pub terminal_joint_penalized_quasi_laplace: f64,
pub terminal_joint_loss: SaeManifoldLoss,
}
#[derive(Clone, Debug)]
pub struct StagewiseResult {
pub term: SaeManifoldTerm,
pub rho: SaeManifoldRho,
pub report: StagewiseReport,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StagewiseEventKind {
SeedReady,
BirthRoundStarted,
ResidualModelStarted,
ResidualModelFitted,
CurrentEvidenceStarted,
CurrentEvidenceFinished,
CandidateStarted,
CandidateFinished,
BirthAccepted,
BirthRejected,
BackfitSweepStarted,
BackfitSweepAccepted,
BackfitSweepRejected,
TerminalEvidenceCompleted,
}
pub struct StagewiseProgress<'a> {
pub event: StagewiseEventKind,
pub birth_round: usize,
pub backfit_sweep: usize,
pub candidate: Option<BirthKind>,
pub accepted: Option<bool>,
pub checkpoint: bool,
pub k_atoms: usize,
pub births_accepted: usize,
pub births_rejected: usize,
pub ev: Option<f64>,
pub factor_energy: Option<f64>,
pub joint_penalized_quasi_laplace_before: Option<f64>,
pub joint_penalized_quasi_laplace_after: Option<f64>,
pub terminal_joint_penalized_quasi_laplace: Option<f64>,
pub term: &'a SaeManifoldTerm,
pub rho: &'a SaeManifoldRho,
}
pub type StagewiseProgressCallback<'cb> =
dyn for<'event> FnMut(StagewiseProgress<'event>) -> Result<(), String> + 'cb;
fn emit_stagewise_progress(
progress: &mut Option<&mut StagewiseProgressCallback<'_>>,
event: StagewiseProgress<'_>,
) -> Result<(), String> {
match event.event {
StagewiseEventKind::BirthAccepted | StagewiseEventKind::BirthRejected => {
let fmt = |v: Option<f64>| v.map_or_else(|| "-".to_string(), |x| format!("{x:.4}"));
log::warn!(
"[stagewise] birth round {} {:?}: K={} accepted={} rejected={} ev={} penalized_quasi_laplace {} -> {}",
event.birth_round,
event.event,
event.k_atoms,
event.births_accepted,
event.births_rejected,
fmt(event.ev),
fmt(event.joint_penalized_quasi_laplace_before),
fmt(event.joint_penalized_quasi_laplace_after),
);
}
StagewiseEventKind::SeedReady
| StagewiseEventKind::BirthRoundStarted
| StagewiseEventKind::ResidualModelStarted
| StagewiseEventKind::ResidualModelFitted
| StagewiseEventKind::CurrentEvidenceStarted
| StagewiseEventKind::CurrentEvidenceFinished
| StagewiseEventKind::CandidateStarted
| StagewiseEventKind::CandidateFinished
| StagewiseEventKind::BackfitSweepStarted
| StagewiseEventKind::BackfitSweepAccepted
| StagewiseEventKind::BackfitSweepRejected
| StagewiseEventKind::TerminalEvidenceCompleted => {}
}
if let Some(callback) = progress.as_deref_mut() {
callback(event)?;
}
Ok(())
}
fn current_residual(
term: &SaeManifoldTerm,
target: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
let fitted = term.try_fitted()?;
Ok(&target.to_owned() - &fitted)
}
fn refresh_terminal_row_metric(
term: &mut SaeManifoldTerm,
target: ArrayView2<'_, f64>,
config: &StagewiseConfig,
) -> Result<(), String> {
if !config.structured_whitening {
return Ok(());
}
let residual = current_residual(term, target)?;
match fit_residual_covariance_on(term, residual, config) {
Ok(Some((_, model))) => term.set_row_metric(model.row_metric(target.nrows())?)?,
Ok(None) => {}
Err(err) => {
log::debug!("stagewise terminal Σ refresh skipped (degenerate final residual): {err}");
}
}
Ok(())
}
pub fn frozen_joint_penalized_quasi_laplace(
term: &mut SaeManifoldTerm,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<(f64, SaeManifoldLoss), String> {
term.assignment.validate_rho_domain(rho)?;
term.penalized_quasi_laplace_criterion(
target,
rho,
registry,
0,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)
.map_err(|error| error.to_string())
}
fn ev_of(term: &SaeManifoldTerm, target: ArrayView2<'_, f64>) -> f64 {
match term.try_fitted() {
Ok(fitted) => reconstruction_explained_variance(target, fitted.view()).unwrap_or(f64::NAN),
Err(_) => f64::NAN,
}
}
fn activity_of(term: &SaeManifoldTerm) -> Array1<f64> {
let assignments = term.assignment.assignments();
let n = assignments.nrows();
(0..n).map(|r| assignments.row(r).sum()).collect()
}
fn fit_residual_covariance_on(
term: &SaeManifoldTerm,
residual: Array2<f64>,
config: &StagewiseConfig,
) -> Result<Option<(Array2<f64>, StructuredResidualModel)>, String> {
let (n, p) = residual.dim();
if n == 0 || p < 2 {
return Ok(None);
}
let activity = activity_of(term);
let max_rank = config.max_factor_rank.min(p.saturating_sub(1)).max(1);
StructuredResidualModel::fit(ResidualFactorInput {
residuals: residual.view(),
activity: activity.view(),
max_factor_rank: max_rank,
})
.map(|model| Some((residual, model)))
.map_err(|err| format!("fit_residual_covariance: structured residual fit failed: {err}"))
}
fn birth_mining_residual(
term: &SaeManifoldTerm,
target: ArrayView2<'_, f64>,
config: &StagewiseConfig,
) -> Result<Array2<f64>, String> {
let pooled = current_residual(term, target)?;
let (n, p) = pooled.dim();
if n == 0 || p < 2 {
return Ok(pooled);
}
let k_router = (term.k_atoms() + config.max_births).max(2);
let floor = crate::routability::routability_floor(p, k_router, 1, 1.0);
let min_routable = crate::routability::minimum_routable_energy(&floor);
let all_rows: Vec<usize> = (0..n).collect();
let pooled_fraction = dominant_energy_fraction(pooled.view(), &all_rows);
if pooled_fraction >= min_routable {
return Ok(pooled);
}
match stratum_local_birth_residual(pooled.view(), &floor) {
Some(pick) => Ok(pick.masked_residual),
None => Ok(pooled),
}
}
fn fit_single_atom_response_in_place(
term: &mut SaeManifoldTerm,
rho: &mut SaeManifoldRho,
atom_idx: usize,
response: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<(), String> {
let n = term.n_obs();
let k = term.k_atoms();
if atom_idx >= k {
return Err(format!(
"fit_single_atom_response_in_place: atom {atom_idx} out of range (K={k})"
));
}
let sub_atom = term.atoms[atom_idx].clone();
let coord_block = term.assignment.coords[atom_idx].clone();
let mut sub_logits = Array2::<f64>::zeros((n, 1));
for row in 0..n {
sub_logits[[row, 0]] = term.assignment.logits[[row, atom_idx]];
}
let sub_assignment =
SaeAssignment::with_mode(sub_logits, vec![coord_block], term.assignment.mode)?;
let mut sub_term = SaeManifoldTerm::new(vec![sub_atom], sub_assignment)?;
sub_term.set_guards_enabled(false);
if let Some(w) = term.row_loss_weights().map(|w| w.to_vec()) {
sub_term.set_row_loss_weights(w)?;
}
if let Some(metric) = term.row_metric().cloned() {
sub_term.set_row_metric(metric)?;
}
if sub_term
.row_metric()
.map(|m| !m.whitens_likelihood())
.unwrap_or(true)
{
let frame_rows: Vec<usize> = (0..n).collect();
crate::manifold::activate_residual_frame(
&mut sub_term.atoms[0],
response,
&frame_rows,
&crate::manifold::InFrameCurvedConfig::default(),
)?;
}
let mut sub_rho = SaeManifoldRho::with_per_atom_smooth(
rho.log_lambda_sparse,
vec![*rho.log_lambda_smooth.get(atom_idx).unwrap_or(&0.0)],
vec![
rho.log_ard
.get(atom_idx)
.cloned()
.unwrap_or_else(|| Array1::zeros(0)),
],
)
.for_assignment(sub_term.assignment.mode);
sub_term.assignment.validate_rho_domain(&sub_rho)?;
sub_term.run_joint_fit_arrow_schur(
response,
&mut sub_rho,
registry,
config.inner_max_iter,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)?;
term.atoms[atom_idx] = sub_term.atoms[0].clone();
term.assignment.coords[atom_idx] = sub_term.assignment.coords[0].clone();
for row in 0..n {
term.assignment.logits[[row, atom_idx]] = sub_term.assignment.logits[[row, 0]];
}
if atom_idx < rho.log_lambda_smooth.len() {
rho.log_lambda_smooth[atom_idx] = sub_rho.log_lambda_smooth[0];
}
if atom_idx < rho.log_ard.len() {
rho.log_ard[atom_idx] = sub_rho.log_ard[0].clone();
}
term.assignment.frozen_logits = None;
term.last_row_layout = None;
term.last_frames_active = false;
term.border_hbb_workspace = Array2::<f64>::zeros((0, 0));
Ok(())
}
fn birth_anchor_weights(term: &SaeManifoldTerm) -> Array1<f64> {
let activity = activity_of(term);
let m_max = activity.iter().copied().fold(0.0_f64, f64::max);
if m_max > 0.0 {
activity.mapv(|m| (m_max - m).max(0.0))
} else {
Array1::ones(activity.len())
}
}
struct BirthSeed {
decoder: Array2<f64>,
energy: f64,
kind: BirthSeedKind,
}
enum BirthSeedKind {
ResidualFactor,
Circle(CircleBirthSeed),
}
struct CircleBirthSeed {
geometry: SaeAtomGeometryPlan,
coords: Array2<f64>,
gate: Vec<f64>,
}
impl BirthSeed {
fn circle(&self) -> Option<&CircleBirthSeed> {
match &self.kind {
BirthSeedKind::ResidualFactor => None,
BirthSeedKind::Circle(circle) => Some(circle),
}
}
}
fn template_circle_geometry(term: &SaeManifoldTerm) -> Option<&SaeAtomGeometryPlan> {
term.atoms
.first()
.and_then(SaeManifoldAtom::geometry_plan)
.filter(|plan| plan.kind() == &SaeAtomBasisKind::Periodic && plan.latent_dim() == 1)
}
fn template_accepts_circle_births(term: &SaeManifoldTerm) -> bool {
template_circle_geometry(term).is_some()
}
fn top_factor_birth_decoder(
term: &SaeManifoldTerm,
model: &StructuredResidualModel,
residual: ArrayView2<'_, f64>,
) -> Option<BirthSeed> {
let r = model.factor_rank();
if r == 0 {
return None;
}
let factor = model.factor(); let p = factor.nrows();
let (n, p_res) = residual.dim();
if p_res != p || n == 0 {
return None;
}
if let Some(circle) = residual_principal_birth_candidate(term, residual) {
if circle.circle().is_some() {
return Some(circle);
}
}
let anchor_w = birth_anchor_weights(term);
let anchor_total: f64 = anchor_w.iter().sum();
let use_anchor = anchor_total > 0.0;
let mut best_j = 0usize;
let mut best_score = f64::NEG_INFINITY;
if use_anchor {
for j in 0..r {
let col = factor.column(j);
let energy: f64 = col.iter().map(|v| v * v).sum();
if !(energy > 0.0) {
continue;
}
let inv_norm = 1.0 / energy.sqrt();
let mut num = 0.0_f64; let mut den = 0.0_f64; for i in 0..n {
let mut proj = 0.0_f64;
for out in 0..p {
proj += residual[[i, out]] * col[out];
}
let s = (proj * inv_norm) * (proj * inv_norm);
num += anchor_w[i] * s;
den += s;
}
if den <= 0.0 {
continue;
}
let score = num / den;
if score > best_score {
best_score = score;
best_j = j;
}
}
}
let chosen = if use_anchor { best_j } else { 0 };
let energy: f64 = factor.column(chosen).iter().map(|v| v * v).sum();
if !(energy > 0.0) {
return None;
}
let m = term.atoms[0].basis_size();
let mut decoder = Array2::<f64>::zeros((m, p));
for out in 0..p {
decoder[[0, out]] = factor[[out, chosen]];
}
Some(BirthSeed {
decoder,
energy,
kind: BirthSeedKind::ResidualFactor,
})
}
fn residual_principal_birth_candidate(
term: &SaeManifoldTerm,
residual: ArrayView2<'_, f64>,
) -> Option<BirthSeed> {
let (n, p) = residual.dim();
if n < 2 || p == 0 || term.atoms.is_empty() {
return None;
}
let parts = isa_eigen_parts(residual).ok()??;
let anchor_w = birth_anchor_weights(term);
let mut best = parts.above[0];
if anchor_w.iter().sum::<f64>() > 0.0 {
let mut best_score = f64::NEG_INFINITY;
for &k in &parts.above {
let col = parts.evecs.column(k); let mut num = 0.0_f64;
let mut den = 0.0_f64;
for i in 0..n {
let mut proj = 0.0_f64;
for j in 0..p {
proj += residual[[i, j]] * col[j];
}
let si = proj * proj;
num += anchor_w[i] * si;
den += si;
}
if den > 0.0 {
let score = num / den;
if score > best_score {
best_score = score;
best = k;
}
}
}
}
let energy = parts.evals[best].max(0.0);
if !(energy > 0.0) {
return None;
}
let m = term.atoms[0].basis_size();
if template_accepts_circle_births(term) && parts.above.len() >= 2 {
let joint_span_parts = capture_signal_span(residual, parts.above.len())
.ok()
.flatten();
if let Some(cand) = joint_span_parts
.as_ref()
.and_then(|span| isa_extract_certified_plane(residual, span, &IsaSeedConfig::default()))
{
let mut decoder = Array2::<f64>::zeros((m, p));
for j in 0..p {
decoder[[1, j]] = cand.amplitudes[0] * cand.basis[[j, 0]];
decoder[[2, j]] = cand.amplitudes[1] * cand.basis[[j, 1]];
}
return Some(BirthSeed {
decoder,
energy,
kind: BirthSeedKind::Circle(CircleBirthSeed {
geometry: template_circle_geometry(term)?.clone(),
coords: cand.phases_turns,
gate: cand.gate_logits,
}),
});
}
}
let amp = energy.sqrt();
let mut decoder = Array2::<f64>::zeros((m, p));
for j in 0..p {
decoder[[0, j]] = amp * parts.evecs[[j, best]];
}
Some(BirthSeed {
decoder,
energy,
kind: BirthSeedKind::ResidualFactor,
})
}
fn isa_birth_seed_batch(
term: &SaeManifoldTerm,
residual: ArrayView2<'_, f64>,
max_planes: usize,
) -> Result<Vec<BirthSeed>, String> {
if max_planes == 0 || !template_accepts_circle_births(term) {
return Ok(Vec::new());
}
let harvest = isa_deflationary_producer(residual, max_planes, &IsaSeedConfig::default())?;
harvest
.planes
.iter()
.map(|cand| plane_to_birth_seed(term, cand))
.collect()
}
fn refit_single_atom_in_place(
term: &mut SaeManifoldTerm,
rho: &SaeManifoldRho,
atom_idx: usize,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<(), String> {
term.assignment.validate_rho_domain(rho)?;
let n = term.n_obs();
let p = term.output_dim();
let k = term.k_atoms();
if atom_idx >= k {
return Err(format!(
"refit_single_atom_in_place: atom {atom_idx} out of range (K={k})"
));
}
let full = term.try_fitted_for_rho(rho)?;
let mut e_k = &target.to_owned() - &full;
let mut g_buf = vec![0.0_f64; p];
for row in 0..n {
let weights = term.assignment.try_assignments_row(row)?;
let a_k = weights[atom_idx];
if a_k == 0.0 {
continue;
}
term.atoms[atom_idx].fill_decoded_row(row, &mut g_buf);
let mut e_row = e_k.row_mut(row);
for out in 0..p {
e_row[out] += a_k * g_buf[out];
}
}
let mut rho_scratch = rho.clone();
fit_single_atom_response_in_place(
term,
&mut rho_scratch,
atom_idx,
e_k.view(),
registry,
config,
)
}
fn backfit_sweep(
term: &mut SaeManifoldTerm,
rho: &mut SaeManifoldRho,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<(), String> {
term.assignment.validate_rho_domain(rho)?;
term.set_guards_enabled(false);
if let Err(err) = term.run_fixed_decoder_arrow_schur(
target,
rho,
registry,
1,
config.learning_rate,
config.ridge_ext_coord,
) {
log::debug!("stagewise routing step declined; keeping the last good iterate: {err}");
}
term.run_joint_fit_arrow_schur(
target,
rho,
registry,
config.inner_max_iter,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)?;
Ok(())
}
pub fn fit_stagewise(
seed: SaeManifoldTerm,
mut rho: SaeManifoldRho,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
sample_weights: Option<&[f64]>,
config: &StagewiseConfig,
mut progress: Option<&mut StagewiseProgressCallback<'_>>,
cancel: Option<&std::sync::atomic::AtomicBool>,
) -> Result<StagewiseResult, String> {
rho = rho.for_assignment(seed.assignment.mode);
seed.assignment.validate_rho_domain(&rho)?;
let n = target.nrows();
if seed.k_atoms() != 1 {
return Err(format!(
"fit_stagewise: seed must be a single-atom (K=1) term; got K={}",
seed.k_atoms()
));
}
if seed.n_obs() != n {
return Err(format!(
"fit_stagewise: seed n_obs {} != target rows {n}",
seed.n_obs()
));
}
let mut term = seed;
term.set_guards_enabled(false);
if let Some(w) = sample_weights {
if w.len() != n {
return Err(format!(
"fit_stagewise: sample_weights length {} != target rows {n}",
w.len()
));
}
term.set_row_loss_weights(w.to_vec())?;
}
let mut ev_trace = vec![ev_of(&term, target)];
let mut birth_records: Vec<BirthRecord> = Vec::new();
let mut births_accepted = 0usize;
let mut births_rejected = 0usize;
let mut consecutive_rejections = 0usize;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::SeedReady,
birth_round: 0,
backfit_sweep: 0,
candidate: None,
accepted: Some(true),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: ev_trace.last().copied(),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let mut birth_round = 0usize;
let stopped_reason = loop {
if cancel.is_some_and(|c| c.load(std::sync::atomic::Ordering::Relaxed)) {
break StagewiseStop::Cancelled;
}
if births_accepted >= config.max_births {
break StagewiseStop::MaxBirths;
}
if consecutive_rejections >= 2 {
break StagewiseStop::TwoConsecutiveRejections;
}
let round = birth_round;
birth_round += 1;
let entry_ev = ev_of(&term, target);
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BirthRoundStarted,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: None,
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(entry_ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::ResidualModelStarted,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(entry_ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let mining_residual = birth_mining_residual(&term, target, config)?;
let Some((residual, model)) = fit_residual_covariance_on(&term, mining_residual, config)?
else {
break StagewiseStop::NoResidualStructure;
};
let seed = if let Some(seed) = isa_birth_seed_batch(&term, residual.view(), 1)?
.into_iter()
.next()
{
seed
} else {
let Some(seed) = top_factor_birth_decoder(&term, &model, residual.view())
.or_else(|| residual_principal_birth_candidate(&term, residual.view()))
else {
break StagewiseStop::NoResidualStructure;
};
seed
};
let factor_energy = seed.energy;
if config.structured_whitening {
term.set_row_metric(model.row_metric(n)?)?;
}
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::ResidualModelFitted,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(entry_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CurrentEvidenceStarted,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(entry_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let (current_penalized_quasi_laplace, _) =
frozen_joint_penalized_quasi_laplace(&mut term, target, &rho, registry, config)?;
let cur_ev = ev_of(&term, target);
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CurrentEvidenceFinished,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(cur_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateStarted,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::NewAtom),
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(cur_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let born_move = match &seed.kind {
BirthSeedKind::Circle(circle) => crate::structure_harvest::born_circle_atom(
&term,
&rho,
circle.geometry.clone(),
seed.decoder.clone(),
circle.coords.clone(),
circle.gate.clone(),
),
BirthSeedKind::ResidualFactor => apply_structure_move(
&term,
&rho,
&StructureMove::Birth { candidate: 0 },
std::slice::from_ref(&seed.decoder),
),
};
let cand_a = born_move.and_then(|(mut cand_term, mut cand_rho)| {
cand_term.set_guards_enabled(false);
let born = cand_term.k_atoms() - 1;
fit_single_atom_response_in_place(
&mut cand_term,
&mut cand_rho,
born,
residual.view(),
registry,
config,
)?;
let (penalized_quasi_laplace, _) = frozen_joint_penalized_quasi_laplace(
&mut cand_term,
target,
&cand_rho,
registry,
config,
)?;
let ev = ev_of(&cand_term, target);
Ok((cand_term, cand_rho, penalized_quasi_laplace, ev))
});
if let Ok((cand_term, cand_rho, penalized_quasi_laplace, ev)) = cand_a.as_ref() {
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateFinished,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::NewAtom),
accepted: None,
checkpoint: false,
k_atoms: cand_term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(*ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: Some(*penalized_quasi_laplace),
terminal_joint_penalized_quasi_laplace: None,
term: cand_term,
rho: cand_rho,
},
)?;
} else {
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateFinished,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::NewAtom),
accepted: Some(false),
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: None,
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
}
let cand_b = if term.k_atoms() > 1 {
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateStarted,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::ChartExtension),
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(cur_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let last = term.k_atoms() - 1;
let mut cand_term = term.clone();
let mut cand_rho = rho.clone();
let built = (|| -> Result<(SaeManifoldTerm, SaeManifoldRho, f64, f64), String> {
refit_single_atom_in_place(
&mut cand_term,
&cand_rho,
last,
target,
registry,
config,
)?;
cand_term.set_guards_enabled(false);
cand_term.run_joint_fit_arrow_schur(
target,
&mut cand_rho,
registry,
config.inner_max_iter,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)?;
let (penalized_quasi_laplace, _) = frozen_joint_penalized_quasi_laplace(
&mut cand_term,
target,
&cand_rho,
registry,
config,
)?;
let ev = ev_of(&cand_term, target);
Ok((cand_term, cand_rho, penalized_quasi_laplace, ev))
})();
if let Ok((cand_term, cand_rho, penalized_quasi_laplace, ev)) = built.as_ref() {
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateFinished,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::ChartExtension),
accepted: None,
checkpoint: false,
k_atoms: cand_term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(*ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: Some(*penalized_quasi_laplace),
terminal_joint_penalized_quasi_laplace: None,
term: cand_term,
rho: cand_rho,
},
)?;
} else {
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::CandidateFinished,
birth_round: round,
backfit_sweep: 0,
candidate: Some(BirthKind::ChartExtension),
accepted: Some(false),
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: None,
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
}
Some(built)
} else {
None
};
let mut attempts = vec![match cand_a.as_ref() {
Ok((_, _, criterion, ev)) => BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy,
criterion: *criterion,
ev: *ev,
},
Err(error) => BirthCandidateAttempt::FitFailed {
kind: BirthKind::NewAtom,
factor_energy,
error: error.clone(),
},
}];
if let Some(attempt) = cand_b.as_ref() {
attempts.push(match attempt {
Ok((_, _, criterion, ev)) => BirthCandidateAttempt::Measured {
kind: BirthKind::ChartExtension,
factor_energy,
criterion: *criterion,
ev: *ev,
},
Err(error) => BirthCandidateAttempt::FitFailed {
kind: BirthKind::ChartExtension,
factor_energy,
error: error.clone(),
},
});
}
let selected_index = best_passing_birth_candidate(
&attempts,
cur_ev,
current_penalized_quasi_laplace,
config.min_effect_ev,
);
let selected_indices = match selected_index {
Some(index) => vec![index],
None => Vec::new(),
};
let candidates = record_birth_candidates(
&attempts,
&selected_indices,
PassingCandidateDisposition::Outranked,
cur_ev,
current_penalized_quasi_laplace,
config.min_effect_ev,
)?;
birth_records.push(BirthRecord {
joint_penalized_quasi_laplace_before: current_penalized_quasi_laplace,
min_effect_ev: config.min_effect_ev,
candidates,
});
let Some(selected_index) = selected_index else {
births_rejected += 1;
consecutive_rejections += 1;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BirthRejected,
birth_round: round,
backfit_sweep: 0,
candidate: None,
accepted: Some(false),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(cur_ev),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(
current_penalized_quasi_laplace,
),
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
continue;
};
let (kind, (cand_term, cand_rho, penalized_quasi_laplace_after, ev_after)) =
match selected_index {
0 => match cand_a {
Ok(candidate) => (BirthKind::NewAtom, candidate),
Err(error) => {
return Err(format!(
"serial birth selected failed candidate A: {error}"
));
}
},
1 => match cand_b {
Some(Ok(candidate)) => (BirthKind::ChartExtension, candidate),
Some(Err(error)) => {
return Err(format!(
"serial birth selected failed candidate B: {error}"
));
}
None => {
return Err(
"serial birth selected candidate B without attempting it".to_string()
);
}
},
_ => {
return Err(format!(
"serial birth selected candidate index {selected_index} outside two arms"
));
}
};
term = cand_term;
rho = cand_rho;
births_accepted += 1;
consecutive_rejections = 0;
ev_trace.push(ev_after);
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BirthAccepted,
birth_round: round,
backfit_sweep: 0,
candidate: Some(kind),
accepted: Some(true),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(ev_after),
factor_energy: Some(factor_energy),
joint_penalized_quasi_laplace_before: Some(current_penalized_quasi_laplace),
joint_penalized_quasi_laplace_after: Some(penalized_quasi_laplace_after),
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
};
let mut backfit_ev_trace: Vec<f64> = Vec::new();
let mut prev_ev = *ev_trace.last().unwrap_or(&f64::NEG_INFINITY);
for sweep in 0..config.max_backfit_sweeps {
if cancel.is_some_and(|c| c.load(std::sync::atomic::Ordering::Relaxed)) {
break;
}
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BackfitSweepStarted,
birth_round,
backfit_sweep: sweep,
candidate: None,
accepted: None,
checkpoint: false,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(prev_ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
let term_snapshot = term.clone();
let rho_snapshot = rho.clone();
backfit_sweep(&mut term, &mut rho, target, registry, config)?;
let ev = ev_of(&term, target);
if ev > prev_ev {
backfit_ev_trace.push(ev);
prev_ev = ev;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BackfitSweepAccepted,
birth_round,
backfit_sweep: sweep,
candidate: None,
accepted: Some(true),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
} else {
term = term_snapshot;
rho = rho_snapshot;
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::BackfitSweepRejected,
birth_round,
backfit_sweep: sweep,
candidate: None,
accepted: Some(false),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(prev_ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: None,
terminal_joint_penalized_quasi_laplace: None,
term: &term,
rho: &rho,
},
)?;
break;
}
}
refresh_terminal_row_metric(&mut term, target, config)?;
let (terminal_joint_penalized_quasi_laplace, terminal_joint_loss) =
frozen_joint_penalized_quasi_laplace(&mut term, target, &rho, registry, config)?;
term.set_guards_enabled(true);
emit_stagewise_progress(
&mut progress,
StagewiseProgress {
event: StagewiseEventKind::TerminalEvidenceCompleted,
birth_round,
backfit_sweep: backfit_ev_trace.len(),
candidate: None,
accepted: Some(true),
checkpoint: true,
k_atoms: term.k_atoms(),
births_accepted,
births_rejected,
ev: Some(prev_ev),
factor_energy: None,
joint_penalized_quasi_laplace_before: None,
joint_penalized_quasi_laplace_after: Some(terminal_joint_penalized_quasi_laplace),
terminal_joint_penalized_quasi_laplace: Some(terminal_joint_penalized_quasi_laplace),
term: &term,
rho: &rho,
},
)?;
Ok(StagewiseResult {
term,
rho,
report: StagewiseReport {
births_accepted,
births_rejected,
birth_records,
ev_trace,
backfit_ev_trace,
stopped_reason,
terminal_joint_penalized_quasi_laplace,
terminal_joint_loss,
},
})
}
#[derive(Clone, Copy, Debug)]
pub struct BatchedStagewiseConfig {
pub base: StagewiseConfig,
pub max_candidates_per_round: usize,
}
impl Default for BatchedStagewiseConfig {
fn default() -> Self {
Self {
base: StagewiseConfig::default(),
max_candidates_per_round: StagewiseConfig::default().max_births,
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct BatchRoundRecord {
pub candidates_generated: usize,
pub candidates_passing_gate: usize,
pub co_accepted: usize,
pub requeued_overlap: usize,
}
#[derive(Clone, Debug)]
pub struct BatchedStagewiseResult {
pub term: SaeManifoldTerm,
pub rho: SaeManifoldRho,
pub report: StagewiseReport,
pub batch_records: Vec<BatchRoundRecord>,
}
#[derive(Clone)]
struct RacedCandidate {
born_atom: SaeManifoldAtom,
born_coord: LatentCoordValues,
born_logit_col: Vec<f64>,
born_ard: Array1<f64>,
born_log_lambda_smooth: f64,
support: Vec<usize>,
out_support: Vec<usize>,
penalized_quasi_laplace: f64,
ev: f64,
energy: f64,
}
fn plane_to_birth_seed(
term: &SaeManifoldTerm,
cand: &IsaPlaneCandidate,
) -> Result<BirthSeed, String> {
let geometry = template_circle_geometry(term).cloned().ok_or_else(|| {
"plane_to_birth_seed: certified circle requires the template atom's persisted periodic geometry plan"
.to_string()
})?;
let m = term.atoms[0].basis_size();
let p = term.output_dim();
let mut decoder = Array2::<f64>::zeros((m, p));
for j in 0..p {
decoder[[1, j]] = cand.amplitudes[0] * cand.basis[[j, 0]];
decoder[[2, j]] = cand.amplitudes[1] * cand.basis[[j, 1]];
}
let energy = cand.amplitudes[0].powi(2) + cand.amplitudes[1].powi(2);
Ok(BirthSeed {
decoder,
energy,
kind: BirthSeedKind::Circle(CircleBirthSeed {
geometry,
coords: cand.phases_turns.clone(),
gate: cand.gate_logits.clone(),
}),
})
}
fn race_birth_seed(
term: &SaeManifoldTerm,
rho: &SaeManifoldRho,
seed: &BirthSeed,
residual: ArrayView2<'_, f64>,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<RacedCandidate, String> {
let faer_seq_race_guard = gam_linalg::faer_ndarray::FaerSequentialScope::enter();
let k = term.k_atoms();
let n = term.assignment.logits.nrows();
let born_move = match &seed.kind {
BirthSeedKind::Circle(circle) => crate::structure_harvest::born_circle_atom(
term,
rho,
circle.geometry.clone(),
seed.decoder.clone(),
circle.coords.clone(),
circle.gate.clone(),
),
BirthSeedKind::ResidualFactor => apply_structure_move(
term,
rho,
&StructureMove::Birth { candidate: 0 },
std::slice::from_ref(&seed.decoder),
),
};
let (mut cand_term, mut cand_rho) = born_move?;
cand_term.set_guards_enabled(false);
fit_single_atom_response_in_place(
&mut cand_term,
&mut cand_rho,
k,
residual,
registry,
config,
)?;
let (penalized_quasi_laplace, _) =
frozen_joint_penalized_quasi_laplace(&mut cand_term, target, &cand_rho, registry, config)?;
let ev = ev_of(&cand_term, target);
let support: Vec<usize> = match &seed.kind {
BirthSeedKind::Circle(circle) => {
(0..n).filter(|&row| circle.gate[row].is_finite()).collect()
}
BirthSeedKind::ResidualFactor => (0..n).collect(),
};
let decoder = cand_term.atoms[k].decoder_coefficients();
let p_out = decoder.ncols();
let mut col_energy = vec![0.0_f64; p_out];
for row in 0..decoder.nrows() {
for j in 0..p_out {
col_energy[j] += decoder[[row, j]] * decoder[[row, j]];
}
}
let total_energy: f64 = col_energy.iter().sum();
let out_thresh = 1e-2 * total_energy;
let out_support: Vec<usize> = (0..p_out).filter(|&j| col_energy[j] > out_thresh).collect();
let born_logit_col: Vec<f64> = (0..n)
.map(|r| cand_term.assignment.logits[[r, k]])
.collect();
let raced_candidate = RacedCandidate {
born_atom: cand_term.atoms[k].clone(),
born_coord: cand_term.assignment.coords[k].clone(),
born_logit_col,
born_ard: cand_rho.log_ard[k].clone(),
born_log_lambda_smooth: cand_rho.log_lambda_smooth[k],
support,
out_support,
penalized_quasi_laplace,
ev,
energy: seed.energy,
};
drop(faer_seq_race_guard);
Ok(raced_candidate)
}
fn append_fitted_atom(
term: &SaeManifoldTerm,
rho: &SaeManifoldRho,
atom: SaeManifoldAtom,
coord: LatentCoordValues,
logit_col: &[f64],
ard: Array1<f64>,
log_lambda_smooth: f64,
) -> Result<(SaeManifoldTerm, SaeManifoldRho), String> {
let k = term.k_atoms();
let n = term.assignment.logits.nrows();
if logit_col.len() != n {
return Err(format!(
"append_fitted_atom: logit column length {} != n_obs {n}",
logit_col.len()
));
}
let mut atoms = term.atoms.clone();
atoms.push(atom);
let mut logits = Array2::<f64>::zeros((n, k + 1));
for row in 0..n {
for col in 0..k {
logits[[row, col]] = term.assignment.logits[[row, col]];
}
logits[[row, k]] = logit_col[row];
}
let mut coords = term.assignment.coords.clone();
coords.push(coord);
let assignment = SaeAssignment::with_mode(logits, coords, term.assignment.mode)?;
let child = SaeManifoldTerm::new(atoms, assignment)?;
let mut child_rho = rho.clone();
child_rho.log_ard.push(ard);
child_rho.log_lambda_smooth.push(log_lambda_smooth);
child_rho.append_curvature_atom(
k,
child.atoms[k]
.geometry_plan()
.and_then(SaeAtomGeometryPlan::constant_curvature),
)?;
Ok((child, child_rho))
}
fn select_disjoint_batch(
raced: &[RacedCandidate],
order: &[usize],
max_accept: usize,
) -> (Vec<usize>, usize) {
let intersects = |a: &[usize], b: &[usize]| -> bool {
let (small, large) = if a.len() <= b.len() { (a, b) } else { (b, a) };
let set: std::collections::HashSet<usize> = large.iter().copied().collect();
small.iter().any(|i| set.contains(i))
};
let mut accepted: Vec<usize> = Vec::new();
let mut accepted_supports: Vec<(&[usize], &[usize])> = Vec::new();
let mut requeued = 0usize;
for &idx in order {
if accepted.len() >= max_accept {
requeued += 1;
continue;
}
let c = &raced[idx];
let conflict = accepted_supports
.iter()
.any(|(ar, ad)| intersects(&c.support, ar) && intersects(&c.out_support, ad));
if conflict {
requeued += 1;
continue;
}
accepted_supports.push((&c.support, &c.out_support));
accepted.push(idx);
}
(accepted, requeued)
}
pub fn fit_stagewise_batched(
seed: SaeManifoldTerm,
mut rho: SaeManifoldRho,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
sample_weights: Option<&[f64]>,
config: &BatchedStagewiseConfig,
) -> Result<BatchedStagewiseResult, String> {
rho = rho.for_assignment(seed.assignment.mode);
seed.assignment.validate_rho_domain(&rho)?;
let n = target.nrows();
if seed.k_atoms() != 1 {
return Err(format!(
"fit_stagewise_batched: seed must be a single-atom (K=1) term; got K={}",
seed.k_atoms()
));
}
if seed.n_obs() != n {
return Err(format!(
"fit_stagewise_batched: seed n_obs {} != target rows {n}",
seed.n_obs()
));
}
let base = &config.base;
let mut term = seed;
term.set_guards_enabled(false);
if let Some(w) = sample_weights {
if w.len() != n {
return Err(format!(
"fit_stagewise_batched: sample_weights length {} != target rows {n}",
w.len()
));
}
term.set_row_loss_weights(w.to_vec())?;
}
let mut ev_trace = vec![ev_of(&term, target)];
let mut birth_records: Vec<BirthRecord> = Vec::new();
let mut batch_records: Vec<BatchRoundRecord> = Vec::new();
let mut births_accepted = 0usize;
let mut births_rejected = 0usize;
let mut consecutive_reject_rounds = 0usize;
let stopped_reason = loop {
if births_accepted >= base.max_births {
break StagewiseStop::MaxBirths;
}
if consecutive_reject_rounds >= 2 {
break StagewiseStop::TwoConsecutiveRejections;
}
let mining_residual = birth_mining_residual(&term, target, base)?;
let Some((residual, model)) = fit_residual_covariance_on(&term, mining_residual, base)?
else {
break StagewiseStop::NoResidualStructure;
};
if base.structured_whitening {
term.set_row_metric(model.row_metric(n)?)?;
}
let harvest = isa_deflationary_producer(
residual.view(),
config.max_candidates_per_round,
&IsaSeedConfig::default(),
)?;
let mut seeds: Vec<BirthSeed> = harvest
.planes
.iter()
.map(|c| plane_to_birth_seed(&term, c))
.collect::<Result<Vec<_>, _>>()?;
if seeds.is_empty() {
let Some(fallback) = top_factor_birth_decoder(&term, &model, residual.view())
.or_else(|| residual_principal_birth_candidate(&term, residual.view()))
else {
break StagewiseStop::NoResidualStructure;
};
seeds.push(fallback);
}
let candidates_generated = seeds.len();
let (current_penalized_quasi_laplace, _) =
frozen_joint_penalized_quasi_laplace(&mut term, target, &rho, registry, base)?;
let cur_ev = ev_of(&term, target);
let raced_attempts: Vec<Result<RacedCandidate, String>> = seeds
.iter()
.map(|seed| {
race_birth_seed(
&term,
&rho,
seed,
residual.view(),
target,
registry,
base,
)
})
.collect();
let mut raced = Vec::with_capacity(raced_attempts.len());
let mut source_index_by_raced = Vec::with_capacity(raced_attempts.len());
let mut attempts = Vec::with_capacity(raced_attempts.len());
for (source_index, attempt) in raced_attempts.into_iter().enumerate() {
match attempt {
Ok(candidate) => {
attempts.push(BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy: candidate.energy,
criterion: candidate.penalized_quasi_laplace,
ev: candidate.ev,
});
source_index_by_raced.push(source_index);
raced.push(candidate);
}
Err(error) => attempts.push(BirthCandidateAttempt::FitFailed {
kind: BirthKind::NewAtom,
factor_energy: seeds[source_index].energy,
error,
}),
}
}
let passes = |c: &RacedCandidate| -> bool {
classify_birth_candidate(
c.penalized_quasi_laplace,
c.ev,
cur_ev,
current_penalized_quasi_laplace,
base.min_effect_ev,
)
.is_none()
};
let mut order: Vec<usize> = (0..raced.len()).filter(|&i| passes(&raced[i])).collect();
let candidates_passing_gate = order.len();
order.sort_by(|&a, &b| {
raced[a]
.penalized_quasi_laplace
.partial_cmp(&raced[b].penalized_quasi_laplace)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
let remaining = base.max_births.saturating_sub(births_accepted);
let (accepted_idx, requeued_overlap) = select_disjoint_batch(&raced, &order, remaining);
let accepted_source_indices: Vec<usize> = accepted_idx
.iter()
.map(|&index| {
source_index_by_raced.get(index).copied().ok_or_else(|| {
format!(
"batched birth selected raced index {index} outside {} successful races",
source_index_by_raced.len()
)
})
})
.collect::<Result<_, _>>()?;
let candidates = record_birth_candidates(
&attempts,
&accepted_source_indices,
PassingCandidateDisposition::DeferredByBatchSelection,
cur_ev,
current_penalized_quasi_laplace,
base.min_effect_ev,
)?;
let mut co_accepted = 0usize;
for &idx in &accepted_idx {
let c = &raced[idx];
let (next_term, next_rho) = append_fitted_atom(
&term,
&rho,
c.born_atom.clone(),
c.born_coord.clone(),
&c.born_logit_col,
c.born_ard.clone(),
c.born_log_lambda_smooth,
)?;
term = next_term;
rho = next_rho;
births_accepted += 1;
co_accepted += 1;
let running_ev = ev_of(&term, target);
ev_trace.push(running_ev);
}
batch_records.push(BatchRoundRecord {
candidates_generated,
candidates_passing_gate,
co_accepted,
requeued_overlap,
});
birth_records.push(BirthRecord {
joint_penalized_quasi_laplace_before: current_penalized_quasi_laplace,
min_effect_ev: base.min_effect_ev,
candidates,
});
if co_accepted == 0 {
births_rejected += 1;
consecutive_reject_rounds += 1;
} else {
consecutive_reject_rounds = 0;
}
if births_accepted >= base.max_births {
break StagewiseStop::MaxBirths;
}
};
let mut prev_ev = *ev_trace.last().unwrap_or(&f64::NEG_INFINITY);
for _sweep in 0..base.max_backfit_sweeps {
let term_snapshot = term.clone();
let rho_snapshot = rho.clone();
backfit_sweep(&mut term, &mut rho, target, registry, base)?;
let ev = ev_of(&term, target);
if ev > prev_ev {
prev_ev = ev;
} else {
term = term_snapshot;
rho = rho_snapshot;
break;
}
}
refresh_terminal_row_metric(&mut term, target, base)?;
let (terminal_joint_penalized_quasi_laplace, terminal_joint_loss) =
frozen_joint_penalized_quasi_laplace(&mut term, target, &rho, registry, base)?;
term.set_guards_enabled(true);
Ok(BatchedStagewiseResult {
term,
rho,
report: StagewiseReport {
births_accepted,
births_rejected,
birth_records,
ev_trace,
backfit_ev_trace: Vec::new(),
stopped_reason,
terminal_joint_penalized_quasi_laplace,
terminal_joint_loss,
},
batch_records,
})
}
pub fn terminal_joint_assembly(
primary: SaeManifoldTerm,
primary_rho: &SaeManifoldRho,
secondary: SaeManifoldTerm,
secondary_rho: &SaeManifoldRho,
target: ArrayView2<'_, f64>,
registry: Option<&AnalyticPenaltyRegistry>,
config: &StagewiseConfig,
) -> Result<(SaeManifoldTerm, SaeManifoldRho, f64, SaeManifoldLoss), String> {
primary.assignment.validate_rho_domain(primary_rho)?;
secondary.assignment.validate_rho_domain(secondary_rho)?;
let (mut merged, merged_rho) =
SaeManifoldTerm::merge_tiers(primary, primary_rho, secondary, secondary_rho)?;
merged.assignment.validate_rho_domain(&merged_rho)?;
merged.set_guards_enabled(false);
let (penalized_quasi_laplace, loss) =
frozen_joint_penalized_quasi_laplace(&mut merged, target, &merged_rho, registry, config)?;
merged.set_guards_enabled(true);
Ok((merged, merged_rho, penalized_quasi_laplace, loss))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::manifold::{AssignmentMode, SaeAssignment, SaeAtomBasisKind, SaeManifoldAtom};
fn fit_residual_covariance(
term: &SaeManifoldTerm,
target: ArrayView2<'_, f64>,
config: &StagewiseConfig,
) -> Result<Option<(Array2<f64>, StructuredResidualModel)>, String> {
let residual = current_residual(term, target)?;
fit_residual_covariance_on(term, residual, config)
}
use gam_terms::latent::LatentManifold;
use ndarray::Array2;
const ON: f64 = 6.0;
const OFF: f64 = -6.0;
fn test_config() -> StagewiseConfig {
StagewiseConfig {
inner_max_iter: 24,
learning_rate: 1.0,
ridge_ext_coord: 1e-6,
ridge_beta: 1e-6,
max_births: 3,
max_backfit_sweeps: 2,
min_effect_ev: 0.0,
max_factor_rank: 3,
structured_whitening: false,
}
}
fn circle_atom(
name: &str,
coords: &Array2<f64>,
dir_a: usize,
dir_b: usize,
p: usize,
) -> (SaeManifoldAtom, Array2<f64>) {
let geometry = SaeAtomGeometryPlan::new(
SaeAtomBasisKind::Periodic,
1,
SaeBasisResolution::PeriodicHarmonics { order: 1 },
SaeReferenceMetricPlan::UnitCircle,
)
.unwrap();
let bundle = geometry.evaluate_bundle(coords.view()).unwrap();
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, dir_a % p]] = 1.0;
decoder[[2, dir_b % p]] = 1.0;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
name.to_string(),
SaeAtomBasisKind::Periodic,
1,
bundle.basis_values,
bundle.basis_jacobian,
decoder,
bundle.reference_penalty,
)
.unwrap()
.with_basis_second_jet(bundle.evaluator)
.with_geometry_plan(geometry)
.unwrap();
(atom, coords.clone())
}
fn build_term(
atoms: Vec<SaeManifoldAtom>,
coord_blocks: Vec<Array2<f64>>,
active: &[Vec<bool>],
) -> (SaeManifoldTerm, SaeManifoldRho) {
let n = active.len();
let k = atoms.len();
let mut logits = Array2::<f64>::zeros((n, k));
for (row, atom_active) in active.iter().enumerate() {
for (atom, &on) in atom_active.iter().enumerate() {
logits[[row, atom]] = if on { ON } else { OFF };
}
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_blocks,
vec![LatentManifold::Circle { period: 1.0 }; k],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); k]);
(term, rho)
}
fn fitted_seed(
mut seed: SaeManifoldTerm,
mut rho: SaeManifoldRho,
target: ArrayView2<'_, f64>,
config: &StagewiseConfig,
) -> (SaeManifoldTerm, SaeManifoldRho) {
seed.set_guards_enabled(false);
seed.run_joint_fit_arrow_schur(
target,
&mut rho,
None,
config.inner_max_iter,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)
.expect("test seed K=1 fit must complete before stagewise entry");
(seed, rho)
}
fn is_non_decreasing(xs: &[f64]) -> bool {
xs.windows(2).all(|w| {
let tol = 1e-9 * (1.0 + w[0].abs());
w[1] >= w[0] - tol
})
}
#[test]
fn residual_covariance_propagates_invalid_residual_errors() {
let n = 8usize;
let p = 2usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("seed", &coords, 0, 1, p);
let (term, _rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let mut target = Array2::<f64>::zeros((n, p));
target[[3, 1]] = f64::NAN;
let err = fit_residual_covariance(&term, target.view(), &test_config())
.expect_err("non-finite residuals must be reported, not downgraded to None");
assert!(
err.contains("structured residual fit failed")
&& err.contains("residuals must be finite"),
"unexpected error: {err}"
);
}
#[test]
fn stagewise_recovers_planted_two_circles_ev_monotone() {
let n = 48usize;
let p = 4usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (atom1, cb1) = circle_atom("t1", &coords, 2, 3, p);
let active_truth: Vec<Vec<bool>> = (0..n).map(|r| vec![r < n / 2, r >= n / 2]).collect();
let (truth, _truth_rho) = build_term(
vec![atom0.clone(), atom1.clone()],
vec![cb0.clone(), cb1.clone()],
&active_truth,
);
let target = truth.fitted();
let config = test_config();
let (seed, rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let (seed, rho) = fitted_seed(seed, rho, target.view(), &config);
let result = fit_stagewise(seed, rho, target.view(), None, None, &config, None, None)
.expect("fit_stagewise must complete on planted two-circles");
assert!(
is_non_decreasing(&result.report.ev_trace),
"EV must be monotone non-decreasing in births by construction; got {:?}",
result.report.ev_trace
);
assert!(
result
.report
.terminal_joint_penalized_quasi_laplace
.is_finite(),
"terminal frozen joint penalized quasi-Laplace must be finite"
);
let seed_ev = result.report.ev_trace[0];
let final_ev = *result.report.ev_trace.last().unwrap();
assert!(
final_ev >= seed_ev - 1e-9,
"final EV {final_ev} must not fall below the seed EV {seed_ev}"
);
assert_eq!(
result.term.k_atoms(),
1 + result.report.births_accepted,
"K must equal the seed atom plus the accepted new-atom births"
);
}
#[test]
fn duplicate_atom_birth_is_rejected() {
let n = 40usize;
let p = 4usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (truth, _rho) =
build_term(vec![atom0.clone()], vec![cb0.clone()], &vec![vec![true]; n]);
let target = truth.fitted();
let config = StagewiseConfig {
min_effect_ev: 0.01,
..test_config()
};
let (seed, rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let (seed, rho) = fitted_seed(seed, rho, target.view(), &config);
let result = fit_stagewise(seed, rho, target.view(), None, None, &config, None, None)
.expect("fit_stagewise must complete on a fully-explained target");
assert_eq!(
result.report.births_accepted, 0,
"a duplicate/empty residual must yield no accepted births"
);
assert_eq!(
result.term.k_atoms(),
1,
"K must stay at the single seed atom"
);
assert!(
is_non_decreasing(&result.report.ev_trace),
"EV trace must remain monotone"
);
}
#[test]
fn anchor_scored_birth_prefers_uncontested_factor_2080() {
use gam_solve::inference::residual_factor::{ResidualFactorInput, StructuredResidualModel};
let n = 120usize;
let p = 6usize;
let h = n / 2; let inv_sqrt2 = 1.0 / 2.0_f64.sqrt();
let d_a = [inv_sqrt2, inv_sqrt2, 0.0, 0.0, 0.0, 0.0];
let d_b = [0.0, 0.0, inv_sqrt2, inv_sqrt2, 0.0, 0.0];
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
let s = (std::f64::consts::TAU * i as f64 / 11.0).cos(); let (dir, amp) = if i < h { (&d_a, 3.0) } else { (&d_b, 2.0) };
for j in 0..p {
residual[[i, j]] = amp * s * dir[j];
residual[[i, j]] += 0.04 * ((i * 7 + j * 13) as f64).sin();
}
}
let uniform_act = Array1::<f64>::ones(n);
let model = StructuredResidualModel::fit(ResidualFactorInput {
residuals: residual.view(),
activity: uniform_act.view(),
max_factor_rank: 2,
})
.unwrap();
assert!(model.factor_rank() >= 2, "need both planted factors");
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let build_ordered_beta_bernoulli = |logit: &dyn Fn(usize) -> f64| -> SaeManifoldTerm {
let mut logits = Array2::<f64>::zeros((n, 1));
for row in 0..n {
logits[[row, 0]] = logit(row);
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![cb0.clone()],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false),
)
.unwrap();
SaeManifoldTerm::new(vec![atom0.clone()], assignment).unwrap()
};
let contrast_term = build_ordered_beta_bernoulli(&|row| if row < h { 3.0 } else { -3.0 });
let act = activity_of(&contrast_term);
assert!(
act[0] > act[n - 1] + 1e-6,
"ordered Beta--Bernoulli activity must be higher on contested rows (got {} vs {})",
act[0],
act[n - 1]
);
let decoder = top_factor_birth_decoder(&contrast_term, &model, residual.view())
.unwrap()
.decoder;
let pick_strong = decoder[[0, 0]].hypot(decoder[[0, 1]]); let pick_anchor = decoder[[0, 2]].hypot(decoder[[0, 3]]); assert!(
pick_anchor > pick_strong,
"anchor-scored birth must pick the UNCONTESTED (dB, channels 2,3) factor, not the \
dominant-variance (dA, channels 0,1) one: |dB|={pick_anchor:.4} |dA|={pick_strong:.4}"
);
let uniform_term = build_ordered_beta_bernoulli(&|_| 0.5);
let decoder_u = top_factor_birth_decoder(&uniform_term, &model, residual.view())
.unwrap()
.decoder;
let u_strong = decoder_u[[0, 0]].hypot(decoder_u[[0, 1]]);
let u_anchor = decoder_u[[0, 2]].hypot(decoder_u[[0, 3]]);
assert!(
u_strong > u_anchor,
"uniform routing must fall back to the dominant-energy factor (dA, channels 0,1): \
|dA|={u_strong:.4} |dB|={u_anchor:.4}"
);
}
#[test]
fn residual_principal_fallback_fires_on_disjoint_not_noise_2080() {
let n = 400usize;
let p = 8usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (term, _rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let mut state = 0xC0FFEE_1234_5678_u64;
let mut rng = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f64) / ((1u64 << 31) as f64) - 1.0
};
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
let a = rng();
let b = rng();
residual[[i, 0]] = 2.0 * a;
residual[[i, 1]] = 2.0 * a; residual[[i, 2]] = 1.5 * b;
residual[[i, 3]] = 1.5 * b; for j in 0..p {
residual[[i, j]] += 0.03 * rng();
}
}
let seed = residual_principal_birth_candidate(&term, residual.view()).expect(
"disjoint block-diagonal residual must yield a fallback candidate \
(structure above the derived MP noise floor)",
);
assert!(
seed.circle().is_none(),
"unequal independent signals must NOT be seeded as a circle"
);
let (decoder, energy) = (seed.decoder, seed.energy);
assert!(energy > 0.0 && energy.is_finite());
let sig_mass: f64 = (0..4).map(|j| decoder[[0, j]].powi(2)).sum();
let noise_mass: f64 = (4..p).map(|j| decoder[[0, j]].powi(2)).sum();
assert!(
sig_mass > noise_mass,
"fallback birth direction must land on the signal block (0-3), not noise: \
sig={sig_mass:.3e} noise={noise_mass:.3e}"
);
let mut noise = Array2::<f64>::zeros((n, p));
for i in 0..n {
for j in 0..p {
noise[[i, j]] = rng();
}
}
assert!(
residual_principal_birth_candidate(&term, noise.view()).is_none(),
"pure-noise residual must be below the derived MP floor ⇒ no candidate (stop)"
);
}
#[test]
fn residual_principal_seeds_circle_as_rank2_not_dc_2101() {
let n = 240usize;
let p = 8usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (term, _rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let mut state = 0x5EED_2101_u64;
let mut rng = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f64) / ((1u64 << 31) as f64)
};
let mut residual = Array2::<f64>::zeros((n, p));
let mut planted = vec![0.0_f64; n];
for i in 0..n {
let theta = std::f64::consts::TAU * rng();
planted[i] = theta;
residual[[i, 2]] = theta.cos();
residual[[i, 3]] = theta.sin();
for j in 0..p {
residual[[i, j]] += 0.02 * (rng() - 0.5);
}
}
let seed = residual_principal_birth_candidate(&term, residual.view())
.expect("a real circle residual must yield a birth candidate");
let born_coords = seed
.circle()
.expect("a circle residual must be seeded as a rank-2 circle, not a DC direction")
.coords
.clone();
let dc: f64 = (0..p)
.map(|j| seed.decoder[[0, j]].powi(2))
.sum::<f64>()
.sqrt();
let harm: f64 = (0..p)
.map(|j| seed.decoder[[1, j]].powi(2) + seed.decoder[[2, j]].powi(2))
.sum::<f64>()
.sqrt();
assert!(
harm > 10.0 * dc.max(1e-9),
"circle seed must put mass on the cos/sin rows, not the DC row: harm={harm:.3} dc={dc:.3}"
);
let on_plane: f64 = [2usize, 3]
.iter()
.map(|&j| seed.decoder[[1, j]].powi(2) + seed.decoder[[2, j]].powi(2))
.sum();
let off_plane: f64 = (0..p)
.filter(|&j| j != 2 && j != 3)
.map(|j| seed.decoder[[1, j]].powi(2) + seed.decoder[[2, j]].powi(2))
.sum();
assert!(
on_plane > off_plane,
"circle seed 2-plane must land on the planted channels (2,3): on={on_plane:.3} off={off_plane:.3}"
);
let cmin = born_coords.iter().copied().fold(f64::INFINITY, f64::min);
let cmax = born_coords
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
assert!(
cmax - cmin > 0.5,
"seeded coordinate must span the circle (breaks the stationary point); range={:.3}",
cmax - cmin
);
let mut best_rmse = f64::INFINITY;
for &sign in &[1.0_f64, -1.0] {
let (mut cs, mut sn) = (0.0_f64, 0.0_f64);
for i in 0..n {
let r = std::f64::consts::TAU * born_coords[[i, 0]] - sign * planted[i];
cs += r.cos();
sn += r.sin();
}
let phase = sn.atan2(cs);
let mut sse = 0.0_f64;
for i in 0..n {
let mut e =
(std::f64::consts::TAU * born_coords[[i, 0]] - sign * planted[i] - phase)
.rem_euclid(std::f64::consts::TAU);
if e > std::f64::consts::PI {
e -= std::f64::consts::TAU;
}
sse += e * e;
}
best_rmse = best_rmse.min((sse / n as f64).sqrt());
}
assert!(
best_rmse < 0.15,
"seeded coordinate must recover the planted circle phase up to gauge; \
gauge-aligned circular RMSE = {best_rmse:.3} rad"
);
}
#[test]
fn certificate_rejects_two_circle_blend_2111() {
let n = 400usize;
let p = 8usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (term, _rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let mut state = 0x2111_B1E4_u64;
let mut rng = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 11) as f64) / ((1u64 << 53) as f64)
};
let mut clean = Array2::<f64>::zeros((n, p));
for i in 0..n {
let th = std::f64::consts::TAU * rng();
clean[[i, 2]] = th.cos();
clean[[i, 3]] = th.sin();
for j in 0..p {
clean[[i, j]] += 0.02 * (rng() - 0.5);
}
}
let clean_seed = residual_principal_birth_candidate(&term, clean.view())
.expect("clean circle must yield a birth candidate");
assert!(
clean_seed.circle().is_some(),
"positive control: a clean single circle must be seeded as a rank-2 circle"
);
let mut blend = Array2::<f64>::zeros((n, p));
for i in 0..n {
let a = std::f64::consts::TAU * rng();
let b = std::f64::consts::TAU * rng();
blend[[i, 0]] = a.cos() + b.cos();
blend[[i, 1]] = a.sin() + b.sin();
for j in 0..p {
blend[[i, j]] += 0.02 * (rng() - 0.5);
}
}
let blend_seed = residual_principal_birth_candidate(&term, blend.view())
.expect("blend residual still yields a (rank-1) birth candidate");
assert!(
blend_seed.circle().is_none(),
"κ-null certificate must REJECT the two-circle blend (κ≈1.5 > analytic-anchor \
gate) and fall through to the rank-1 seed, not born it as a clean circle"
);
}
#[test]
fn kappa_deflation_extracts_clean_circle_from_dense_torus_2111() {
let n = 320usize;
let p = 16usize;
let ncirc = 6usize;
let amps: Vec<f64> = (0..ncirc)
.map(|c| 1.0 - 0.45 * (c as f64) / ((ncirc - 1) as f64))
.collect();
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (term, _rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let mut s = 0x2111_D0BE_u64;
let mut rng = || {
s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((s >> 11) as f64) / ((1u64 << 53) as f64)
};
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
for c in 0..ncirc {
let th = std::f64::consts::TAU * rng();
residual[[i, 2 * c]] += amps[c] * th.cos();
residual[[i, 2 * c + 1]] += amps[c] * th.sin();
}
for j in 0..p {
residual[[i, j]] += 0.05 * (rng() - 0.5);
}
}
let seed = residual_principal_birth_candidate(&term, residual.view())
.expect("dense torus must yield a birth candidate");
let dec = seed
.circle()
.map(|_| &seed.decoder)
.expect("dense torus must be seeded as a CLEAN rank-2 circle (κ-deflation), not DC");
let total: f64 = (0..p)
.map(|j| dec[[1, j]].powi(2) + dec[[2, j]].powi(2))
.sum();
assert!(total > 0.0, "born plane must carry mass");
let mut fracs: Vec<f64> = (0..ncirc)
.map(|c| {
let e = dec[[1, 2 * c]].powi(2)
+ dec[[2, 2 * c]].powi(2)
+ dec[[1, 2 * c + 1]].powi(2)
+ dec[[2, 2 * c + 1]].powi(2);
e / total
})
.collect();
fracs.sort_by(|a, b| b.total_cmp(a));
assert!(
fracs[0] > 0.80 && fracs[1] < 0.20,
"κ-deflation must isolate ONE clean circle from the dense torus: top channel-pair \
energy fraction {:.3} (want > 0.80), second {:.3} (want < 0.20) — a blended plane \
would spread across circles",
fracs[0],
fracs[1]
);
}
#[test]
fn born_circle_survives_on_incumbent_sparse_rows_2109() {
let n = 160usize;
let p = 8usize;
let h = n / 2; let mut state = 0x2109_5A17_0000_0001u64;
let mut rng = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f64) / ((1u64 << 31) as f64)
};
let m = n - h;
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
if i >= h {
let theta = std::f64::consts::TAU * ((i - h) as f64) / m as f64;
residual[[i, 0]] = theta.cos();
residual[[i, 1]] = theta.sin();
}
for j in 0..p {
residual[[i, j]] += 0.02 * (rng() - 0.5);
}
}
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (inc_atom, inc_cb) = circle_atom("inc", &coords, 4, 5, p);
let mut inc_logits = Array2::<f64>::zeros((n, 1));
for row in 0..n {
inc_logits[[row, 0]] = if row < h { 4.0 } else { -6.0 };
}
let inc_assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
inc_logits,
vec![inc_cb],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, false),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![inc_atom], inc_assignment).unwrap();
term.set_guards_enabled(false);
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
let seed = residual_principal_birth_candidate(&term, residual.view())
.expect("an incumbent-sparse circle must still yield a birth candidate");
let circle = seed
.circle()
.expect("residual must be seeded as a rank-2 circle");
let born_geometry = circle.geometry.clone();
let born_coords = circle.coords.clone();
let gate = circle.gate.clone();
let present_on_circle = (h..n).filter(|&i| gate[i].is_finite()).count();
let present_off_circle = (0..h).filter(|&i| gate[i].is_finite()).count();
assert!(
present_on_circle > (n - h) / 2 && present_off_circle < h / 4,
"own-presence must fire on the circle's rows, not the empty ones: \
on={present_on_circle}/{} off={present_off_circle}/{h}",
n - h
);
let (child, mut child_rho) = crate::structure_harvest::born_circle_atom(
&term,
&rho,
born_geometry,
seed.decoder.clone(),
born_coords,
gate,
)
.expect("born_circle_atom");
let born = child.k_atoms() - 1;
let mut min_born_logit_on_circle = f64::INFINITY;
for row in h..n {
min_born_logit_on_circle =
min_born_logit_on_circle.min(child.assignment.logits[[row, born]]);
}
assert!(
min_born_logit_on_circle > 1.0,
"born circle on incumbent-sparse rows must seed a STRONG own-presence gate \
(>1), not the incumbent's negative inc_max (−6); got min={min_born_logit_on_circle:.3}"
);
let config = StagewiseConfig {
inner_max_iter: 40,
learning_rate: 1.0,
ridge_ext_coord: 1e-6,
ridge_beta: 1e-6,
max_births: 1,
max_backfit_sweeps: 1,
min_effect_ev: 0.0,
max_factor_rank: 3,
structured_whitening: false,
};
let mut child = child;
fit_single_atom_response_in_place(
&mut child,
&mut child_rho,
born,
residual.view(),
None,
&config,
)
.expect("K=1 born-circle sub-fit must complete");
let born_norm = child.atoms[born]
.decoder_coefficients()
.iter()
.map(|v| v * v)
.sum::<f64>()
.sqrt();
assert!(
born_norm.is_finite() && born_norm > 0.3,
"born circle must SURVIVE the ordered Beta--Bernoulli sub-fit on incumbent-sparse rows \
(‖B‖ O(1)); got ‖B‖={born_norm:.3e} (a collapse to ~1e-4 is the #2109 bug)"
);
}
#[test]
fn top_factor_birth_mirrors_circle_seed_2109() {
use gam_solve::inference::residual_factor::{ResidualFactorInput, StructuredResidualModel};
let n = 200usize;
let p = 8usize;
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
let theta = std::f64::consts::TAU * (i as f64) / n as f64;
let (c, s) = (theta.cos(), theta.sin());
residual[[i, 0]] = c;
residual[[i, 1]] = c;
residual[[i, 2]] = s;
residual[[i, 3]] = s;
for j in 0..p {
residual[[i, j]] += 0.02 * ((i * 7 + j * 5) as f64).sin();
}
}
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let logits = Array2::<f64>::from_elem((n, 1), 0.5);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![cb0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, false),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom0], assignment).unwrap();
let activity = activity_of(&term);
let model = StructuredResidualModel::fit(ResidualFactorInput {
residuals: residual.view(),
activity: activity.view(),
max_factor_rank: 2,
})
.unwrap();
assert!(
model.factor_rank() >= 1,
"the correlated cos/sin axes must make the factor model rank ≥ 1 so \
top_factor_birth_decoder is the active path; got rank {}",
model.factor_rank()
);
let seed = top_factor_birth_decoder(&term, &model, residual.view())
.expect("the entangled path must yield a birth seed");
assert!(
seed.circle().is_some(),
"top_factor_birth_decoder must MIRROR the #2101 circle seed on a degenerate \
2-plane residual (circle_coords Some), not the flat DC seed"
);
let gate = &seed
.circle()
.expect("the mirrored circle seed must carry typed circle state")
.gate;
assert!(
gate.iter().filter(|g| g.is_finite()).count() > n / 2,
"the mirrored circle must mark its present rows with a finite own-presence gate"
);
let dc: f64 = (0..p)
.map(|j| seed.decoder[[0, j]].powi(2))
.sum::<f64>()
.sqrt();
let harm_on: f64 = [0usize, 1, 2, 3]
.iter()
.map(|&j| seed.decoder[[1, j]].powi(2) + seed.decoder[[2, j]].powi(2))
.sum();
let harm_off: f64 = (4..p)
.map(|j| seed.decoder[[1, j]].powi(2) + seed.decoder[[2, j]].powi(2))
.sum();
assert!(
harm_on > harm_off && harm_on.sqrt() > 10.0 * dc.max(1e-9),
"mirrored circle seed must put its 2-plane on the cos/sin rows of channels \
(0,1,2,3): harm_on={harm_on:.3} harm_off={harm_off:.3} dc={dc:.3}"
);
}
#[test]
fn progress_callback_emits_pre_birth_checkpoints() {
let n = 32usize;
let p = 4usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (atom1, cb1) = circle_atom("t1", &coords, 2, 3, p);
let active_truth: Vec<Vec<bool>> = (0..n).map(|r| vec![r < n / 2, r >= n / 2]).collect();
let (truth, _rho) = build_term(
vec![atom0.clone(), atom1],
vec![cb0.clone(), cb1],
&active_truth,
);
let target = truth.fitted();
let config = StagewiseConfig {
max_births: 1,
max_backfit_sweeps: 0,
..test_config()
};
let (seed, rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let (seed, rho) = fitted_seed(seed, rho, target.view(), &config);
let mut events: Vec<(StagewiseEventKind, bool, usize, Option<BirthKind>)> = Vec::new();
let mut progress = |event: StagewiseProgress<'_>| -> Result<(), String> {
events.push((
event.event,
event.checkpoint,
event.k_atoms,
event.candidate,
));
Ok(())
};
fit_stagewise(
seed,
rho,
target.view(),
None,
None,
&config,
Some(&mut progress),
None,
)
.expect("fit_stagewise must complete while emitting progress");
assert_eq!(
events.first().map(|event| event.0),
Some(StagewiseEventKind::SeedReady),
"the first callback must expose the fitted K=1 seed"
);
assert_eq!(
events.get(1).map(|event| event.0),
Some(StagewiseEventKind::BirthRoundStarted),
"the second callback must expose a durable birth-round checkpoint"
);
assert_eq!(events[0].1, true, "seed_ready must be checkpointable");
assert_eq!(
events[1].1, true,
"birth_round_started must be checkpointable before residual work"
);
assert_eq!(events[0].2, 1, "seed checkpoint must be K=1");
let pos = |kind: StagewiseEventKind| -> usize {
events
.iter()
.position(|event| event.0 == kind)
.expect("expected progress event")
};
assert!(
pos(StagewiseEventKind::ResidualModelStarted)
< pos(StagewiseEventKind::CurrentEvidenceStarted),
"residual-fit progress must precede current-evidence progress"
);
assert!(
pos(StagewiseEventKind::CurrentEvidenceStarted)
< pos(StagewiseEventKind::CandidateStarted),
"current-evidence progress must precede candidate fitting"
);
assert!(
events
.iter()
.any(|event| event.0 == StagewiseEventKind::CandidateStarted
&& event.3 == Some(BirthKind::NewAtom)),
"first birth must report the new-atom candidate"
);
}
#[test]
fn backfitting_ev_is_monotone() {
let n = 48usize;
let p = 4usize;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (atom0, cb0) = circle_atom("t0", &coords, 0, 1, p);
let (atom1, cb1) = circle_atom("t1", &coords, 2, 3, p);
let active_truth: Vec<Vec<bool>> = (0..n).map(|r| vec![r < n / 2, r >= n / 2]).collect();
let (truth, _rho) = build_term(
vec![atom0.clone(), atom1.clone()],
vec![cb0.clone(), cb1.clone()],
&active_truth,
);
let target = truth.fitted();
let config = test_config();
let (seed, rho) = build_term(vec![atom0], vec![cb0], &vec![vec![true]; n]);
let (seed, rho) = fitted_seed(seed, rho, target.view(), &config);
let result = fit_stagewise(seed, rho, target.view(), None, None, &config, None, None)
.expect("fit_stagewise must complete");
assert!(
is_non_decreasing(&result.report.backfit_ev_trace),
"backfitting EV must be monotone non-decreasing; got {:?}",
result.report.backfit_ev_trace
);
}
fn build_term_gate(
atoms: Vec<SaeManifoldAtom>,
coord_blocks: Vec<Array2<f64>>,
active: &[Vec<bool>],
mode: AssignmentMode,
) -> (SaeManifoldTerm, SaeManifoldRho) {
let n = active.len();
let k = atoms.len();
let mut logits = Array2::<f64>::zeros((n, k));
for (row, atom_active) in active.iter().enumerate() {
for (atom, &on) in atom_active.iter().enumerate() {
logits[[row, atom]] = if on { ON } else { OFF };
}
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_blocks,
vec![LatentManifold::Circle { period: 1.0 }; k],
mode,
)
.unwrap();
let term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); k]);
(term, rho)
}
fn disjoint_circle_seed(n: usize, p: usize, a: usize, b: usize, rows: &[usize]) -> BirthSeed {
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, a]] = 1.0;
decoder[[2, b]] = 1.0;
let mut gate = vec![f64::NEG_INFINITY; n];
let mut phases = Array2::<f64>::zeros((n, 1));
for (pos, &r) in rows.iter().enumerate() {
gate[r] = 2.0;
phases[[r, 0]] = pos as f64 / rows.len() as f64;
}
BirthSeed {
decoder,
energy: 1.0,
kind: BirthSeedKind::Circle(CircleBirthSeed {
geometry: SaeAtomGeometryPlan::new(
SaeAtomBasisKind::Periodic,
1,
SaeBasisResolution::PeriodicHarmonics { order: 1 },
SaeReferenceMetricPlan::UnitCircle,
)
.unwrap(),
coords: phases,
gate,
}),
}
}
fn lcg_uniform(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*state >> 11) as f64) / ((1u64 << 53) as f64)
}
fn lcg_normal(state: &mut u64) -> f64 {
let u1 = lcg_uniform(state).max(1e-12);
let u2 = lcg_uniform(state);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn planted_axis_dense_circles(
n: usize,
p: usize,
k: usize,
amp: f64,
sigma: f64,
seed: u64,
) -> Array2<f64> {
let mut state = seed;
let mut data = Array2::<f64>::zeros((n, p));
for i in 0..n {
for c in 0..k {
let th = std::f64::consts::TAU * lcg_uniform(&mut state);
data[[i, 2 * c]] += amp * th.cos();
data[[i, 2 * c + 1]] += amp * th.sin();
}
for j in 0..p {
data[[i, j]] += sigma * lcg_normal(&mut state);
}
}
data
}
#[test]
fn batched_disjoint_birth_fit_matches_serial_bit_for_bit() {
let n = 40usize;
let p = 6usize;
let h = n / 2;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(r, _)| r as f64 / n as f64);
let mut target = Array2::<f64>::zeros((n, p));
for r in 0..h {
let th = std::f64::consts::TAU * (r as f64) / (h as f64);
target[[r, 0]] = th.cos();
target[[r, 1]] = th.sin();
}
for r in h..n {
let th = std::f64::consts::TAU * ((r - h) as f64) / ((n - h) as f64);
target[[r, 2]] = th.cos();
target[[r, 3]] = th.sin();
}
let mode = AssignmentMode::threshold_gate(1.0, -3.0);
let (atom0, cb0) = circle_atom("seed", &coords, 0, 1, p);
let active: Vec<Vec<bool>> = (0..n).map(|r| vec![r < h]).collect();
let (mut seed, mut rho) = build_term_gate(vec![atom0], vec![cb0], &active, mode);
seed.set_guards_enabled(false);
let config = test_config();
seed.run_joint_fit_arrow_schur(
target.view(),
&mut rho,
None,
config.inner_max_iter,
config.learning_rate,
config.ridge_ext_coord,
config.ridge_beta,
)
.expect("seed K=1 fit");
let rows_a: Vec<usize> = (0..h).collect();
let rows_b: Vec<usize> = (h..n).collect();
let seed_a = disjoint_circle_seed(n, p, 0, 1, &rows_a);
let seed_b = disjoint_circle_seed(n, p, 2, 3, &rows_b);
let r0 = current_residual(&seed, target.view()).unwrap();
let b_batched = race_birth_seed(
&seed,
&rho,
&seed_b,
r0.view(),
target.view(),
None,
&config,
)
.expect("race B against R0");
let a = race_birth_seed(
&seed,
&rho,
&seed_a,
r0.view(),
target.view(),
None,
&config,
)
.expect("race A against R0");
let (term_a, rho_a) = append_fitted_atom(
&seed,
&rho,
a.born_atom.clone(),
a.born_coord.clone(),
&a.born_logit_col,
a.born_ard.clone(),
a.born_log_lambda_smooth,
)
.expect("append A");
let r1 = current_residual(&term_a, target.view()).unwrap();
let mut rows_b_diff = 0.0_f64;
for &r in &rows_b {
for j in 0..p {
rows_b_diff += (r0[[r, j]] - r1[[r, j]]).abs();
}
}
assert!(
rows_b_diff < 1e-9,
"R0 and R1 must be identical on B's disjoint rows; L1 diff {rows_b_diff}"
);
let b_serial = race_birth_seed(
&term_a,
&rho_a,
&seed_b,
r1.view(),
target.view(),
None,
&config,
)
.expect("race B against R1");
let d_batched = b_batched.born_atom.decoder_coefficients();
let d_serial = b_serial.born_atom.decoder_coefficients();
let decoder_diff = (d_batched - d_serial).mapv(f64::abs).sum();
assert!(
decoder_diff < 1e-6,
"disjoint birth B must fit IDENTICALLY against R0 (batched) and R1 (serial); \
decoder L1 diff {decoder_diff}"
);
let seed_ev = ev_of(&seed, target.view());
let term_a_ev = ev_of(&term_a, target.view());
let charge_batched = b_batched.ev - seed_ev;
let charge_serial = b_serial.ev - term_a_ev;
assert!(
(charge_batched - charge_serial).abs() < 1e-6,
"disjoint birth B marginal ΔEV charge must match batched vs serial; \
{charge_batched} vs {charge_serial}"
);
}
#[test]
fn batched_driver_matches_serial_and_batches() {
let n = 900usize;
let p = 16usize;
let q = 3usize; let coords = Array2::<f64>::from_shape_fn((n, 1), |(r, _)| r as f64 / n as f64);
let target = planted_axis_dense_circles(n, p, q, 1.0, 0.03, 0x2111_A11E_u64);
let mode = AssignmentMode::softmax(1.0);
let mut config = test_config();
config.max_births = 8;
config.max_backfit_sweeps = 1;
let build_seed = || {
let (atom0, cb0) = circle_atom("seed", &coords, 0, 1, p);
let (mut seed, rho) =
build_term_gate(vec![atom0], vec![cb0], &vec![vec![true]; n], mode);
seed.set_guards_enabled(false);
(seed, rho)
};
let (seed_s, rho_s) = build_seed();
let serial = fit_stagewise(
seed_s,
rho_s,
target.view(),
None,
None,
&config,
None,
None,
)
.expect("serial driver");
assert!(
serial.report.births_accepted >= 2,
"serial driver must grow K on the planted {q}-circle image for the parity \
comparison to be meaningful; births_accepted={}",
serial.report.births_accepted
);
let (seed_b, rho_b) = build_seed();
let batch_config = BatchedStagewiseConfig {
base: config,
max_candidates_per_round: 8,
};
let batched =
fit_stagewise_batched(seed_b, rho_b, target.view(), None, None, &batch_config)
.expect("batched driver");
assert!(
batched.report.births_accepted >= 2,
"batched driver must also grow K on the planted {q}-circle image; \
births_accepted={}",
batched.report.births_accepted
);
assert!(
is_non_decreasing(&batched.report.ev_trace),
"batched EV must be monotone non-decreasing; got {:?}",
batched.report.ev_trace
);
assert!(
batched
.report
.terminal_joint_penalized_quasi_laplace
.is_finite(),
"batched terminal joint penalized quasi-Laplace must be finite"
);
let serial_k = serial.term.k_atoms() as i64;
let batched_k = batched.term.k_atoms() as i64;
assert!(
(serial_k - batched_k).abs() <= 1,
"batched atom count {batched_k} must match serial {serial_k} within 1"
);
let serial_ev = *serial.report.ev_trace.last().unwrap();
let batched_ev = *batched.report.ev_trace.last().unwrap();
assert!(
(serial_ev - batched_ev).abs() < 0.1,
"batched final EV {batched_ev} must match serial {serial_ev} within 0.1"
);
}
#[test]
fn batched_round_co_accepts_via_both_routes() {
let n = 40usize;
let p = 8usize;
let h = n / 2;
let coords = Array2::<f64>::from_shape_fn((n, 1), |(r, _)| r as f64 / n as f64);
let config = test_config();
let (atom0, cb0) = circle_atom("seed", &coords, 0, 1, p);
let (mut seed_t, rho_t) = build_term_gate(
vec![atom0],
vec![cb0],
&vec![vec![true]; n],
AssignmentMode::threshold_gate(1.0, -3.0),
);
seed_t.set_guards_enabled(false);
let target = Array2::<f64>::from_shape_fn((n, p), |(r, j)| {
if j < 2 {
let th = std::f64::consts::TAU * (r as f64) / (n as f64);
if j == 0 { th.cos() } else { th.sin() }
} else {
0.0
}
});
let r0 = current_residual(&seed_t, target.view()).unwrap();
let sa = disjoint_circle_seed(n, p, 0, 1, &(0..n).collect::<Vec<_>>());
let template = race_birth_seed(
&seed_t,
&rho_t,
&sa,
r0.view(),
target.view(),
None,
&config,
)
.expect("race template");
let with_support = |rows: Vec<usize>, dims: Vec<usize>| -> RacedCandidate {
let mut c = template.clone();
c.support = rows;
c.out_support = dims;
c
};
let rows_a: Vec<usize> = (0..h).collect();
let rows_b: Vec<usize> = (h..n).collect();
let raced_rows = vec![
with_support(rows_a.clone(), vec![0, 1]),
with_support(rows_b.clone(), vec![0, 1]),
];
let (accepted, requeued) = select_disjoint_batch(&raced_rows, &[0, 1], 8);
assert_eq!(
accepted.len(),
2,
"two ROW-disjoint candidates must co-accept (accepted={accepted:?}, requeued={requeued})"
);
let all_rows: Vec<usize> = (0..n).collect();
let raced_dims = vec![
with_support(all_rows.clone(), vec![0, 1]),
with_support(all_rows.clone(), vec![2, 3]),
];
let (accepted2, requeued2) = select_disjoint_batch(&raced_dims, &[0, 1], 8);
assert_eq!(
accepted2.len(),
2,
"two OUTPUT-DIM-disjoint candidates must co-accept \
(accepted={accepted2:?}, requeued={requeued2})"
);
let raced_conflict = vec![
with_support(all_rows.clone(), vec![0, 1]),
with_support(all_rows.clone(), vec![0, 1]),
];
let (accepted3, requeued3) = select_disjoint_batch(&raced_conflict, &[0, 1], 8);
assert_eq!(
(accepted3.len(), requeued3),
(1, 1),
"candidates overlapping on BOTH rows and output dims must NOT co-accept; \
accepted={accepted3:?} requeued={requeued3}"
);
}
}
#[cfg(test)]
mod birth_rejection_reason_tests {
use super::*;
#[test]
fn every_birth_gate_clause_names_itself() {
assert_eq!(
classify_birth_candidate(1.0, 0.5, 0.4, 2.0, 0.05),
None,
"a candidate that beats the criterion and clears the floor must pass"
);
assert_eq!(
classify_birth_candidate(2.0, 0.5, 0.4, 2.0, 0.0),
Some(BirthRejection::EvidenceNotImproved {
criterion: 2.0,
must_be_below: 2.0
}),
"equality is NOT improvement — the gate is strict"
);
assert_eq!(
classify_birth_candidate(1.0, 0.41, 0.4, 2.0, 0.05),
Some(BirthRejection::EffectBelowFloor {
delta_ev: 0.41_f64 - 0.4_f64,
floor: 0.05
})
);
let negative = classify_birth_candidate(1.0, 0.3, 0.4, 2.0, 0.0);
match negative {
Some(BirthRejection::EffectBelowFloor { delta_ev, floor }) => {
assert!(delta_ev < 0.0, "a real negative ΔEV must survive as negative");
assert_eq!(floor, 0.0);
}
other => panic!("expected the effect floor to reject a negative ΔEV; got {other:?}"),
}
assert_eq!(
classify_birth_candidate(f64::NAN, 0.5, 0.4, 2.0, 0.0),
Some(BirthRejection::NonFiniteCriterion)
);
assert_eq!(
classify_birth_candidate(1.0, f64::NAN, 0.4, 2.0, 0.0),
Some(BirthRejection::NonFiniteEv)
);
assert!(
classify_birth_candidate(1.0, 0.5, f64::NAN, 2.0, 0.0).is_some(),
"a NaN ΔEV must be rejected, not passed"
);
}
#[test]
fn serial_birth_ledger_retains_errors_and_selects_the_best_arm_2556() {
let current_criterion = 10.0;
let cur_ev = 0.5;
let floor = 0.1;
let passing = vec![
BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy: 1.0,
criterion: 8.0,
ev: 0.8,
},
BirthCandidateAttempt::Measured {
kind: BirthKind::ChartExtension,
factor_energy: 2.0,
criterion: 7.0,
ev: 0.9,
},
];
let selected =
best_passing_birth_candidate(&passing, cur_ev, current_criterion, floor);
assert_eq!(
selected,
Some(1),
"the lower-criterion second arm, not generation-order arm zero, must win"
);
let ledger = record_birth_candidates(
&passing,
&[selected.expect("one serial arm must pass")],
PassingCandidateDisposition::Outranked,
cur_ev,
current_criterion,
floor,
)
.expect("a valid serial decision must produce a ledger");
assert_eq!(ledger.len(), 2);
assert_eq!(ledger[0].decision, BirthCandidateDecision::Outranked);
assert_eq!(ledger[1].decision, BirthCandidateDecision::Accepted);
assert_eq!(ledger[0].delta_ev, Some(0.8 - cur_ev));
assert_eq!(ledger[1].joint_penalized_quasi_laplace, Some(7.0));
let failed_and_rejected = vec![
BirthCandidateAttempt::FitFailed {
kind: BirthKind::NewAtom,
factor_energy: 3.0,
error: "new-atom fit refused".to_string(),
},
BirthCandidateAttempt::Measured {
kind: BirthKind::ChartExtension,
factor_energy: 4.0,
criterion: current_criterion,
ev: 0.9,
},
];
assert_eq!(
best_passing_birth_candidate(
&failed_and_rejected,
cur_ev,
current_criterion,
floor,
),
None
);
let rejected_ledger = record_birth_candidates(
&failed_and_rejected,
&[],
PassingCandidateDisposition::Outranked,
cur_ev,
current_criterion,
floor,
)
.expect("failed and gate-rejected arms are valid durable outcomes");
assert_eq!(rejected_ledger[0].delta_ev, None);
assert_eq!(rejected_ledger[0].joint_penalized_quasi_laplace, None);
assert_eq!(
rejected_ledger[0].decision,
BirthCandidateDecision::FitFailed("new-atom fit refused".to_string())
);
assert_eq!(
rejected_ledger[1].decision,
BirthCandidateDecision::GateRejected(BirthRejection::EvidenceNotImproved {
criterion: current_criterion,
must_be_below: current_criterion,
})
);
}
#[test]
fn batch_birth_ledger_preserves_harvest_order_and_exact_dispositions_2556() {
let current_criterion = 10.0;
let cur_ev = 0.5;
let floor = 0.1;
let attempts = vec![
BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy: 1.0,
criterion: current_criterion,
ev: 0.8,
},
BirthCandidateAttempt::FitFailed {
kind: BirthKind::NewAtom,
factor_energy: 2.0,
error: "seed one failed".to_string(),
},
BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy: 3.0,
criterion: 6.0,
ev: 0.9,
},
BirthCandidateAttempt::Measured {
kind: BirthKind::NewAtom,
factor_energy: 4.0,
criterion: 7.0,
ev: 0.85,
},
];
assert_eq!(
best_passing_birth_candidate(&attempts, cur_ev, current_criterion, floor),
Some(2),
"the first harvested candidate is not a proxy for the best adjudicated candidate"
);
let ledger = record_birth_candidates(
&attempts,
&[2],
PassingCandidateDisposition::DeferredByBatchSelection,
cur_ev,
current_criterion,
floor,
)
.expect("a valid batch selection must produce a ledger");
assert_eq!(
ledger
.iter()
.map(|candidate| candidate.factor_energy)
.collect::<Vec<_>>(),
vec![1.0, 2.0, 3.0, 4.0],
"fit failures must not compact or reorder the harvest ledger"
);
assert_eq!(
ledger[0].decision,
BirthCandidateDecision::GateRejected(BirthRejection::EvidenceNotImproved {
criterion: current_criterion,
must_be_below: current_criterion,
})
);
assert_eq!(
ledger[1].decision,
BirthCandidateDecision::FitFailed("seed one failed".to_string())
);
assert_eq!(ledger[2].decision, BirthCandidateDecision::Accepted);
assert_eq!(
ledger[3].decision,
BirthCandidateDecision::DeferredByBatchSelection
);
assert!(
record_birth_candidates(
&attempts,
&[0],
PassingCandidateDisposition::DeferredByBatchSelection,
cur_ev,
current_criterion,
floor,
)
.is_err(),
"the ledger must refuse a selected index which did not clear the gate"
);
}
}