use ndarray::{Array2, ArrayView2, ArrayView3};
use super::block::{
block_sparse_dictionary_block_coords, block_sparse_dictionary_project_residual,
block_sparse_dictionary_project_residual_with_codes, reconstruct_block_sparse_rows,
};
#[derive(Clone, Debug)]
pub struct BlockChartComposeConfig {
pub block_size: usize,
pub block_topk: usize,
pub gamma: f32,
pub residual_target: bool,
pub min_firings: usize,
pub max_blocks: usize,
pub crossfit_folds: usize,
pub min_effect: f64,
pub whitening_ridge: f64,
pub pair_screen: bool,
pub pair_top_blocks: usize,
pub max_pairs: usize,
pub pair_min_cofirings: usize,
pub pair_min_score: f64,
pub block_tile: usize,
}
pub const CHART_FDR_ALPHA: f64 = 0.1;
impl Default for BlockChartComposeConfig {
fn default() -> Self {
Self {
block_size: 4,
block_topk: 32,
gamma: 1.0,
residual_target: true,
min_firings: 64,
max_blocks: 256,
crossfit_folds: 2,
min_effect: 0.0,
whitening_ridge: 1.0e-8,
pair_screen: true,
pair_top_blocks: 64,
max_pairs: 128,
pair_min_cofirings: 64,
pair_min_score: 0.20,
block_tile: 1024,
}
}
}
#[derive(Clone, Debug)]
pub struct ChartEvidence {
pub n_rows: usize,
pub n_effective: f64,
pub linear_loss: f64,
pub chart_loss: f64,
pub deviance_gain: f64,
pub mean_delta: f64,
pub se: f64,
pub ci_low: f64,
pub ci_high: f64,
pub charge: f64,
pub margin: f64,
pub selected_by_bic: bool,
pub log_e: f64,
pub fdr_selected: bool,
}
#[derive(Clone, Debug)]
pub struct BlockChartRecord {
pub block0: usize,
pub block1: Option<usize>,
pub screen_score: f64,
pub evidence: ChartEvidence,
}
#[derive(Clone, Debug)]
pub struct BlockChartComposeResult {
pub reconstructed: Array2<f32>,
pub block_records: Vec<BlockChartRecord>,
pub pair_records: Vec<BlockChartRecord>,
pub selected_blocks: Vec<usize>,
pub selected_chart_blocks: Vec<usize>,
pub selected_chart_pairs: Vec<(usize, usize)>,
pub fdr_alpha: f64,
pub fdr_selected_chart_blocks: Vec<usize>,
pub fdr_selected_chart_pairs: Vec<(usize, usize)>,
}
#[derive(Clone, Debug)]
pub struct BlockSeedManifestConfig {
pub block_size: usize,
pub block_topk: usize,
pub gamma: f32,
pub residual_target: bool,
pub n_basis_chart: usize,
pub include_bases: bool,
pub name_prefix: String,
pub block_tile: usize,
}
#[derive(Clone, Debug)]
pub struct MdlFeaturizerRow {
pub name: String,
pub kind: String,
pub total_var: f64,
pub n_tokens: usize,
pub n_firings: usize,
pub n_params: usize,
pub coded_var: Vec<f64>,
pub g_dict: usize,
pub k_active: usize,
pub block_name: Option<String>,
pub chart_name: Option<String>,
}
#[derive(Clone, Debug)]
pub struct BlockSeedRecord {
pub block: usize,
pub block_dim: usize,
pub n_firings: usize,
pub utilization: f32,
pub stable_rank: f32,
pub coded_var: Vec<f64>,
pub total_var: f64,
pub block_linear_ev: f64,
pub basis: Option<Vec<Vec<f32>>>,
pub mdl_block: MdlFeaturizerRow,
pub mdl_chart: MdlFeaturizerRow,
pub matched_dl_flat: crate::description_length::MatchedDl,
pub matched_dl_chart: crate::description_length::MatchedDl,
pub matched_dl_delta_bits: f64,
}
#[derive(Clone, Debug)]
pub struct BlockSeedManifest {
pub n_blocks: usize,
pub block_size: usize,
pub block_topk: usize,
pub ambient_p: usize,
pub gamma: f32,
pub explained_variance: f64,
pub residual_target: bool,
pub n_basis_chart: usize,
pub blocks: Vec<BlockSeedRecord>,
}
#[derive(Clone, Debug)]
struct Whitening {
mean: Vec<f64>,
eigvec: Vec<f64>,
scale: Vec<f64>,
dim: usize,
}
#[derive(Clone, Debug)]
struct CandidateFit {
rows: Vec<usize>,
blocks: Vec<usize>,
fitted_coords: Array2<f64>,
screen_score: f64,
evidence: ChartEvidence,
}
pub fn compose_block_coordinate_charts(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
config: &BlockChartComposeConfig,
) -> Result<BlockChartComposeResult, String> {
validate_inputs(x, decoder, blocks, codes, config)?;
let base = reconstruct_block_sparse_rows(decoder, blocks, codes, config.block_size)?;
let selected = select_blocks(x, decoder, blocks, codes, config)?;
let mut singles = Vec::new();
for &g in &selected {
let rows = rows_for_block(blocks, g);
let coords_all = block_coords_for_config(x, decoder, blocks, codes, config, g)?;
let coords = take_rows(&coords_all, &rows);
let evidence = crossfit_evidence(&coords, config)?;
let fitted_coords = fit_radial_chart_all(&coords, config.whitening_ridge)?;
singles.push(CandidateFit {
rows,
blocks: vec![g],
fitted_coords,
screen_score: 1.0,
evidence,
});
}
let mut pairs = Vec::new();
if config.pair_screen {
let screens = screen_pairs(x, decoder, blocks, codes, config, &selected)?;
for (g0, g1, score) in screens {
if score < config.pair_min_score {
continue;
}
let rows = rows_for_pair(blocks, g0, g1);
let coords0_all = block_coords_for_config(x, decoder, blocks, codes, config, g0)?;
let coords1_all = block_coords_for_config(x, decoder, blocks, codes, config, g1)?;
let coords0 = take_rows(&coords0_all, &rows);
let coords1 = take_rows(&coords1_all, &rows);
let coords = hstack(&coords0, &coords1);
let evidence = crossfit_evidence(&coords, config)?;
let fitted_coords = fit_radial_chart_all(&coords, config.whitening_ridge)?;
pairs.push(CandidateFit {
rows,
blocks: vec![g0, g1],
fitted_coords,
screen_score: score,
evidence,
});
}
}
let mut reconstructed = base;
let mut replaced = vec![vec![false; decoder.nrows() / config.block_size]; x.nrows()];
for pair in &pairs {
if !pair.evidence.selected_by_bic {
continue;
}
let b = config.block_size;
for (local_row, &row) in pair.rows.iter().enumerate() {
for (slot, &g) in pair.blocks.iter().enumerate() {
subtract_block_contribution(
&mut reconstructed,
decoder,
blocks,
codes,
config.block_size,
row,
g,
);
add_lifted_coords(
&mut reconstructed,
decoder,
config.block_size,
row,
g,
pair.fitted_coords
.row(local_row)
.as_slice()
.ok_or("non-contiguous pair row")?,
slot * b,
);
replaced[row][g] = true;
}
}
}
for single in &singles {
if !single.evidence.selected_by_bic {
continue;
}
let g = single.blocks[0];
for (local_row, &row) in single.rows.iter().enumerate() {
if replaced[row][g] {
continue;
}
subtract_block_contribution(
&mut reconstructed,
decoder,
blocks,
codes,
config.block_size,
row,
g,
);
add_lifted_coords(
&mut reconstructed,
decoder,
config.block_size,
row,
g,
single
.fitted_coords
.row(local_row)
.as_slice()
.ok_or("non-contiguous single row")?,
0,
);
}
}
let mut family_log_e: Vec<f64> = Vec::with_capacity(singles.len() + pairs.len());
family_log_e.extend(singles.iter().map(|c| c.evidence.log_e));
family_log_e.extend(pairs.iter().map(|c| c.evidence.log_e));
let fdr = super::split_lr_fdr::family_fdr_certificate(family_log_e, CHART_FDR_ALPHA)
.map_err(|error| format!("block chart e-BH certificate: {error}"))?;
for &idx in &fdr.rejected {
if idx < singles.len() {
singles[idx].evidence.fdr_selected = true;
} else {
pairs[idx - singles.len()].evidence.fdr_selected = true;
}
}
let fdr_selected_chart_blocks = singles
.iter()
.filter(|c| c.evidence.fdr_selected)
.map(|c| c.blocks[0])
.collect::<Vec<_>>();
let fdr_selected_chart_pairs = pairs
.iter()
.filter(|c| c.evidence.fdr_selected)
.map(|c| (c.blocks[0], c.blocks[1]))
.collect::<Vec<_>>();
let block_records = singles
.iter()
.map(|c| BlockChartRecord {
block0: c.blocks[0],
block1: None,
screen_score: c.screen_score,
evidence: c.evidence.clone(),
})
.collect::<Vec<_>>();
let pair_records = pairs
.iter()
.map(|c| BlockChartRecord {
block0: c.blocks[0],
block1: Some(c.blocks[1]),
screen_score: c.screen_score,
evidence: c.evidence.clone(),
})
.collect::<Vec<_>>();
let selected_chart_blocks = singles
.iter()
.filter(|c| c.evidence.selected_by_bic)
.map(|c| c.blocks[0])
.collect::<Vec<_>>();
let selected_chart_pairs = pairs
.iter()
.filter(|c| c.evidence.selected_by_bic)
.map(|c| (c.blocks[0], c.blocks[1]))
.collect::<Vec<_>>();
Ok(BlockChartComposeResult {
reconstructed,
block_records,
pair_records,
selected_blocks: selected,
selected_chart_blocks,
selected_chart_pairs,
fdr_alpha: CHART_FDR_ALPHA,
fdr_selected_chart_blocks,
fdr_selected_chart_pairs,
})
}
pub fn block_sparse_dictionary_firings(
blocks: ArrayView2<'_, u32>,
n_blocks: usize,
) -> Result<Vec<usize>, String> {
let mut counts = vec![0usize; n_blocks];
for i in 0..blocks.nrows() {
for j in 0..blocks.ncols() {
let g = blocks[[i, j]] as usize;
if g >= n_blocks {
return Err(format!(
"block firings: block index {g} out of range 0..{n_blocks}"
));
}
counts[g] += 1;
}
}
Ok(counts)
}
pub fn block_sparse_dictionary_seed_manifest(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
block_utilization: &[f32],
block_stable_rank: &[f32],
explained_variance: f64,
config: &BlockSeedManifestConfig,
) -> Result<BlockSeedManifest, String> {
if config.block_size == 0 {
return Err("block seed manifest: block_size must be >= 1".to_string());
}
if decoder.nrows() == 0 || decoder.nrows() % config.block_size != 0 {
return Err(
"block seed manifest: decoder rows must be a positive multiple of block_size"
.to_string(),
);
}
if x.ncols() != decoder.ncols() {
return Err(format!(
"block seed manifest: X has P={} but decoder has P={}",
x.ncols(),
decoder.ncols()
));
}
if blocks.nrows() != x.nrows() {
return Err(format!(
"block seed manifest: blocks has {} rows but X has {}",
blocks.nrows(),
x.nrows()
));
}
let n_blocks = decoder.nrows() / config.block_size;
if block_utilization.len() != n_blocks || block_stable_rank.len() != n_blocks {
return Err(format!(
"block seed manifest: block reports must have length {n_blocks}, got {} and {}",
block_utilization.len(),
block_stable_rank.len()
));
}
let firings = block_sparse_dictionary_firings(blocks, n_blocks)?;
let ambient_var = centered_energy_view(x).max(f64::MIN_POSITIVE);
let mut records = Vec::with_capacity(n_blocks);
for g in 0..n_blocks {
let coords = block_coords_for_seed_config(x, decoder, config, g)?;
let coded_var = coordinate_spectrum(&coords)?;
let total_var = coded_var.iter().sum::<f64>();
let total_var_report = total_var.max(f64::MIN_POSITIVE);
let block_ev = centered_energy(&coords) / ambient_var;
let base = format!("{}{}", config.name_prefix, g);
let n_firings = firings[g].max(1);
let mdl_block = MdlFeaturizerRow {
name: format!("{base}-linear-{}d", config.block_size),
kind: "block".to_string(),
total_var: total_var_report,
n_tokens: x.nrows(),
n_firings,
n_params: config.block_size * decoder.ncols(),
coded_var: coded_var.clone(),
g_dict: n_blocks,
k_active: config.block_topk,
block_name: None,
chart_name: None,
};
let chart_name = format!("{base}-circle-chart");
let block_name = mdl_block.name.clone();
let mdl_chart = MdlFeaturizerRow {
name: chart_name.clone(),
kind: "chart".to_string(),
total_var: total_var_report,
n_tokens: x.nrows(),
n_firings,
n_params: config.n_basis_chart * decoder.ncols(),
coded_var: vec![total_var_report],
g_dict: n_blocks,
k_active: config.block_topk,
block_name: Some(block_name),
chart_name: Some(chart_name),
};
let firing_rows = rows_for_block(blocks, g);
let firing_coords = take_rows(&coords, &firing_rows);
let (matched_dl_flat, matched_dl_chart) =
matched_dl_for_block(&firing_coords, config, decoder.ncols(), block_ev);
let matched_dl_delta_bits =
crate::description_length::matched_dl_delta(&matched_dl_flat, &matched_dl_chart);
records.push(BlockSeedRecord {
block: g,
block_dim: config.block_size,
n_firings: firings[g],
utilization: block_utilization[g],
stable_rank: block_stable_rank[g],
coded_var,
total_var,
block_linear_ev: block_ev,
basis: config
.include_bases
.then(|| block_basis(decoder, config.block_size, g)),
mdl_block,
mdl_chart,
matched_dl_flat,
matched_dl_chart,
matched_dl_delta_bits,
});
}
Ok(BlockSeedManifest {
n_blocks,
block_size: config.block_size,
block_topk: config.block_topk,
ambient_p: decoder.ncols(),
gamma: config.gamma,
explained_variance,
residual_target: config.residual_target,
n_basis_chart: config.n_basis_chart,
blocks: records,
})
}
fn matched_dl_for_block(
firing_coords: &Array2<f32>,
config: &BlockSeedManifestConfig,
ambient_p: usize,
ev: f64,
) -> (
crate::description_length::MatchedDl,
crate::description_length::MatchedDl,
) {
use crate::description_length::{matched_dl, se_resolution_bits};
let n_fire = firing_coords.nrows();
let mut norms: Vec<f64> = Vec::with_capacity(n_fire);
for row in firing_coords.rows() {
norms.push(
row.iter()
.map(|&v| (v as f64) * (v as f64))
.sum::<f64>()
.sqrt(),
);
}
let sigma_hat = if n_fire >= 2 {
let mean = norms.iter().sum::<f64>() / n_fire as f64;
let ss: f64 = norms.iter().map(|&r| (r - mean) * (r - mean)).sum();
(ss / (n_fire - 1) as f64).sqrt()
} else {
0.0
};
let phase_ses: Vec<f64> = norms
.iter()
.map(|&nz| super::coordinate::phase_coordinate_se(sigma_hat, nz))
.collect();
let amp_ses: Vec<f64> = norms
.iter()
.map(|&nz| {
if nz > 0.0 {
sigma_hat / nz
} else {
f64::INFINITY
}
})
.collect();
let mean_rate = |ses: &[f64]| -> f64 {
if ses.is_empty() {
0.0
} else {
ses.iter().map(|&se| se_resolution_bits(se)).sum::<f64>() / ses.len() as f64
}
};
let l_param_flat = mean_rate(&_ses);
let l_param_chart = mean_rate(&phase_ses);
let flat = matched_dl(
config.block_size as i64,
config.block_size as i64,
ambient_p as i64,
l_param_flat,
&_ses,
ev,
);
let chart = matched_dl(
config.n_basis_chart as i64,
1,
ambient_p as i64,
l_param_chart,
&phase_ses,
ev,
);
(flat, chart)
}
fn validate_inputs(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
config: &BlockChartComposeConfig,
) -> Result<(), String> {
if config.block_size == 0 {
return Err("block chart compose: block_size must be >= 1".to_string());
}
if decoder.nrows() == 0 || decoder.nrows() % config.block_size != 0 {
return Err(
"block chart compose: decoder rows must be a positive multiple of block_size"
.to_string(),
);
}
if x.ncols() != decoder.ncols() {
return Err(format!(
"block chart compose: X has P={} but decoder has P={}",
x.ncols(),
decoder.ncols()
));
}
let (n, k) = blocks.dim();
if n != x.nrows() || codes.shape() != [n, k, config.block_size] {
return Err(format!(
"block chart compose: blocks/codes shapes must be ({}, k) and ({}, k, {}), got {:?} and {:?}",
x.nrows(),
x.nrows(),
config.block_size,
blocks.dim(),
codes.shape()
));
}
if config.block_tile == 0 {
return Err("block chart compose: block_tile must be >= 1".to_string());
}
Ok(())
}
fn block_coords_for_config(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
config: &BlockChartComposeConfig,
block: usize,
) -> Result<Array2<f32>, String> {
if config.residual_target {
block_sparse_dictionary_project_residual_with_codes(
x,
decoder,
blocks,
codes,
config.block_size,
block,
)
} else {
block_sparse_dictionary_block_coords(x, decoder, config.block_size, block)
}
}
fn block_coords_for_seed_config(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
config: &BlockSeedManifestConfig,
block: usize,
) -> Result<Array2<f32>, String> {
if config.residual_target {
block_sparse_dictionary_project_residual(
x,
decoder,
config.gamma,
config.block_size,
config.block_topk,
config.block_tile,
block,
)
} else {
block_sparse_dictionary_block_coords(x, decoder, config.block_size, block)
}
}
fn select_blocks(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
config: &BlockChartComposeConfig,
) -> Result<Vec<usize>, String> {
let g_total = decoder.nrows() / config.block_size;
let mut scored = Vec::<(f64, usize, usize)>::new();
for g in 0..g_total {
let rows = rows_for_block(blocks, g);
if rows.len() < config.min_firings {
continue;
}
let coords_all = block_coords_for_config(x, decoder, blocks, codes, config, g)?;
let coords = take_rows(&coords_all, &rows);
let energy = centered_energy(&coords);
scored.push((energy, rows.len(), g));
}
scored.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(config.max_blocks.min(scored.len()));
Ok(scored.into_iter().map(|(_, _, g)| g).collect())
}
fn screen_pairs(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
config: &BlockChartComposeConfig,
selected: &[usize],
) -> Result<Vec<(usize, usize, f64)>, String> {
let top = selected.len().min(config.pair_top_blocks);
let mut out = Vec::new();
for i in 0..top {
let g0 = selected[i];
let z0_all = block_coords_for_config(x, decoder, blocks, codes, config, g0)?;
for &g1 in selected.iter().take(top).skip(i + 1) {
let rows = rows_for_pair(blocks, g0, g1);
if rows.len() < config.pair_min_cofirings {
continue;
}
let z0 = take_rows(&z0_all, &rows);
let z1_all = block_coords_for_config(x, decoder, blocks, codes, config, g1)?;
let z1 = take_rows(&z1_all, &rows);
out.push((g0, g1, pair_score(&z0, &z1)?));
}
}
out.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal));
out.truncate(config.max_pairs.min(out.len()));
Ok(out)
}
fn crossfit_evidence(
coords: &Array2<f32>,
config: &BlockChartComposeConfig,
) -> Result<ChartEvidence, String> {
let n = coords.nrows();
if n < 2 {
return Err("block chart evidence requires at least two rows".to_string());
}
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);
}
}
let train_coords = take_rows(coords, &train);
let eval_coords = take_rows(coords, &eval);
let whitening = fit_whitening(&train_coords, config.whitening_ridge)?;
let train_w = whitening.transform(&to_f64(&train_coords));
let eval_w = whitening.transform(&to_f64(&eval_coords));
let linear_pred_w = pca_reconstruct(&train_w, &eval_w, 1)?;
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);
let eval_f = to_f64(&eval_coords);
for (pos, &row) in eval.iter().enumerate() {
linear_loss[row] = row_sse(&eval_f, &linear_pred, pos);
chart_loss[row] = row_sse(&eval_f, &chart_pred, pos);
}
}
let linear_total = linear_loss.iter().sum::<f64>();
let chart_total = chart_loss.iter().sum::<f64>();
let s2_floor = 1.0e-12 * (linear_total.max(chart_total) / n as f64).max(1.0e-300);
let s2_lin = (linear_total / n as f64).max(s2_floor);
let s2_chart = (chart_total / n as f64).max(s2_floor);
let half_log_ratio = 0.5 * (s2_lin / s2_chart).ln();
let q = coords.ncols().max(1) as f64;
let deviance: Vec<f64> = (0..n)
.map(|i| {
q * (half_log_ratio + linear_loss[i] / (2.0 * s2_lin)
- chart_loss[i] / (2.0 * s2_chart))
})
.collect();
let n_eff = autocorr_ess(&deviance);
let mean_delta = deviance.iter().sum::<f64>() / n as f64;
let se = newey_west_se(&deviance);
let ci_low = mean_delta - 1.959963984540054 * se;
let ci_high = mean_delta + 1.959963984540054 * se;
let d_eff = (2 * coords.ncols()).max(1) as f64;
let gain = n_eff.max(2.0) * mean_delta;
let charge = 0.5 * d_eff * n_eff.max(2.0).ln();
let margin = gain - charge;
let coords64 = to_f64(coords);
let log_e =
super::split_lr_fdr::shell_vs_ring_log_evalue(&coords64, folds, config.whitening_ridge)?;
Ok(ChartEvidence {
n_rows: n,
n_effective: n_eff,
linear_loss: linear_total,
chart_loss: chart_total,
deviance_gain: gain,
mean_delta,
se,
ci_low,
ci_high,
charge,
margin,
selected_by_bic: margin >= config.min_effect && ci_low > 0.0,
log_e,
fdr_selected: false,
})
}
fn fit_radial_chart_all(coords: &Array2<f32>, ridge: f64) -> Result<Array2<f64>, String> {
let whitening = fit_whitening(coords, ridge)?;
let z = whitening.transform(&to_f64(coords));
let pred = radial_predict(&z, &z);
Ok(whitening.inverse(&pred))
}
fn fit_whitening(coords: &Array2<f32>, ridge: f64) -> Result<Whitening, String> {
let x = to_f64(coords);
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 = vec![0.0; 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 * d + b] += va * (x[[i, b]] - mean[b]);
}
}
}
let denom = (n.saturating_sub(1)).max(1) as f64;
for v in &mut cov {
*v /= denom;
}
let (vals, eigvec) = jacobi_eigh(cov, d)?;
let max_eval = vals.iter().copied().fold(0.0, f64::max);
if !(ridge.is_finite() && ridge >= 0.0) {
return Err(format!(
"whitening ridge fraction must be finite and non-negative; got {ridge}"
));
}
let spectral_scale = max_eval.max(f64::MIN_POSITIVE);
let floor = (ridge * spectral_scale).max(f64::EPSILON * spectral_scale);
let scale = vals
.into_iter()
.map(|v| (v.max(0.0) + floor).sqrt())
.collect();
Ok(Whitening {
mean,
eigvec,
scale,
dim: d,
})
}
impl Whitening {
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 * self.dim + 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 * self.dim + k];
}
out[[i, j]] = v;
}
}
out
}
}
fn pca_reconstruct(
train: &Array2<f64>,
eval: &Array2<f64>,
rank: usize,
) -> 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 = vec![0.0; 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 * d + b] += va * (train[[i, b]] - mean[b]);
}
}
}
let denom = n.saturating_sub(1).max(1) as f64;
for v in &mut cov {
*v /= denom;
}
let (vals, eigvec) = jacobi_eigh(cov, d)?;
let mut order = (0..d).collect::<Vec<_>>();
order.sort_by(|&a, &b| {
vals[b]
.partial_cmp(&vals[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let r = rank.min(d);
let mut out = Array2::<f64>::zeros((eval.nrows(), d));
for i in 0..eval.nrows() {
for j in 0..d {
out[[i, j]] = mean[j];
}
for &k in order.iter().take(r) {
let mut score = 0.0;
for j in 0..d {
score += (eval[[i, j]] - mean[j]) * eigvec[j * d + k];
}
for j in 0..d {
out[[i, j]] += score * eigvec[j * d + k];
}
}
}
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
}
pub(crate) fn jacobi_eigh(mut a: Vec<f64>, n: usize) -> Result<(Vec<f64>, Vec<f64>), String> {
if a.len() != n * n {
return Err("jacobi_eigh: matrix length mismatch".to_string());
}
let mut v = vec![0.0; n * n];
for i in 0..n {
v[i * n + i] = 1.0;
}
for _ in 0..(64 * n.max(1) * n.max(1)) {
let mut p = 0usize;
let mut q = 0usize;
let mut max_off = 0.0;
for i in 0..n {
for j in (i + 1)..n {
let val = a[i * n + j].abs();
if val > max_off {
max_off = val;
p = i;
q = j;
}
}
}
if max_off <= 1.0e-12 {
break;
}
let app = a[p * n + p];
let aqq = a[q * n + q];
let apq = a[p * n + q];
let tau = (aqq - app) / (2.0 * apq);
let t = tau.signum() / (tau.abs() + (1.0 + tau * tau).sqrt());
let c = 1.0 / (1.0 + t * t).sqrt();
let s = t * c;
for k in 0..n {
let akp = a[k * n + p];
let akq = a[k * n + q];
a[k * n + p] = c * akp - s * akq;
a[k * n + q] = s * akp + c * akq;
}
for k in 0..n {
let apk = a[p * n + k];
let aqk = a[q * n + k];
a[p * n + k] = c * apk - s * aqk;
a[q * n + k] = s * apk + c * aqk;
}
for k in 0..n {
let vkp = v[k * n + p];
let vkq = v[k * n + q];
v[k * n + p] = c * vkp - s * vkq;
v[k * n + q] = s * vkp + c * vkq;
}
}
let vals = (0..n).map(|i| a[i * n + i]).collect::<Vec<_>>();
Ok((vals, v))
}
fn pair_score(z0: &Array2<f32>, z1: &Array2<f32>) -> Result<f64, String> {
let joint = hstack(z0, z1);
let whitening = fit_whitening(&joint, 1.0e-8)?;
let z = whitening.transform(&to_f64(&joint));
let mut radii = Vec::with_capacity(z.nrows());
for i in 0..z.nrows() {
let mut ss = 0.0;
for j in 0..z.ncols() {
ss += z[[i, j]] * z[[i, j]];
}
radii.push(ss.sqrt());
}
let mean = radii.iter().sum::<f64>() / radii.len().max(1) as f64;
if mean <= 0.0 {
return Ok(0.0);
}
let var =
radii.iter().map(|r| (r - mean) * (r - mean)).sum::<f64>() / radii.len().max(1) as f64;
Ok((1.0 - var.sqrt() / mean).clamp(0.0, 1.0))
}
fn rows_for_block(blocks: ArrayView2<'_, u32>, block: usize) -> Vec<usize> {
let mut out = Vec::new();
for i in 0..blocks.nrows() {
if (0..blocks.ncols()).any(|j| blocks[[i, j]] as usize == block) {
out.push(i);
}
}
out
}
fn rows_for_pair(blocks: ArrayView2<'_, u32>, g0: usize, g1: usize) -> Vec<usize> {
let mut out = Vec::new();
for i in 0..blocks.nrows() {
let has0 = (0..blocks.ncols()).any(|j| blocks[[i, j]] as usize == g0);
let has1 = (0..blocks.ncols()).any(|j| blocks[[i, j]] as usize == g1);
if has0 && has1 {
out.push(i);
}
}
out
}
fn take_rows(a: &Array2<f32>, rows: &[usize]) -> Array2<f32> {
let mut out = Array2::<f32>::zeros((rows.len(), a.ncols()));
for (i, &row) in rows.iter().enumerate() {
for j in 0..a.ncols() {
out[[i, j]] = a[[row, j]];
}
}
out
}
fn hstack(a: &Array2<f32>, b: &Array2<f32>) -> Array2<f32> {
let mut out = Array2::<f32>::zeros((a.nrows(), a.ncols() + b.ncols()));
for i in 0..a.nrows() {
for j in 0..a.ncols() {
out[[i, j]] = a[[i, j]];
}
for j in 0..b.ncols() {
out[[i, a.ncols() + j]] = b[[i, j]];
}
}
out
}
fn centered_energy(a: &Array2<f32>) -> f64 {
let mut means = vec![0.0; a.ncols()];
for j in 0..a.ncols() {
for i in 0..a.nrows() {
means[j] += a[[i, j]] as f64;
}
means[j] /= a.nrows().max(1) as f64;
}
let mut e = 0.0;
for i in 0..a.nrows() {
for j in 0..a.ncols() {
let v = a[[i, j]] as f64 - means[j];
e += v * v;
}
}
e
}
fn centered_energy_view(a: ArrayView2<'_, f32>) -> f64 {
let mut means = vec![0.0; a.ncols()];
for j in 0..a.ncols() {
for i in 0..a.nrows() {
means[j] += a[[i, j]] as f64;
}
means[j] /= a.nrows().max(1) as f64;
}
let mut e = 0.0;
for i in 0..a.nrows() {
for j in 0..a.ncols() {
let v = a[[i, j]] as f64 - means[j];
e += v * v;
}
}
e
}
fn coordinate_spectrum(coords: &Array2<f32>) -> Result<Vec<f64>, String> {
let n = coords.nrows();
let d = coords.ncols();
let mut means = vec![0.0; d];
for j in 0..d {
for i in 0..n {
means[j] += coords[[i, j]] as f64;
}
means[j] /= n.max(1) as f64;
}
let mut cov = vec![0.0; d * d];
for i in 0..n {
for a in 0..d {
let va = coords[[i, a]] as f64 - means[a];
for b in 0..d {
cov[a * d + b] += va * (coords[[i, b]] as f64 - means[b]);
}
}
}
let denom = n.max(1) as f64;
for v in &mut cov {
*v /= denom;
}
let (vals, _) = jacobi_eigh(cov, d)?;
let mut spectrum = vals.into_iter().map(|v| v.max(0.0)).collect::<Vec<_>>();
spectrum.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
Ok(spectrum)
}
fn block_basis(decoder: ArrayView2<'_, f32>, block_size: usize, block: usize) -> Vec<Vec<f32>> {
let mut basis = vec![vec![0.0; block_size]; decoder.ncols()];
for p in 0..decoder.ncols() {
for r in 0..block_size {
basis[p][r] = decoder[[block * block_size + r, p]];
}
}
basis
}
fn to_f64(a: &Array2<f32>) -> Array2<f64> {
let mut out = Array2::<f64>::zeros(a.dim());
for i in 0..a.nrows() {
for j in 0..a.ncols() {
out[[i, j]] = a[[i, j]] as f64;
}
}
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()
}
fn subtract_block_contribution(
out: &mut Array2<f32>,
decoder: ArrayView2<'_, f32>,
blocks: ArrayView2<'_, u32>,
codes: ArrayView3<'_, f32>,
b: usize,
row: usize,
block: usize,
) {
for j in 0..blocks.ncols() {
if blocks[[row, j]] as usize != block {
continue;
}
for r in 0..b {
let code = codes[[row, j, r]];
let atom = decoder.row(block * b + r);
for c in 0..out.ncols() {
out[[row, c]] -= code * atom[c];
}
}
}
}
fn add_lifted_coords(
out: &mut Array2<f32>,
decoder: ArrayView2<'_, f32>,
b: usize,
row: usize,
block: usize,
coords: &[f64],
offset: usize,
) {
for r in 0..b {
let code = coords[offset + r] as f32;
let atom = decoder.row(block * b + r);
for c in 0..out.ncols() {
out[[row, c]] += code * atom[c];
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn unif(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(*state >> 33) as f64 / (1u64 << 31) as f64
}
fn annulus(n: usize) -> Array2<f32> {
let mut st = 0x1234_5678u64;
let mut z = Array2::<f32>::zeros((n, 2));
for i in 0..n {
let theta = std::f64::consts::TAU * (i as f64 / n as f64) + 0.01 * unif(&mut st);
let r = 1.0 + 0.05 * (unif(&mut st) - 0.5);
z[[i, 0]] = (r * theta.cos()) as f32;
z[[i, 1]] = (r * theta.sin()) as f32;
}
z
}
#[test]
fn crossfit_bic_selection_is_scale_invariant() {
let coords = annulus(400);
let config = BlockChartComposeConfig::default();
assert!(
config.whitening_ridge > 0.0,
"the default ridge must be exercised (a nonzero absolute ridge would break \
scale-equivariance; the relative-fraction ridge does not)"
);
let base = crossfit_evidence(&coords, &config).expect("evidence");
assert!(base.deviance_gain > 0.0, "chart must beat linear on a ring");
assert!(
base.selected_by_bic,
"the ring must be selected by held-out BIC"
);
let mut scaled = coords.clone();
scaled.mapv_inplace(|v| v * 10.0);
let big = crossfit_evidence(&scaled, &config).expect("evidence scaled");
assert_eq!(
base.selected_by_bic, big.selected_by_bic,
"accept/reject verdict must be scale-invariant"
);
let scale_invariant =
|a: f64, b: f64| (a - b).abs() <= 1e-6 * a.abs().max(b.abs()).max(1.0);
assert!(
scale_invariant(base.margin, big.margin),
"margin must be scale-invariant: {} vs {}",
base.margin,
big.margin
);
assert!(
scale_invariant(base.deviance_gain, big.deviance_gain),
"deviance gain must be scale-invariant: {} vs {}",
base.deviance_gain,
big.deviance_gain
);
assert!(
scale_invariant(base.ci_low, big.ci_low) && scale_invariant(base.ci_high, big.ci_high),
"the deviance-scale CI must be scale-invariant"
);
let old_gain_base = base.linear_loss - base.chart_loss;
let old_gain_big = big.linear_loss - big.chart_loss;
assert!(
(old_gain_big - 100.0 * old_gain_base).abs() < 1e-3 * old_gain_big.abs().max(1.0),
"raw SSE gain is scale-dependent (×100) — the currency the fix replaces"
);
}
}