use ndarray::ArrayView2;
use crate::migration_ledger::{BirthSeed, MoveEvidence, MoveReason, MoveStage, SaeMigrationLedger};
use crate::sparse_dict::{
BlockSparseConfig, BlockSparseFit, CofitConfig, CofitReport, cofit_block_and_curved,
fit_block_sparse_dictionary,
};
use crate::tiered::Tier0Mean;
#[derive(Clone, Debug)]
pub struct TieredFitConfig {
pub tier1: BlockSparseConfig,
pub tier2_enabled: bool,
pub cofit: CofitConfig,
}
impl TieredFitConfig {
pub fn linear_bulk(n_blocks: usize, block_size: usize) -> Self {
Self {
tier1: BlockSparseConfig::new(n_blocks, block_size),
tier2_enabled: false,
cofit: CofitConfig::default(),
}
}
pub fn tiered(n_blocks: usize, block_size: usize) -> Self {
Self {
tier1: BlockSparseConfig::new(n_blocks, block_size),
tier2_enabled: true,
cofit: CofitConfig::default(),
}
}
}
#[derive(Clone, Debug)]
pub struct TieredFitReport {
pub tier0: Tier0Mean,
pub tier1: BlockSparseFit,
pub tier2: Option<CofitReport>,
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 tier1 = fit_block_sparse_dictionary(r0_f32.view(), &config.tier1)?;
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 report = cofit_block_and_curved(
r0_f32.view(),
tier1.decoder.view(),
tier1.blocks.view(),
tier1.codes.view(),
tier1.gamma,
&config.cofit,
)?;
record_cofit_moves(&mut ledger, &report);
let ev = report.explained_variance;
(Some(report), ev)
} else {
(None, tier1.explained_variance)
};
Ok(TieredFitReport {
tier0,
tier1,
tier2,
ledger,
explained_variance,
})
}
fn record_cofit_moves(ledger: &mut SaeMigrationLedger, report: &CofitReport) {
let mut prev_accepted = 0usize;
for round in &report.rounds {
let accepted = round.n_accepted_charts;
if accepted > prev_accepted {
let count = accepted - prev_accepted;
ledger.birth(
MoveStage::Curved,
BirthSeed::LinearAtom,
count,
Some(round.round),
MoveEvidence::from_dl_bits(round.curved_charge),
round.objective,
);
} else if accepted < prev_accepted {
let count = prev_accepted - accepted;
ledger.refuse(
MoveStage::Curved,
MoveReason::EvidenceInsufficient,
count,
Some(round.round),
MoveEvidence::none(),
round.objective,
);
} else if round.round > 0 && !round.curved_committed {
ledger.refuse(
MoveStage::Curved,
MoveReason::EvidenceInsufficient,
1,
Some(round.round),
MoveEvidence::none(),
round.objective,
);
}
prev_accepted = accepted;
}
}
#[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.max_epochs = 8;
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");
}
#[test]
fn tiered_beats_linear_on_six_circle_mixture_and_records_a_promotion() {
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.max_epochs = 20;
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.max_epochs = 20;
let report = fit_tiered(z.view(), &tiered).expect("tiered fit runs");
assert_eq!(
report.ledger.pc_reseed_events, 0,
"the tiered path must never PC-reseed"
);
assert!(
report.ledger.n_births >= 1,
"the migration ledger must record >=1 curved birth; got {} (moves: {:?})",
report.ledger.n_births,
report.ledger.moves
);
assert!(
report.explained_variance > ev_lin,
"tiered EV {} must beat pure-linear EV {}",
report.explained_variance,
ev_lin
);
}
}