use ndarray::{Array1, Array2, ArrayView2};
use super::weight_frame_catalog::{WeightFrameCatalog, WeightFrameSource};
use crate::frames::{
GrassmannCrossMoment, GrassmannFrame, SAE_FRAME_ACTIVATION_MARGIN,
SAE_FRAME_MIN_AUTO_OUTPUT_DIM, SAE_FRAME_RANK_CUTOFF,
};
use faer::Side;
use gam_linalg::faer_ndarray::{FaerEigh, FaerSvd, fast_ab, fast_abt};
const BYTES_PER_F64: usize = 8;
#[derive(Clone, Debug)]
pub struct InFrameCurvedConfig {
pub frame_rank_min: usize,
pub frame_rank_max: usize,
pub rank_cutoff: f64,
pub crossfit_folds: usize,
pub min_effect: f64,
pub whitening_ridge: f64,
pub min_rows: usize,
pub frame_refresh: bool,
}
impl Default for InFrameCurvedConfig {
fn default() -> Self {
Self {
frame_rank_min: 2,
frame_rank_max: 32,
rank_cutoff: SAE_FRAME_RANK_CUTOFF,
crossfit_folds: 4,
min_effect: 0.0,
whitening_ridge: 1.0e-8,
min_rows: 32,
frame_refresh: false,
}
}
}
#[derive(Clone, Debug)]
pub struct CurvedRegion {
pub rows: Vec<usize>,
pub basis_size: usize,
}
#[derive(Clone, Debug)]
pub struct RegionEvidence {
pub n_rows: usize,
pub n_effective: f64,
pub frame_rank: usize,
pub linear_loss: f64,
pub chart_loss: f64,
pub deviance_gain: f64,
pub mean_delta: f64,
pub se: f64,
pub ci_low: f64,
pub charge: f64,
pub margin: f64,
pub selected_by_bic: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ChartOccupancyStatus {
Occupied,
ChartableUnoccupied,
InsufficientOccupancy,
}
#[derive(Clone, Debug)]
pub struct RegionRecord {
pub region: usize,
pub frame_source: Option<WeightFrameSource>,
pub frame_catalog_index: Option<usize>,
pub occupancy_status: ChartOccupancyStatus,
pub basis_size: usize,
pub frame_rank: usize,
pub frame_manifold_dim: usize,
pub evidence: RegionEvidence,
pub inframe_border_coeffs: usize,
pub dense_border_coeffs: usize,
}
#[derive(Clone, Debug)]
pub struct CascadeMemoryLedger {
pub p: usize,
pub n_regions_selected: usize,
pub dense_border_coeffs: usize,
pub inframe_border_coeffs: usize,
pub dense_cov_bytes: usize,
pub inframe_cov_bytes: usize,
pub global_frame_charge: f64,
}
impl CascadeMemoryLedger {
pub fn border_shrink(&self) -> f64 {
if self.inframe_border_coeffs == 0 {
return f64::INFINITY;
}
self.dense_border_coeffs as f64 / self.inframe_border_coeffs as f64
}
pub fn cov_shrink(&self) -> f64 {
if self.inframe_cov_bytes == 0 {
return f64::INFINITY;
}
self.dense_cov_bytes as f64 / self.inframe_cov_bytes as f64
}
}
#[derive(Clone, Debug)]
pub struct InFrameCurvedResult {
pub curved_prediction: InFrameCurvedPrediction,
pub records: Vec<RegionRecord>,
pub selected_regions: Vec<usize>,
pub ledger: CascadeMemoryLedger,
}
#[derive(Clone, Debug)]
pub struct WeightFrameOccupancy {
pub frame_index: usize,
pub rows: Vec<usize>,
pub basis_size: usize,
}
#[derive(Clone, Debug)]
pub struct InFrameCurvedRegionPrediction {
pub rows: Vec<usize>,
pub frame: GrassmannFrame,
pub frame_source: Option<WeightFrameSource>,
pub fitted_coords: Array2<f64>,
}
impl InFrameCurvedRegionPrediction {
pub fn frame_rank(&self) -> usize {
self.frame.rank()
}
pub fn frame_source(&self) -> Option<&WeightFrameSource> {
self.frame_source.as_ref()
}
pub fn inframe_entries(&self) -> usize {
self.fitted_coords.len()
}
pub fn ambient_entries_if_materialized(&self) -> usize {
self.rows.len().saturating_mul(self.frame.output_dim())
}
pub fn materialize_ambient(&self) -> Array2<f64> {
fast_abt(&self.fitted_coords, &self.frame.frame().to_owned())
}
fn fill_row_into(&self, local_row: usize, out: &mut [f64]) {
let u = self.frame.frame();
let r = self.frame.rank();
for c in 0..self.frame.output_dim() {
let mut acc = 0.0;
for axis in 0..r {
acc += self.fitted_coords[[local_row, axis]] * u[[c, axis]];
}
out[c] = acc;
}
}
}
#[derive(Clone, Debug)]
pub struct InFrameCurvedPrediction {
n_rows: usize,
output_dim: usize,
regions: Vec<InFrameCurvedRegionPrediction>,
}
impl InFrameCurvedPrediction {
pub fn new(
n_rows: usize,
output_dim: usize,
regions: Vec<InFrameCurvedRegionPrediction>,
) -> Self {
Self {
n_rows,
output_dim,
regions,
}
}
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn output_dim(&self) -> usize {
self.output_dim
}
pub fn regions(&self) -> &[InFrameCurvedRegionPrediction] {
&self.regions
}
pub fn inframe_entries(&self) -> usize {
self.regions
.iter()
.map(InFrameCurvedRegionPrediction::inframe_entries)
.sum()
}
pub fn ambient_entries_if_materialized(&self) -> usize {
self.n_rows.saturating_mul(self.output_dim)
}
pub fn accepted_ambient_entries_if_eager(&self) -> usize {
self.regions
.iter()
.map(InFrameCurvedRegionPrediction::ambient_entries_if_materialized)
.sum()
}
pub fn materialize_rows(&self, rows: &[usize]) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((rows.len(), self.output_dim));
let mut row_buf = vec![0.0_f64; self.output_dim];
for (out_row, &global_row) in rows.iter().enumerate() {
for region in &self.regions {
if let Some(local_row) = region.rows.iter().position(|&row| row == global_row) {
region.fill_row_into(local_row, &mut row_buf);
for c in 0..self.output_dim {
out[[out_row, c]] = row_buf[c];
}
break;
}
}
}
out
}
pub fn materialize_ambient(&self) -> Array2<f64> {
let rows: Vec<usize> = (0..self.n_rows).collect();
self.materialize_rows(&rows)
}
}
pub fn fit_inframe_curved_regions(
residual: ArrayView2<'_, f64>,
regions: &[CurvedRegion],
n_tokens_total: usize,
config: &InFrameCurvedConfig,
) -> Result<InFrameCurvedResult, String> {
let p = residual.ncols();
if p == 0 {
return Err("fit_inframe_curved_regions: residual must have p >= 1".to_string());
}
if config.frame_rank_min == 0 || config.frame_rank_max < config.frame_rank_min {
return Err(
"fit_inframe_curved_regions: require 1 <= frame_rank_min <= frame_rank_max".to_string(),
);
}
let mut fits: Vec<Option<RegionFit>> = Vec::with_capacity(regions.len());
for region in regions {
fits.push(fit_one_region(residual, region, config)?);
}
let mut prediction_regions = Vec::new();
let mut records = Vec::with_capacity(regions.len());
let mut selected_regions = Vec::new();
let mut dense_border = 0usize;
let mut inframe_border = 0usize;
let mut dense_cov = 0usize;
let mut inframe_cov = 0usize;
let mut frame_charge = 0.0f64;
let ln_n = (n_tokens_total.max(2) as f64).ln();
for (i, (region, fit)) in regions.iter().zip(fits.iter()).enumerate() {
let Some(fit) = fit else {
continue;
};
let m = region.basis_size;
let r = fit.frame_rank;
let inframe_border_coeffs = m.saturating_mul(r);
let dense_border_coeffs = m.saturating_mul(p);
let manifold_dim = r.saturating_mul(p.saturating_sub(r));
records.push(RegionRecord {
region: i,
frame_source: None,
frame_catalog_index: None,
occupancy_status: ChartOccupancyStatus::Occupied,
basis_size: m,
frame_rank: r,
frame_manifold_dim: manifold_dim,
evidence: fit.evidence.clone(),
inframe_border_coeffs,
dense_border_coeffs,
});
if !fit.evidence.selected_by_bic {
continue;
}
selected_regions.push(i);
dense_border += dense_border_coeffs;
inframe_border += inframe_border_coeffs;
dense_cov += dense_border_coeffs
.saturating_mul(dense_border_coeffs)
.saturating_mul(BYTES_PER_F64);
inframe_cov += inframe_border_coeffs
.saturating_mul(inframe_border_coeffs)
.saturating_mul(BYTES_PER_F64);
frame_charge += 0.5 * manifold_dim as f64 * ln_n;
prediction_regions.push(InFrameCurvedRegionPrediction {
rows: region.rows.clone(),
frame: fit.frame.clone(),
frame_source: None,
fitted_coords: fit.fitted_coords.clone(),
});
}
let ledger = CascadeMemoryLedger {
p,
n_regions_selected: selected_regions.len(),
dense_border_coeffs: dense_border,
inframe_border_coeffs: inframe_border,
dense_cov_bytes: dense_cov,
inframe_cov_bytes: inframe_cov,
global_frame_charge: frame_charge,
};
Ok(InFrameCurvedResult {
curved_prediction: InFrameCurvedPrediction::new(residual.nrows(), p, prediction_regions),
records,
selected_regions,
ledger,
})
}
pub fn fit_inframe_curved_weight_frame_catalog(
residual: ArrayView2<'_, f64>,
catalog: &WeightFrameCatalog,
occupancies: &[WeightFrameOccupancy],
n_tokens_total: usize,
config: &InFrameCurvedConfig,
) -> Result<InFrameCurvedResult, String> {
let p = residual.ncols();
if p == 0 {
return Err(
"fit_inframe_curved_weight_frame_catalog: residual must have p >= 1".to_string(),
);
}
if catalog.output_dim() != p {
return Err(format!(
"fit_inframe_curved_weight_frame_catalog: catalog output dim {} != residual dim {p}",
catalog.output_dim()
));
}
let mut fits: Vec<Option<RegionFit>> = Vec::with_capacity(occupancies.len());
let mut statuses = Vec::with_capacity(occupancies.len());
for occupancy in occupancies {
let entry = catalog.entry(occupancy.frame_index).ok_or_else(|| {
format!(
"fit_inframe_curved_weight_frame_catalog: frame index {} out of range",
occupancy.frame_index
)
})?;
if occupancy.rows.is_empty() {
fits.push(None);
statuses.push(ChartOccupancyStatus::ChartableUnoccupied);
continue;
}
let region = CurvedRegion {
rows: occupancy.rows.clone(),
basis_size: occupancy.basis_size,
};
let fit = fit_one_region_in_frame(residual, ®ion, &entry.frame, config)?;
let status = if fit.is_some() {
ChartOccupancyStatus::Occupied
} else {
ChartOccupancyStatus::InsufficientOccupancy
};
fits.push(fit);
statuses.push(status);
}
let mut prediction_regions = Vec::new();
let mut records = Vec::with_capacity(occupancies.len());
let mut selected_regions = Vec::new();
let mut dense_border = 0usize;
let mut inframe_border = 0usize;
let mut dense_cov = 0usize;
let mut inframe_cov = 0usize;
let mut frame_charge = 0.0f64;
let ln_n = (n_tokens_total.max(2) as f64).ln();
for (i, occupancy) in occupancies.iter().enumerate() {
let entry = catalog
.entry(occupancy.frame_index)
.expect("validated catalog index");
let r = entry.frame.rank();
let m = occupancy.basis_size;
let inframe_border_coeffs = m.saturating_mul(r);
let dense_border_coeffs = m.saturating_mul(p);
let manifold_dim = r.saturating_mul(p.saturating_sub(r));
let evidence = fits[i]
.as_ref()
.map(|fit| fit.evidence.clone())
.unwrap_or_else(|| empty_evidence(occupancy.rows.len(), r));
records.push(RegionRecord {
region: i,
frame_source: Some(entry.source.clone()),
frame_catalog_index: Some(occupancy.frame_index),
occupancy_status: statuses[i].clone(),
basis_size: m,
frame_rank: r,
frame_manifold_dim: manifold_dim,
evidence,
inframe_border_coeffs,
dense_border_coeffs,
});
let Some(fit) = fits[i].as_ref() else {
continue;
};
if !fit.evidence.selected_by_bic {
continue;
}
selected_regions.push(i);
dense_border += dense_border_coeffs;
inframe_border += inframe_border_coeffs;
dense_cov += dense_border_coeffs
.saturating_mul(dense_border_coeffs)
.saturating_mul(BYTES_PER_F64);
inframe_cov += inframe_border_coeffs
.saturating_mul(inframe_border_coeffs)
.saturating_mul(BYTES_PER_F64);
frame_charge += 0.5 * manifold_dim as f64 * ln_n;
prediction_regions.push(InFrameCurvedRegionPrediction {
rows: occupancy.rows.clone(),
frame: fit.frame.clone(),
frame_source: Some(entry.source.clone()),
fitted_coords: fit.fitted_coords.clone(),
});
}
let n_regions_selected = selected_regions.len();
Ok(InFrameCurvedResult {
curved_prediction: InFrameCurvedPrediction::new(residual.nrows(), p, prediction_regions),
records,
selected_regions,
ledger: CascadeMemoryLedger {
p,
n_regions_selected,
dense_border_coeffs: dense_border,
inframe_border_coeffs: inframe_border,
dense_cov_bytes: dense_cov,
inframe_cov_bytes: inframe_cov,
global_frame_charge: frame_charge,
},
})
}
pub fn inframe_curved_region_prediction(
residual: ArrayView2<'_, f64>,
rows: &[usize],
config: &InFrameCurvedConfig,
) -> Result<Option<InFrameCurvedRegionPrediction>, String> {
if rows.len() < 2 {
return Ok(None);
}
let r_g = take_rows(residual, rows);
let Some(frame) = learn_frame(&r_g, config)? else {
return Ok(None);
};
let z = fast_ab(&r_g, &frame.frame().to_owned());
let fitted = fit_radial_all(&z, config.whitening_ridge)?;
Ok(Some(InFrameCurvedRegionPrediction {
rows: rows.to_vec(),
frame,
frame_source: None,
fitted_coords: fitted,
}))
}
pub fn residual_span_frame(
residual: ArrayView2<'_, f64>,
rows: &[usize],
config: &InFrameCurvedConfig,
) -> Result<Option<GrassmannFrame>, String> {
if rows.len() < 2 {
return Ok(None);
}
let r_g = take_rows(residual, rows);
learn_frame(&r_g, config)
}
pub fn activate_residual_frame(
atom: &mut crate::manifold::SaeManifoldAtom,
residual: ArrayView2<'_, f64>,
rows: &[usize],
config: &InFrameCurvedConfig,
) -> Result<Option<usize>, String> {
let p = residual.ncols();
if p < SAE_FRAME_MIN_AUTO_OUTPUT_DIM {
atom.decoder_frame = None;
return Ok(None);
}
let Some(frame) = residual_span_frame(residual, rows, config)? else {
atom.decoder_frame = None;
return Ok(None);
};
let r = frame.rank();
if (r as f64) > (p as f64) * (1.0 - SAE_FRAME_ACTIVATION_MARGIN) {
atom.decoder_frame = None;
return Ok(None);
}
let u = frame.frame().to_owned(); let c_proj = fast_ab(&atom.decoder_coefficients, &u); atom.decoder_coefficients = fast_abt(&c_proj, &u); atom.decoder_frame = Some(frame);
Ok(Some(r))
}
struct RegionFit {
frame: GrassmannFrame,
frame_rank: usize,
fitted_coords: Array2<f64>,
evidence: RegionEvidence,
}
fn fit_one_region(
residual: ArrayView2<'_, f64>,
region: &CurvedRegion,
config: &InFrameCurvedConfig,
) -> Result<Option<RegionFit>, String> {
let n_g = region.rows.len();
if n_g < config.min_rows || n_g < 2 * config.crossfit_folds.max(2) {
return Ok(None);
}
let r_g = take_rows(residual, ®ion.rows);
let frame = match learn_frame(&r_g, config)? {
Some(f) => f,
None => return Ok(None),
};
let r = frame.rank();
let mut z = fast_ab(&r_g, &frame.frame().to_owned());
let evidence = crossfit_evidence(&z, config, n_g)?;
let mut frame = frame;
if config.frame_refresh {
if let Some((refreshed, refreshed_z, refreshed_ev)) =
try_frame_refresh(&r_g, &z, &frame, config, n_g)?
{
if refreshed_ev.deviance_gain > evidence.deviance_gain {
frame = refreshed;
z = refreshed_z;
let fitted_coords = fit_radial_all(&z, config.whitening_ridge)?;
return Ok(Some(RegionFit {
frame,
frame_rank: r,
fitted_coords,
evidence: refreshed_ev,
}));
}
}
}
let fitted_coords = fit_radial_all(&z, config.whitening_ridge)?;
Ok(Some(RegionFit {
frame,
frame_rank: r,
fitted_coords,
evidence,
}))
}
fn fit_one_region_in_frame(
residual: ArrayView2<'_, f64>,
region: &CurvedRegion,
frame: &GrassmannFrame,
config: &InFrameCurvedConfig,
) -> Result<Option<RegionFit>, String> {
let n_g = region.rows.len();
if n_g < config.min_rows || n_g < 2 * config.crossfit_folds.max(2) {
return Ok(None);
}
if frame.output_dim() != residual.ncols() {
return Err(format!(
"fit_one_region_in_frame: frame output dim {} != residual dim {}",
frame.output_dim(),
residual.ncols()
));
}
let r_g = take_rows(residual, ®ion.rows);
let z = fast_ab(&r_g, &frame.frame().to_owned());
let evidence = crossfit_evidence(&z, config, n_g)?;
let fitted_coords = fit_radial_all(&z, config.whitening_ridge)?;
Ok(Some(RegionFit {
frame: frame.clone(),
frame_rank: frame.rank(),
fitted_coords,
evidence,
}))
}
fn empty_evidence(n_rows: usize, frame_rank: usize) -> RegionEvidence {
RegionEvidence {
n_rows,
n_effective: n_rows as f64,
frame_rank,
linear_loss: 0.0,
chart_loss: 0.0,
deviance_gain: 0.0,
mean_delta: 0.0,
se: f64::INFINITY,
ci_low: f64::NEG_INFINITY,
charge: 0.0,
margin: 0.0,
selected_by_bic: false,
}
}
fn learn_frame(
r_g: &Array2<f64>,
config: &InFrameCurvedConfig,
) -> Result<Option<GrassmannFrame>, String> {
let (n_g, p) = r_g.dim();
let (_u, sv, vt_opt) = r_g
.svd(false, true)
.map_err(|e| format!("inframe learn_frame: SVD failed: {e}"))?;
let vt =
vt_opt.ok_or_else(|| "inframe learn_frame: SVD returned no right factor".to_string())?;
let max_sv = sv.iter().copied().fold(0.0_f64, f64::max);
if !(max_sv > 0.0) {
return Ok(None);
}
let tol = config.rank_cutoff * max_sv;
let numerical_rank = sv.iter().filter(|&&v| v > tol).count();
let available = vt.nrows().min(n_g).min(p.saturating_sub(1));
let r = numerical_rank
.max(config.frame_rank_min)
.min(config.frame_rank_max)
.min(available);
if r == 0 || p.saturating_sub(r) == 0 {
return Ok(None);
}
let mut frame = Array2::<f64>::zeros((p, r));
for col in 0..r {
for row in 0..p {
frame[[row, col]] = vt[[col, row]];
}
}
let mut gauge = Array1::<f64>::zeros(r);
for i in 0..r {
gauge[i] = sv.get(i).copied().unwrap_or(0.0);
}
Ok(Some(GrassmannFrame::from_oriented(frame, gauge)))
}
fn try_frame_refresh(
r_g: &Array2<f64>,
z: &Array2<f64>,
frame: &GrassmannFrame,
config: &InFrameCurvedConfig,
n_g: usize,
) -> Result<Option<(GrassmannFrame, Array2<f64>, RegionEvidence)>, String> {
let (p, r) = (frame.output_dim(), frame.rank());
let mut cross = GrassmannCrossMoment::new(p, r);
cross.accumulate(r_g.view(), z.view())?;
let refreshed = cross.polar_frame()?;
let z_new = fast_ab(r_g, &refreshed.frame().to_owned());
let ev = crossfit_evidence(&z_new, config, n_g)?;
Ok(Some((refreshed, z_new, ev)))
}
fn crossfit_evidence(
z: &Array2<f64>,
config: &InFrameCurvedConfig,
n_g: usize,
) -> Result<RegionEvidence, String> {
let n = z.nrows();
let r = z.ncols();
let folds = config.crossfit_folds.max(2).min(n);
let mut linear_loss = vec![0.0f64; n];
let mut chart_loss = vec![0.0f64; n];
for fold in 0..folds {
let mut train = Vec::new();
let mut eval = Vec::new();
for i in 0..n {
if i % folds == fold {
eval.push(i);
} else {
train.push(i);
}
}
if train.len() < 2 || eval.is_empty() {
continue;
}
let train_z = take_index(z, &train);
let eval_z = take_index(z, &eval);
let whitening = Whitening::fit(&train_z, config.whitening_ridge)?;
let train_w = whitening.transform(&train_z);
let eval_w = whitening.transform(&eval_z);
let linear_pred_w = pca_rank1_reconstruct(&train_w, &eval_w)?;
let chart_pred_w = radial_predict(&train_w, &eval_w);
let linear_pred = whitening.inverse(&linear_pred_w);
let chart_pred = whitening.inverse(&chart_pred_w);
for (pos, &row) in eval.iter().enumerate() {
linear_loss[row] = row_sse(&eval_z, &linear_pred, pos);
chart_loss[row] = row_sse(&eval_z, &chart_pred, pos);
}
}
let delta: Vec<f64> = linear_loss
.iter()
.zip(chart_loss.iter())
.map(|(l, c)| l - c)
.collect();
let n_eff = autocorr_ess(&delta);
let mean_delta = delta.iter().sum::<f64>() / n as f64;
let se = newey_west_se(&delta);
let ci_low = mean_delta - 1.959963984540054 * se;
let linear_total: f64 = linear_loss.iter().sum();
let chart_total: f64 = chart_loss.iter().sum();
let sse_floor = 1.0e-12 * (linear_total.max(chart_total)).max(1.0e-300);
let gain = 0.5
* (r.max(1) as f64)
* n_eff.max(2.0)
* (linear_total.max(sse_floor) / chart_total.max(sse_floor)).ln();
let d_eff = (2 * r).max(1) as f64;
let charge = 0.5 * d_eff * n_eff.max(2.0).ln();
let margin = gain - charge;
Ok(RegionEvidence {
n_rows: n_g,
n_effective: n_eff,
frame_rank: r,
linear_loss: linear_total,
chart_loss: chart_total,
deviance_gain: gain,
mean_delta,
se,
ci_low,
charge,
margin,
selected_by_bic: margin >= config.min_effect && ci_low > 0.0,
})
}
fn fit_radial_all(z: &Array2<f64>, ridge: f64) -> Result<Array2<f64>, String> {
let whitening = Whitening::fit(z, ridge)?;
let w = whitening.transform(z);
let pred_w = radial_predict(&w, &w);
Ok(whitening.inverse(&pred_w))
}
pub fn dense_ambient_radial_reference(
r_g: ArrayView2<'_, f64>,
ridge: f64,
) -> Result<Array2<f64>, String> {
let owned = r_g.to_owned();
fit_radial_all(&owned, ridge)
}
struct Whitening {
mean: Vec<f64>,
eigvec: Array2<f64>,
scale: Vec<f64>,
dim: usize,
}
impl Whitening {
fn fit(x: &Array2<f64>, ridge: f64) -> Result<Self, String> {
let n = x.nrows();
let d = x.ncols();
let mut mean = vec![0.0; d];
for j in 0..d {
for i in 0..n {
mean[j] += x[[i, j]];
}
mean[j] /= n.max(1) as f64;
}
let mut cov = Array2::<f64>::zeros((d, d));
for i in 0..n {
for a in 0..d {
let va = x[[i, a]] - mean[a];
for b in 0..d {
cov[[a, b]] += va * (x[[i, b]] - mean[b]);
}
}
}
let denom = (n.saturating_sub(1)).max(1) as f64;
cov.mapv_inplace(|v| v / denom);
let (vals, eigvec) = cov
.eigh(Side::Lower)
.map_err(|e| format!("inframe whitening eigh failed: {e}"))?;
if !(ridge.is_finite() && ridge >= 0.0) {
return Err(format!(
"inframe whitening ridge fraction must be finite and non-negative; got {ridge}"
));
}
let max_eval = vals.iter().copied().fold(0.0, f64::max);
let spectral_scale = max_eval.max(f64::MIN_POSITIVE);
let floor = (ridge * spectral_scale).max(f64::EPSILON * spectral_scale);
let scale = vals.iter().map(|&v| (v.max(0.0) + floor).sqrt()).collect();
Ok(Self {
mean,
eigvec,
scale,
dim: d,
})
}
fn transform(&self, x: &Array2<f64>) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((x.nrows(), self.dim));
for i in 0..x.nrows() {
for k in 0..self.dim {
let mut v = 0.0;
for j in 0..self.dim {
v += (x[[i, j]] - self.mean[j]) * self.eigvec[[j, k]];
}
out[[i, k]] = v / self.scale[k];
}
}
out
}
fn inverse(&self, x: &Array2<f64>) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((x.nrows(), self.dim));
for i in 0..x.nrows() {
for j in 0..self.dim {
let mut v = self.mean[j];
for k in 0..self.dim {
v += x[[i, k]] * self.scale[k] * self.eigvec[[j, k]];
}
out[[i, j]] = v;
}
}
out
}
}
fn pca_rank1_reconstruct(train: &Array2<f64>, eval: &Array2<f64>) -> Result<Array2<f64>, String> {
let n = train.nrows();
let d = train.ncols();
let mut mean = vec![0.0; d];
for j in 0..d {
for i in 0..n {
mean[j] += train[[i, j]];
}
mean[j] /= n.max(1) as f64;
}
let mut cov = Array2::<f64>::zeros((d, d));
for i in 0..n {
for a in 0..d {
let va = train[[i, a]] - mean[a];
for b in 0..d {
cov[[a, b]] += va * (train[[i, b]] - mean[b]);
}
}
}
let denom = n.saturating_sub(1).max(1) as f64;
cov.mapv_inplace(|v| v / denom);
let (vals, eigvec) = cov
.eigh(Side::Lower)
.map_err(|e| format!("inframe pca eigh failed: {e}"))?;
let mut top = 0usize;
let mut top_val = f64::NEG_INFINITY;
for (k, &v) in vals.iter().enumerate() {
if v > top_val {
top_val = v;
top = k;
}
}
let mut out = Array2::<f64>::zeros((eval.nrows(), d));
for i in 0..eval.nrows() {
let mut score = 0.0;
for j in 0..d {
score += (eval[[i, j]] - mean[j]) * eigvec[[j, top]];
}
for j in 0..d {
out[[i, j]] = mean[j] + score * eigvec[[j, top]];
}
}
Ok(out)
}
fn radial_predict(train: &Array2<f64>, eval: &Array2<f64>) -> Array2<f64> {
let d = train.ncols();
let mut radius = 0.0;
for i in 0..train.nrows() {
let mut ss = 0.0;
for j in 0..d {
ss += train[[i, j]] * train[[i, j]];
}
radius += ss.sqrt();
}
radius /= train.nrows().max(1) as f64;
let mut out = Array2::<f64>::zeros(eval.dim());
for i in 0..eval.nrows() {
let mut norm = 0.0;
for j in 0..d {
norm += eval[[i, j]] * eval[[i, j]];
}
norm = norm.sqrt().max(1.0e-12);
for j in 0..d {
out[[i, j]] = radius * eval[[i, j]] / norm;
}
}
out
}
fn take_rows(x: ArrayView2<'_, f64>, rows: &[usize]) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((rows.len(), x.ncols()));
for (i, &row) in rows.iter().enumerate() {
for j in 0..x.ncols() {
out[[i, j]] = x[[row, j]];
}
}
out
}
fn take_index(x: &Array2<f64>, rows: &[usize]) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((rows.len(), x.ncols()));
for (i, &row) in rows.iter().enumerate() {
for j in 0..x.ncols() {
out[[i, j]] = x[[row, j]];
}
}
out
}
fn row_sse(a: &Array2<f64>, b: &Array2<f64>, row: usize) -> f64 {
let mut s = 0.0;
for j in 0..a.ncols() {
let d = a[[row, j]] - b[[row, j]];
s += d * d;
}
s
}
fn autocorr_ess(x: &[f64]) -> f64 {
let n = x.len();
if n <= 1 {
return n as f64;
}
let mean = x.iter().sum::<f64>() / n as f64;
let var = x.iter().map(|v| (v - mean) * (v - mean)).sum::<f64>() / n as f64;
if var <= 0.0 {
return n as f64;
}
let lag_cap = (n as f64).sqrt() as usize;
let mut rho_sum = 0.0;
for lag in 1..=lag_cap.max(1).min(n - 1) {
let mut cov = 0.0;
for i in lag..n {
cov += (x[i] - mean) * (x[i - lag] - mean);
}
cov /= (n - lag) as f64;
let rho = cov / var;
if rho <= 0.0 || !rho.is_finite() {
break;
}
rho_sum += rho;
}
(n as f64 / (1.0 + 2.0 * rho_sum)).max(1.0)
}
fn newey_west_se(x: &[f64]) -> f64 {
let n = x.len();
if n <= 1 {
return f64::INFINITY;
}
let mean = x.iter().sum::<f64>() / n as f64;
let lag_cap = (n as f64).sqrt() as usize;
let mut gamma0 = 0.0;
for v in x {
gamma0 += (v - mean) * (v - mean);
}
gamma0 /= n as f64;
let mut var = gamma0;
for lag in 1..=lag_cap.max(1).min(n - 1) {
let mut gamma = 0.0;
for i in lag..n {
gamma += (x[i] - mean) * (x[i - lag] - mean);
}
gamma /= n as f64;
let w = 1.0 - lag as f64 / (lag_cap as f64 + 1.0);
var += 2.0 * w * gamma;
}
(var.max(0.0) / n as f64).sqrt()
}