use ndarray::{Array2, ArrayView2, ArrayView3};
use super::block::{
block_sparse_dictionary_block_coords, block_sparse_dictionary_project_residual,
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 alpha: f64,
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,
}
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,
alpha: 0.10,
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,
}
}
}
#[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 log_e_value: f64,
pub accepted_pre_ebh: bool,
pub accepted: 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 accepted_blocks: Vec<usize>,
pub accepted_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,
}
#[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,
}
#[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, 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, 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, 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, config, g0)?;
let coords1_all = block_coords_for_config(x, decoder, 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,
});
}
}
apply_ebh(&mut singles, &mut pairs, config.alpha, config.min_effect);
let mut reconstructed = base;
let mut replaced = vec![vec![false; decoder.nrows() / config.block_size]; x.nrows()];
for pair in &pairs {
if !pair.evidence.accepted {
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.accepted {
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 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 accepted_blocks = singles
.iter()
.filter(|c| c.evidence.accepted)
.map(|c| c.blocks[0])
.collect::<Vec<_>>();
let accepted_pairs = pairs
.iter()
.filter(|c| c.evidence.accepted)
.map(|c| (c.blocks[0], c.blocks[1]))
.collect::<Vec<_>>();
Ok(BlockChartComposeResult {
reconstructed,
block_records,
pair_records,
selected_blocks: selected,
accepted_blocks,
accepted_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),
};
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,
});
}
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 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.alpha > 0.0 && config.alpha <= 1.0) {
return Err("block chart compose: alpha must be in (0, 1]".to_string());
}
Ok(())
}
fn block_coords_for_config(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
config: &BlockChartComposeConfig,
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,
1024,
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,
1024,
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>,
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, 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>,
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, 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, 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 delta = linear_loss
.iter()
.zip(chart_loss.iter())
.map(|(l, c)| l - c)
.collect::<Vec<_>>();
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 ci_high = mean_delta + 1.959963984540054 * se;
let linear_total = linear_loss.iter().sum::<f64>();
let chart_total = chart_loss.iter().sum::<f64>();
let gain = linear_total - chart_total;
let d_eff = (2 * coords.ncols()).max(1) as f64;
let charge = 0.5 * d_eff * n_eff.max(2.0).ln();
let margin = gain - charge;
let scale = (linear_total / n as f64)
.max(chart_total / n as f64)
.max(1.0e-12);
let log_e_value = (margin / scale).max(0.0);
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,
log_e_value,
accepted_pre_ebh: margin >= config.min_effect && ci_low > 0.0,
accepted: false,
})
}
fn apply_ebh(
singles: &mut [CandidateFit],
pairs: &mut [CandidateFit],
alpha: f64,
min_effect: f64,
) {
let mut logs = Vec::new();
let mut refs = Vec::<(bool, usize)>::new();
for (i, c) in singles.iter().enumerate() {
if c.evidence.accepted_pre_ebh && c.evidence.margin >= min_effect {
logs.push(c.evidence.log_e_value);
refs.push((false, i));
}
}
for (i, c) in pairs.iter().enumerate() {
if c.evidence.accepted_pre_ebh && c.evidence.margin >= min_effect {
logs.push(c.evidence.log_e_value);
refs.push((true, i));
}
}
let keep = ebh(&logs, alpha);
for (ok, (is_pair, idx)) in keep.into_iter().zip(refs.into_iter()) {
if is_pair {
pairs[idx].evidence.accepted = ok;
} else {
singles[idx].evidence.accepted = ok;
}
}
}
fn ebh(logs: &[f64], alpha: f64) -> Vec<bool> {
let m = logs.len();
let mut order = (0..m).collect::<Vec<_>>();
order.sort_by(|&a, &b| {
logs[b]
.partial_cmp(&logs[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut k_star = 0usize;
for (rank0, &idx) in order.iter().enumerate() {
let rank = rank0 + 1;
let threshold = (m as f64 / (alpha * rank as f64)).ln();
if logs[idx] >= threshold {
k_star = rank;
}
}
let mut keep = vec![false; m];
for &idx in order.iter().take(k_star) {
keep[idx] = true;
}
keep
}
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).max(1.0);
let floor = ridge.max(f64::EPSILON * max_eval);
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
}
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];
}
}
}