use ndarray::{Array1, ArrayView2, Axis};
use gam_solve::rho_optimizer::OuterCriterionCertificate;
use crate::front_door::{SaeFitLane, admit_topk_manifold};
use crate::manifold::{
SaeSupportFixedPointReport, SaeSupportOuterRequest, SaeSupportSeedRequest,
SaeSupportSparseTerm, SaeSupportTermSeedRequest, build_sae_support_seed,
build_sae_support_term_seed, run_sae_support_outer, sae_support_effective_atom_dims,
};
use crate::migration_ledger::{BirthSeed, MoveEvidence, MoveReason, MoveStage, SaeMigrationLedger};
use crate::sparse_dict::{
BlockSeedPolicy, BlockSparseConfig, BlockSparseFit, fit_block_sparse_dictionary_with_seed,
};
use crate::tiered::Tier0Mean;
const FARTHEST_POINT_SEED_MAX_OPS: u128 = 1_000_000_000;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum TieredSeedPolicy {
#[default]
Auto,
FarthestPoint,
CoordinatePartition,
}
impl TieredSeedPolicy {
fn resolve(self, n: usize, p: usize, config: &BlockSparseConfig) -> BlockSeedPolicy {
match self {
TieredSeedPolicy::FarthestPoint => BlockSeedPolicy::FarthestPoint,
TieredSeedPolicy::CoordinatePartition => BlockSeedPolicy::CoordinatePartition,
TieredSeedPolicy::Auto => {
let ops = (n as u128)
* (p as u128)
* (config.n_blocks as u128)
* (config.block_size as u128);
if ops > FARTHEST_POINT_SEED_MAX_OPS {
BlockSeedPolicy::CoordinatePartition
} else {
BlockSeedPolicy::FarthestPoint
}
}
}
}
}
#[derive(Clone, Debug)]
pub struct Tier2SupportConfig {
pub atom_basis: String,
pub atom_dim: usize,
pub n_atoms: usize,
pub support_k: usize,
pub initial_smoothness: f64,
pub max_outer_iter: usize,
pub max_inner_iter: usize,
pub inner_tolerance: f64,
pub trust_radius: f64,
pub random_state: u64,
}
impl Default for Tier2SupportConfig {
fn default() -> Self {
Self {
atom_basis: "periodic".to_string(),
atom_dim: 1,
n_atoms: 256,
support_k: 4,
initial_smoothness: 1.0,
max_outer_iter: 64,
max_inner_iter: 256,
inner_tolerance: 1.0e-8,
trust_radius: 1.0,
random_state: 0xC0FF_EE00_D15E_A5E5,
}
}
}
#[derive(Clone, Debug)]
pub struct TieredFitConfig {
pub tier1: BlockSparseConfig,
pub tier1_seed: TieredSeedPolicy,
pub tier2_enabled: bool,
pub tier2: Tier2SupportConfig,
}
impl TieredFitConfig {
pub fn linear_bulk(n_blocks: usize, block_size: usize) -> Self {
Self {
tier1: BlockSparseConfig::new(n_blocks, block_size),
tier1_seed: TieredSeedPolicy::Auto,
tier2_enabled: false,
tier2: Tier2SupportConfig::default(),
}
}
pub fn tiered(n_blocks: usize, block_size: usize) -> Self {
Self {
tier1: BlockSparseConfig::new(n_blocks, block_size),
tier1_seed: TieredSeedPolicy::Auto,
tier2_enabled: true,
tier2: Tier2SupportConfig::default(),
}
}
}
#[derive(Clone, Debug)]
pub struct Tier2SupportFit {
pub mean: Array1<f64>,
pub term: SaeSupportSparseTerm,
pub lambda_smooth: Vec<f64>,
pub criterion: f64,
pub fixed_point: SaeSupportFixedPointReport,
pub outer_certificate: OuterCriterionCertificate,
pub outer_iterations: usize,
pub requested_atoms: usize,
pub retained_atoms: usize,
pub explained_variance: f64,
}
#[derive(Clone, Debug)]
pub struct TieredFitReport {
pub tier0: Tier0Mean,
pub tier1: BlockSparseFit,
pub tier2: Option<Tier2SupportFit>,
pub ledger: SaeMigrationLedger,
pub explained_variance: f64,
}
pub fn fit_tiered(
z: ArrayView2<'_, f64>,
config: &TieredFitConfig,
) -> Result<TieredFitReport, String> {
let tier0 = Tier0Mean::fit(z)?;
let r0 = tier0.apply(z)?;
let r0_f32 = r0.mapv(|v| v as f32);
let seed_policy = config
.tier1_seed
.resolve(r0_f32.nrows(), r0_f32.ncols(), &config.tier1);
let tier1 = fit_block_sparse_dictionary_with_seed(r0_f32.view(), &config.tier1, seed_policy)?;
let mut ledger = SaeMigrationLedger::new();
let n_dead = tier1
.block_utilization
.iter()
.filter(|&&u| u == 0.0)
.count();
if n_dead > 0 {
ledger.death(
MoveStage::Linear,
MoveReason::DeadRouting,
n_dead,
None,
MoveEvidence::none(),
f64::NAN,
);
}
let (tier2, explained_variance) = if config.tier2_enabled {
let fit = fit_tier2_support(r0.view(), &tier1, &config.tier2)?;
record_support_moves(&mut ledger, &fit);
let ev = fit.explained_variance;
(Some(fit), ev)
} else {
(None, tier1.explained_variance)
};
Ok(TieredFitReport {
tier0,
tier1,
tier2,
ledger,
explained_variance,
})
}
fn fit_tier2_support(
r0: ArrayView2<'_, f64>,
tier1: &BlockSparseFit,
config: &Tier2SupportConfig,
) -> Result<Tier2SupportFit, String> {
let (n_obs, output_dim) = r0.dim();
let linear = tier1.reconstruct();
if linear.dim() != (n_obs, output_dim) {
return Err(format!(
"fit_tier2_support: Tier-1 reconstruction {:?} does not match residual ({n_obs}, {output_dim})",
linear.dim()
));
}
let mut residual = r0.to_owned();
for row in 0..n_obs {
for column in 0..output_dim {
residual[[row, column]] -= linear[[row, column]] as f64;
}
}
let mean = residual
.mean_axis(Axis(0))
.ok_or_else(|| "fit_tier2_support: residual mean_axis returned None".to_string())?;
let centered = &residual - &mean.view().insert_axis(Axis(0));
let requested_atoms = config.n_atoms;
let atom_basis = vec![config.atom_basis.clone(); requested_atoms];
let atom_dim = vec![config.atom_dim; requested_atoms];
let effective_dims = sae_support_effective_atom_dims(&atom_basis, &atom_dim)?;
let d_max = effective_dims.iter().copied().max().unwrap_or(1);
let admission =
admit_topk_manifold(n_obs, output_dim, requested_atoms, d_max, config.support_k)?;
if admission.lane != SaeFitLane::CurvedStreaming {
return Err(format!(
"fit_tier2_support: the curved refinement is the overcomplete support-sparse lane, \
which requires K > P (CurvedStreaming admission); got lane {:?} at N={n_obs}, \
P={output_dim}, K={requested_atoms}. Widen the Tier-2 dictionary past the residual \
dimension",
admission.lane
));
}
let seed = build_sae_support_seed(SaeSupportSeedRequest {
target: centered.view(),
atom_basis: &atom_basis,
atom_dim: &atom_dim,
support_k: config.support_k,
random_state: config.random_state,
admission,
})?;
let retained_atom_indices = seed.retained_atom_indices;
let retained_atoms = retained_atom_indices.len();
let retained_basis = retained_atom_indices
.iter()
.map(|&atom| atom_basis[atom].clone())
.collect::<Vec<_>>();
let retained_dim = retained_atom_indices
.iter()
.map(|&atom| atom_dim[atom])
.collect::<Vec<_>>();
let term_seed = build_sae_support_term_seed(SaeSupportTermSeedRequest {
assignment: seed.assignment,
atom_basis: retained_basis,
atom_dim: retained_dim,
output_dim,
random_state: config.random_state,
})?;
let ard_precisions = (0..term_seed.term.k_atoms())
.map(|atom| vec![1.0; term_seed.term.assignment.atom_coord_dim(atom)])
.collect::<Vec<_>>();
let outer = run_sae_support_outer(SaeSupportOuterRequest {
term: term_seed.term,
target: centered.clone(),
initial_smoothness: config.initial_smoothness,
ard_precisions,
max_outer_iter: config.max_outer_iter,
max_inner_iter: config.max_inner_iter,
inner_tolerance: config.inner_tolerance,
trust_radius: config.trust_radius,
random_state: config.random_state,
})
.map_err(|error| error.to_string())?;
let curved_centered = outer.term.reconstruct()?;
let mut rss = 0.0f64;
for row in 0..n_obs {
for column in 0..output_dim {
let delta = centered[[row, column]] - curved_centered[[row, column]];
rss += delta * delta;
}
}
let tss = r0.iter().map(|value| value * value).sum::<f64>();
let explained_variance = if tss > 0.0 { 1.0 - rss / tss } else { f64::NAN };
Ok(Tier2SupportFit {
mean,
term: outer.term,
lambda_smooth: outer.lambda_smooth,
criterion: outer.criterion,
fixed_point: outer.fixed_point,
outer_certificate: outer.outer_certificate,
outer_iterations: outer.outer_iterations,
requested_atoms,
retained_atoms,
explained_variance,
})
}
fn record_support_moves(ledger: &mut SaeMigrationLedger, fit: &Tier2SupportFit) {
if fit.retained_atoms > 0 {
ledger.birth(
MoveStage::Curved,
BirthSeed::LinearAtom,
fit.retained_atoms,
Some(0),
MoveEvidence::none(),
fit.criterion,
);
}
let pruned = fit.requested_atoms - fit.retained_atoms;
if pruned > 0 {
ledger.death(
MoveStage::Curved,
MoveReason::DeadRouting,
pruned,
None,
MoveEvidence::none(),
fit.criterion,
);
}
}
#[cfg(test)]
mod fit_tests {
use super::*;
use ndarray::Array2;
#[test]
fn tiered_driver_runs_and_never_pc_reseeds() {
let n = 64;
let p = 6;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let t = i as f64 / n as f64;
z[[i, 0]] = 1.0 + (t * 6.28).cos();
z[[i, 1]] = 1.0 + (t * 6.28).sin();
z[[i, 2]] = -0.5 + (t * 3.14).cos();
z[[i, 3]] = -0.5 + (t * 3.14).sin();
}
let mut config = TieredFitConfig::linear_bulk(3, 2);
config.tier1.block_topk = 2;
config.tier1.aux_k = 3;
config.tier1.max_epochs = 200;
let report = fit_tiered(z.view(), &config).expect("tiered fit runs");
assert!(
report.explained_variance.is_finite(),
"composed EV must be finite, got {}",
report.explained_variance
);
assert_eq!(
report.ledger.pc_reseed_events, 0,
"the tiered path must never PC-reseed"
);
assert!(report.tier0.mean.iter().all(|m| m.is_finite()));
assert!(report.tier2.is_none(), "linear_bulk disables Tier-2");
assert!(
!report.tier1.convergence.certified,
"an over-complete linear-bulk fit is BEST-EFFORT (certified=false); got certified=true, frame_residual={} tol={}",
report.tier1.convergence.frame_residual, report.tier1.convergence.tolerance
);
assert!(
report.tier1.convergence.frame_residual.is_finite(),
"the open frame residual must be recorded (finite) on a best-effort certificate"
);
}
#[test]
fn tiered_returns_best_effort_open_certificate_at_k_gg_rank_2275() {
let n = 96usize;
let p = 8usize;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let t = (i as f64) * 0.2;
z[[i, 0]] = t.cos();
z[[i, 1]] = t.sin();
}
let mut config = TieredFitConfig::tiered(16, 1); config.tier1.block_topk = 4;
config.tier1.aux_k = 4; config.tier1.max_epochs = 40;
let report = fit_tiered(z.view(), &config)
.expect("#2275: best-effort tiered fit must RETURN at K ≫ rank, not error");
assert!(
!report.tier1.convergence.certified,
"K ≫ rank fit must carry an OPEN certificate; got certified=true (frame_residual={}, tol={})",
report.tier1.convergence.frame_residual, report.tier1.convergence.tolerance
);
assert!(
report.tier1.convergence.frame_residual > report.tier1.convergence.tolerance,
"an open certificate must report frame_residual above tolerance; got {} <= {}",
report.tier1.convergence.frame_residual,
report.tier1.convergence.tolerance
);
assert!(
report.tier1.convergence.ev_residual.is_finite(),
"the plateaued objective residual must be recorded (finite); got {}",
report.tier1.convergence.ev_residual
);
assert!(
report.tier1.explained_variance.is_finite(),
"best-effort Tier-1 EV must be finite"
);
assert!(
report.tier2.is_some(),
"#2275: Tier-2 must run on the best-effort Tier-1 residual"
);
assert!(
report.explained_variance.is_finite(),
"composed EV must be finite on the best-effort path"
);
assert_eq!(
report.tier1.convergence.tolerance, config.tier1.tolerance,
"#2275 must NOT soften tolerance; the open certificate uses the configured tol"
);
}
#[test]
fn block_sparse_open_fixed_point_returns_open_certificate_2275() {
use crate::sparse_dict::{
BlockSeedPolicy, BlockSparseConfig, fit_block_sparse_dictionary_with_seed,
};
let n = 96usize;
let p = 8usize;
let mut x = Array2::<f32>::zeros((n, p));
for i in 0..n {
let t = (i as f32) * 0.2;
x[[i, 0]] = t.cos();
x[[i, 1]] = t.sin();
}
let mut config = BlockSparseConfig::new(16, 1);
config.block_topk = 4;
config.aux_k = 4;
config.max_epochs = 40;
let fit = fit_block_sparse_dictionary_with_seed(
x.view(),
&config,
BlockSeedPolicy::FarthestPoint,
)
.expect("#2275: the block entry must RETURN the objective-converged open fit");
let c = &fit.convergence;
assert!(
!c.certified,
"a K ≫ rank fit must carry an OPEN certificate (certified=false); got certified=true (frame_residual={}, tol={})",
c.frame_residual, c.tolerance
);
assert!(
c.frame_residual > c.tolerance,
"an open certificate must report frame_residual above tolerance; got {} <= {}",
c.frame_residual,
c.tolerance
);
assert!(
c.ev_residual.is_finite(),
"the plateaued objective residual must be recorded (finite); got {}",
c.ev_residual
);
assert_eq!(
c.tolerance, config.tolerance,
"#2275 must NOT soften tolerance"
);
}
#[test]
fn tiered_curved_refinement_is_certified_and_records_promotions() {
let n = 240usize;
let p = 16usize;
let n_circles = 6usize;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let ph = (i as f64) * 0.261_799; for c in 0..n_circles {
let theta = ph * (1.0 + c as f64 * 0.37) + c as f64;
z[[i, 2 * c]] = theta.cos();
z[[i, 2 * c + 1]] = theta.sin();
}
let t = i as f64 / n as f64;
z[[i, 12]] = 2.0 * t - 1.0;
z[[i, 13]] = 1.0 - 2.0 * t;
z[[i, 14]] = 0.01 * (ph * 2.0).sin();
z[[i, 15]] = 0.01 * (ph * 3.0).cos();
}
let mut lin = TieredFitConfig::linear_bulk(8, 2);
lin.tier1.block_topk = 7;
lin.tier1.aux_k = 3;
lin.tier1.max_epochs = 200;
let lin_report = fit_tiered(z.view(), &lin).expect("linear-bulk fit runs");
let ev_lin = lin_report.explained_variance;
let mut tiered = TieredFitConfig::tiered(8, 2);
tiered.tier1.block_topk = 7;
tiered.tier1.aux_k = 3;
tiered.tier1.max_epochs = 200;
tiered.tier2.n_atoms = 24;
tiered.tier2.support_k = 2;
tiered.tier2.max_outer_iter = 24;
tiered.tier2.max_inner_iter = 128;
let report = fit_tiered(z.view(), &tiered).expect("tiered fit runs");
let tier2 = report.tier2.as_ref().expect("Tier-2 curved refinement ran");
assert!(
tier2.outer_certificate.certifies() && tier2.outer_certificate.is_stationary(),
"Tier-2 must carry a certifying outer stationarity certificate"
);
assert!(
tier2.fixed_point.recurred,
"Tier-2 inner fixed point must have recurred"
);
assert!(
tier2.retained_atoms >= 1 && tier2.term.k_atoms() == tier2.retained_atoms,
"Tier-2 must retain >=1 occupied curved atom (got {})",
tier2.retained_atoms
);
assert_eq!(
report.ledger.pc_reseed_events, 0,
"the tiered path must never PC-reseed"
);
assert_eq!(
report.ledger.n_births, tier2.retained_atoms,
"every retained curved atom is a promotion off the linear residual"
);
assert!(
report.explained_variance >= ev_lin - 1.0e-9,
"tiered EV {} must not regress pure-linear EV {}",
report.explained_variance,
ev_lin
);
}
#[test]
fn auto_seed_switches_at_the_farthest_point_budget() {
let small = TieredFitConfig::linear_bulk(8, 2);
assert_eq!(
small.tier1_seed.resolve(240, 16, &small.tier1),
BlockSeedPolicy::FarthestPoint,
"small-K tiered fit must keep the data-aware seed"
);
let large = TieredFitConfig::linear_bulk(2_500, 4);
assert_eq!(
large.tier1_seed.resolve(100_000, 64, &large.tier1),
BlockSeedPolicy::CoordinatePartition,
"large-K tiered fit must switch to the coordinate-partition seed"
);
let mut forced = TieredFitConfig::linear_bulk(2_500, 4);
forced.tier1_seed = TieredSeedPolicy::FarthestPoint;
assert_eq!(
forced.tier1_seed.resolve(100_000, 64, &forced.tier1),
BlockSeedPolicy::FarthestPoint
);
forced.tier1_seed = TieredSeedPolicy::CoordinatePartition;
assert_eq!(
forced.tier1_seed.resolve(240, 16, &forced.tier1),
BlockSeedPolicy::CoordinatePartition
);
}
#[test]
fn coordinate_seed_carries_a_full_tiered_fit() {
let n = 240usize;
let p = 16usize;
let n_circles = 6usize;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let ph = (i as f64) * 0.261_799;
for c in 0..n_circles {
let theta = ph * (1.0 + c as f64 * 0.37) + c as f64;
z[[i, 2 * c]] = theta.cos();
z[[i, 2 * c + 1]] = theta.sin();
}
let t = i as f64 / n as f64;
z[[i, 12]] = 2.0 * t - 1.0;
z[[i, 13]] = 1.0 - 2.0 * t;
z[[i, 14]] = 0.01 * (ph * 2.0).sin();
z[[i, 15]] = 0.01 * (ph * 3.0).cos();
}
let mut config = TieredFitConfig::tiered(8, 2);
config.tier1_seed = TieredSeedPolicy::CoordinatePartition;
config.tier1.block_topk = 7;
config.tier1.aux_k = 3;
config.tier1.max_epochs = 200;
config.tier2.n_atoms = 24;
config.tier2.support_k = 2;
config.tier2.max_outer_iter = 24;
config.tier2.max_inner_iter = 128;
let report =
fit_tiered(z.view(), &config).expect("coordinate-seeded tiered fit runs end to end");
assert!(
report.explained_variance.is_finite() && report.explained_variance > 0.0,
"coordinate-seeded composed EV must be finite and positive, got {}",
report.explained_variance
);
assert_eq!(
report.ledger.pc_reseed_events, 0,
"the coordinate-seeded tiered path must never PC-reseed"
);
let tier2 = report
.tier2
.as_ref()
.expect("tiered config must run Tier-2");
assert!(
tier2.outer_certificate.certifies(),
"the Tier-2 support-sparse refinement must return a certified fit"
);
}
#[test]
fn tier2_branch_constructs_the_support_sparse_path() {
let n = 96usize;
let p = 4usize;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let ph = (i as f64) * 0.19;
z[[i, 0]] = ph.cos();
z[[i, 1]] = ph.sin();
z[[i, 2]] = (1.7 * ph).cos();
z[[i, 3]] = (1.7 * ph).sin();
}
let mut config = TieredFitConfig::tiered(2, 2);
config.tier1.block_topk = 2;
config.tier1.aux_k = 2;
config.tier1.max_epochs = 200;
config.tier2.atom_basis = "periodic".to_string();
config.tier2.atom_dim = 1;
config.tier2.n_atoms = 8;
config.tier2.support_k = 2;
config.tier2.max_outer_iter = 32;
config.tier2.max_inner_iter = 256;
let report = fit_tiered(z.view(), &config).expect("tiny two-circle tiered fit runs");
let tier2 = report
.tier2
.as_ref()
.expect("the Tier-2 curved refinement branch must have run");
assert!(
tier2.outer_certificate.certifies() && tier2.outer_certificate.is_stationary(),
"Tier-2 must return a certifying outer stationarity certificate"
);
assert!(
tier2.fixed_point.recurred,
"Tier-2 inner fixed point must have recurred"
);
assert!(
tier2.retained_atoms >= 1
&& tier2.retained_atoms <= tier2.requested_atoms
&& tier2.term.k_atoms() == tier2.retained_atoms,
"Tier-2 must retain 1..={} occupied curved atoms (got {})",
tier2.requested_atoms,
tier2.retained_atoms
);
assert_eq!(
tier2.lambda_smooth.len(),
tier2.term.k_atoms(),
"each retained curved atom carries its selected smoothing strength"
);
assert_eq!(tier2.mean.len(), p, "the peeled residual mean spans P");
assert!(
tier2.term.atoms.iter().all(|atom| atom
.decoder_coefficients
.iter()
.all(|value| value.is_finite())),
"every retained curved atom must carry finite decoder coefficients"
);
assert_eq!(
report.ledger.n_births, tier2.retained_atoms,
"every retained curved atom is one curved birth"
);
assert_eq!(
report.ledger.pc_reseed_events, 0,
"the support-sparse Tier-2 path must never PC-reseed"
);
assert!(
report.explained_variance.is_finite(),
"composed EV must be finite, got {}",
report.explained_variance
);
}
}