mod code_space;
mod fit;
pub use code_space::{
CensusPairVerdict, CodeSpacePromotionReport, PairChartFit, fit_pair_chart,
harvest_code_space_pair_promotions, harvest_code_space_promotions, linear_distortion_floor,
};
pub use fit::{
LinearPeel, LinearPeelConfig, LinearPeelState, TieredFitConfig, TieredFitReport,
TieredSeedPolicy, fit_tiered,
};
use std::collections::{BTreeMap, BTreeSet};
use ndarray::{Array1, Array2, ArrayView2, Axis};
use crate::manifold::SaeManifoldAtom;
use crate::sparse_dict::{SparseDictConfig, SparseDictFit};
#[derive(Clone, Debug)]
pub struct Tier0Mean {
pub mean: Array1<f64>,
}
impl Tier0Mean {
pub fn fit(z: ArrayView2<'_, f64>) -> Result<Self, String> {
if z.nrows() == 0 || z.ncols() == 0 {
return Err("Tier0Mean::fit requires a non-empty (N, P) matrix".to_string());
}
let mean = z
.mean_axis(Axis(0))
.ok_or_else(|| "Tier0Mean::fit: mean_axis returned None".to_string())?;
Ok(Self { mean })
}
pub fn apply(&self, z: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
if z.ncols() != self.mean.len() {
return Err(format!(
"Tier0Mean::apply: z has P={} but μ has length {}",
z.ncols(),
self.mean.len()
));
}
Ok(&z - &self.mean.view().insert_axis(Axis(0)))
}
pub fn reconstruct(&self, recon: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
if recon.ncols() != self.mean.len() {
return Err(format!(
"Tier0Mean::reconstruct: recon has P={} but μ has length {}",
recon.ncols(),
self.mean.len()
));
}
Ok(&recon + &self.mean.view().insert_axis(Axis(0)))
}
}
#[derive(Clone, Debug)]
pub struct PerContextMean {
pub global: Array1<f64>,
pub group_means: BTreeMap<i64, Array1<f64>>,
}
impl PerContextMean {
pub fn fit(z: ArrayView2<'_, f64>, group_ids: &[i64]) -> Result<Self, String> {
let n = z.nrows();
let p = z.ncols();
if n == 0 || p == 0 {
return Err("PerContextMean::fit requires a non-empty (N, P) matrix".to_string());
}
if group_ids.len() != n {
return Err(format!(
"PerContextMean::fit: group_ids length {} != N {n}",
group_ids.len()
));
}
let global = z
.mean_axis(Axis(0))
.ok_or_else(|| "PerContextMean::fit: global mean_axis returned None".to_string())?;
let mut sums: BTreeMap<i64, (Array1<f64>, usize)> = BTreeMap::new();
for (row, &g) in z.rows().into_iter().zip(group_ids.iter()) {
let entry = sums
.entry(g)
.or_insert_with(|| (Array1::<f64>::zeros(p), 0usize));
entry.0 += &row;
entry.1 += 1;
}
let mut group_means = BTreeMap::new();
for (g, (sum, count)) in sums {
if count > 0 {
group_means.insert(g, sum / count as f64);
}
}
Ok(Self {
global,
group_means,
})
}
pub fn row_mean(&self, group: i64) -> &Array1<f64> {
self.group_means.get(&group).unwrap_or(&self.global)
}
pub fn apply(&self, z: ArrayView2<'_, f64>, group_ids: &[i64]) -> Result<Array2<f64>, String> {
if group_ids.len() != z.nrows() {
return Err(format!(
"PerContextMean::apply: group_ids length {} != N {}",
group_ids.len(),
z.nrows()
));
}
let mut out = z.to_owned();
for (mut row, &g) in out.rows_mut().into_iter().zip(group_ids.iter()) {
row -= self.row_mean(g);
}
Ok(out)
}
pub fn reconstruct(
&self,
recon: ArrayView2<'_, f64>,
group_ids: &[i64],
) -> Result<Array2<f64>, String> {
if group_ids.len() != recon.nrows() {
return Err(format!(
"PerContextMean::reconstruct: group_ids length {} != N {}",
group_ids.len(),
recon.nrows()
));
}
let mut out = recon.to_owned();
for (mut row, &g) in out.rows_mut().into_iter().zip(group_ids.iter()) {
row += self.row_mean(g);
}
Ok(out)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum SinkDelimiterClass {
Bos,
Eos,
Newline,
ChatBoundary,
Separator,
}
impl SinkDelimiterClass {
pub fn label(self) -> &'static str {
match self {
Self::Bos => "bos",
Self::Eos => "eos",
Self::Newline => "newline",
Self::ChatBoundary => "chat_boundary",
Self::Separator => "separator",
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum SinkAnchor {
Semantic,
PositionZero,
Delimiter(SinkDelimiterClass),
}
impl SinkAnchor {
pub fn label(self) -> &'static str {
match self {
Self::Semantic => "semantic_reference",
Self::PositionZero => "position_0",
Self::Delimiter(class) => class.label(),
}
}
}
#[derive(Clone, Debug)]
pub struct Tier05SinkAtomConfig {
pub enabled: bool,
pub include_position_zero: bool,
pub delimiter_classes: Vec<SinkDelimiterClass>,
}
impl Default for Tier05SinkAtomConfig {
fn default() -> Self {
Self {
enabled: false,
include_position_zero: true,
delimiter_classes: Vec::new(),
}
}
}
impl Tier05SinkAtomConfig {
pub fn disabled() -> Self {
Self {
enabled: false,
include_position_zero: true,
delimiter_classes: Vec::new(),
}
}
pub fn anchors(&self) -> Result<Vec<SinkAnchor>, String> {
if !self.enabled {
return Ok(Vec::new());
}
if !self.include_position_zero && self.delimiter_classes.is_empty() {
return Err(
"Tier05SinkAtomConfig::anchors: enabled sink atom needs position-0 or delimiter support"
.to_string(),
);
}
let mut anchors = vec![SinkAnchor::Semantic];
if self.include_position_zero {
anchors.push(SinkAnchor::PositionZero);
}
let unique: BTreeSet<SinkDelimiterClass> = self.delimiter_classes.iter().copied().collect();
for class in unique {
anchors.push(SinkAnchor::Delimiter(class));
}
Ok(anchors)
}
}
#[derive(Clone, Debug)]
pub struct Tier05SinkAtom {
pub atom: SaeManifoldAtom,
pub anchors: Vec<SinkAnchor>,
pub anchor_counts: Vec<usize>,
pub rank_charge: usize,
pub variance_absorbed: f64,
}
impl Tier05SinkAtom {
pub fn reconstruction(&self) -> Array2<f64> {
self.atom.basis_values.dot(self.atom.decoder_coefficients())
}
}
#[derive(Clone, Debug)]
pub struct TieredPrechartResidual {
pub tier0: Tier0Mean,
pub tier05_sink: Option<Tier05SinkAtom>,
pub residual: Array2<f64>,
}
#[derive(Clone, Debug)]
pub struct TieredConfig {
pub tier1: SparseDictConfig,
pub lambda_seed_rank: Option<usize>,
pub tier2_enabled: bool,
pub tier05_sink: Tier05SinkAtomConfig,
}
impl TieredConfig {
pub fn linear_bulk(k_linear: usize) -> Self {
Self {
tier1: SparseDictConfig::new(k_linear),
lambda_seed_rank: None,
tier2_enabled: false,
tier05_sink: Tier05SinkAtomConfig::disabled(),
}
}
}
#[derive(Clone, Debug)]
pub struct InterferenceSubspace {
pub q: Array2<f64>,
pub q_perp: Array2<f64>,
pub scale: Array1<f64>,
}
#[derive(Clone, Debug)]
pub struct WhitenedResidualHandoff {
pub residual: Array2<f64>,
pub interference: InterferenceSubspace,
pub tier1_decoder: Array2<f32>,
pub mean: Array1<f64>,
pub tier05_sink: Option<Tier05SinkAtom>,
}
#[derive(Clone, Debug)]
pub struct TieredSaeFit<T2> {
pub tier0: Tier0Mean,
pub tier05_sink: Option<Tier05SinkAtom>,
pub tier1: SparseDictFit,
pub tier2: Option<T2>,
pub explained_variance: f64,
}
pub fn explained_variance_from_sums(rss: f64, tss: f64) -> f64 {
if tss > 0.0 { 1.0 - rss / tss } else { f64::NAN }
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn tier0_mean_roundtrips() {
let z = array![[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]];
let t0 = Tier0Mean::fit(z.view()).expect("fit");
assert!((t0.mean[0] - 3.0).abs() < 1e-12);
assert!((t0.mean[1] - 4.0).abs() < 1e-12);
let demeaned = t0.apply(z.view()).expect("apply");
let cm = demeaned.mean_axis(Axis(0)).unwrap();
assert!(cm[0].abs() < 1e-12 && cm[1].abs() < 1e-12);
let back = t0.reconstruct(demeaned.view()).expect("reconstruct");
for (a, b) in back.iter().zip(z.iter()) {
assert!((a - b).abs() < 1e-12);
}
}
#[test]
fn per_context_mean_zeros_each_group_and_falls_back() {
let z = array![[11.0, 9.0], [9.0, 11.0], [-4.0, -6.0], [-6.0, -4.0]];
let groups = [0i64, 0, 1, 1];
let pcm = PerContextMean::fit(z.view(), &groups).expect("fit");
assert!((pcm.row_mean(0)[0] - 10.0).abs() < 1e-12);
assert!((pcm.row_mean(1)[0] + 5.0).abs() < 1e-12);
assert!((pcm.row_mean(999)[0] - pcm.global[0]).abs() < 1e-12);
let demeaned = pcm.apply(z.view(), &groups).expect("apply");
let col_sum = demeaned.sum_axis(Axis(0));
assert!(col_sum[0].abs() < 1e-12 && col_sum[1].abs() < 1e-12);
let back = pcm
.reconstruct(demeaned.view(), &groups)
.expect("reconstruct");
for (a, b) in back.iter().zip(z.iter()) {
assert!((a - b).abs() < 1e-12);
}
}
#[test]
fn explained_variance_is_undefined_when_there_is_no_variance_to_explain() {
assert!(explained_variance_from_sums(0.0, 0.0).is_nan());
assert!(explained_variance_from_sums(1.0, 0.0).is_nan());
assert!(explained_variance_from_sums(0.0, -1.0).is_nan());
}
#[test]
fn explained_variance_propagates_a_non_finite_residual() {
assert!(explained_variance_from_sums(f64::NAN, 1.0).is_nan());
assert!(explained_variance_from_sums(f64::INFINITY, 1.0).is_infinite());
}
#[test]
fn explained_variance_recovers_a_known_fraction() {
assert!((explained_variance_from_sums(0.25, 1.0) - 0.75).abs() < 1e-15);
assert!((explained_variance_from_sums(1.0, 1.0) - 0.0).abs() < 1e-15);
assert!(explained_variance_from_sums(2.0, 1.0) < 0.0);
}
}