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 std::sync::Arc;
use gam_linalg::faer_ndarray::FaerEigh;
use ndarray::{Array1, Array2, ArrayView2, Axis};
use crate::basis::{AnchorIndicatorEvaluator, SaeBasisEvaluator};
use crate::manifold::{SaeAtomBasisKind, SaeManifoldAtom, finite_set_rank_charge};
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(),
}
}
fn is_sink(self) -> bool {
!matches!(self, Self::Semantic)
}
}
#[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 position_zero() -> Self {
Self {
enabled: true,
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())
}
pub fn residual_after_sink(
&self,
residual: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
let expected = (self.atom.basis_values.nrows(), self.atom.output_dim());
if residual.dim() != expected {
return Err(format!(
"Tier05SinkAtom::residual_after_sink: residual shape {:?} incompatible with atom rows/output ({}, {})",
residual.dim(),
self.atom.basis_values.nrows(),
self.atom.output_dim()
));
}
Ok(&residual - &self.reconstruction())
}
pub fn reconstruct_with_sink(
&self,
semantic_recon: ArrayView2<'_, f64>,
) -> Result<Array2<f64>, String> {
let expected = (self.atom.basis_values.nrows(), self.atom.output_dim());
if semantic_recon.dim() != expected {
return Err(format!(
"Tier05SinkAtom::reconstruct_with_sink: reconstruction shape {:?} incompatible with atom rows/output ({}, {})",
semantic_recon.dim(),
self.atom.basis_values.nrows(),
self.atom.output_dim()
));
}
Ok(&semantic_recon + &self.reconstruction())
}
}
pub fn fit_tier05_sink_atom(
residual: ArrayView2<'_, f64>,
positions: &[i64],
delimiter_classes: &[Option<SinkDelimiterClass>],
config: &Tier05SinkAtomConfig,
) -> Result<Option<Tier05SinkAtom>, String> {
if !config.enabled {
return Ok(None);
}
let n = residual.nrows();
let p = residual.ncols();
if n == 0 || p == 0 {
return Err("fit_tier05_sink_atom: residual must be a non-empty N×P matrix".to_string());
}
if positions.len() != n {
return Err(format!(
"fit_tier05_sink_atom: positions length {} != N {n}",
positions.len()
));
}
if !delimiter_classes.is_empty() && delimiter_classes.len() != n {
return Err(format!(
"fit_tier05_sink_atom: delimiter_classes length {} must be 0 or N {n}",
delimiter_classes.len()
));
}
if !config.delimiter_classes.is_empty() && delimiter_classes.is_empty() {
return Err(
"fit_tier05_sink_atom: delimiter classes configured but no per-row delimiter labels supplied"
.to_string(),
);
}
let anchors = config.anchors()?;
let delimiter_set: BTreeSet<SinkDelimiterClass> =
config.delimiter_classes.iter().copied().collect();
let mut anchor_lookup = BTreeMap::new();
for (idx, anchor) in anchors.iter().copied().enumerate() {
anchor_lookup.insert(anchor, idx);
}
let mut coords = Array2::<f64>::zeros((n, 1));
let mut counts = vec![0usize; anchors.len()];
for row in 0..n {
let delimiter = if delimiter_classes.is_empty() {
None
} else {
delimiter_classes[row]
};
let anchor = if config.include_position_zero && positions[row] == 0 {
SinkAnchor::PositionZero
} else if let Some(class) = delimiter {
if delimiter_set.contains(&class) {
SinkAnchor::Delimiter(class)
} else {
SinkAnchor::Semantic
}
} else {
SinkAnchor::Semantic
};
let idx = anchor_lookup.get(&anchor).copied().ok_or_else(|| {
format!(
"fit_tier05_sink_atom: support anchor {} was not configured",
anchor.label()
)
})?;
coords[[row, 0]] = idx as f64;
counts[idx] += 1;
}
let sink_rows: usize = anchors
.iter()
.zip(counts.iter())
.filter(|(anchor, _count)| anchor.is_sink())
.map(|(_anchor, &count)| count)
.sum();
if sink_rows == 0 {
return Err(
"fit_tier05_sink_atom: enabled sink atom has no sink-supported rows".to_string(),
);
}
let evaluator = Arc::new(AnchorIndicatorEvaluator::new(anchors.len())?);
let (basis_values, basis_jacobian) = evaluator.evaluate(coords.view())?;
let mut decoder = Array2::<f64>::zeros((anchors.len(), p));
for row in 0..n {
let anchor = coords[[row, 0]] as usize;
for col in 0..p {
decoder[[anchor, col]] += residual[[row, col]];
}
}
for anchor in 0..anchors.len() {
if counts[anchor] > 0 {
let scale = 1.0 / counts[anchor] as f64;
for col in 0..p {
decoder[[anchor, col]] *= scale;
}
}
}
let smooth_penalty = Array2::<f64>::zeros((anchors.len(), anchors.len()));
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"tier0_5_attention_sink",
SaeAtomBasisKind::FiniteSet,
1,
basis_values,
basis_jacobian,
decoder,
smooth_penalty,
)?
.with_basis_second_jet(evaluator);
let reconstruction = atom.basis_values.dot(atom.decoder_coefficients());
let mut rss = 0.0f64;
let mut tss = 0.0f64;
for row in 0..n {
for col in 0..p {
let r = residual[[row, col]] - reconstruction[[row, col]];
rss += r * r;
let v = residual[[row, col]];
tss += v * v;
}
}
let variance_absorbed = if tss <= 0.0 { 0.0 } else { 1.0 - rss / tss };
Ok(Some(Tier05SinkAtom {
atom,
anchors,
anchor_counts: counts,
rank_charge: finite_set_rank_charge(anchor_lookup.len()),
variance_absorbed,
}))
}
pub fn fit_position0_sink_atom(
residual: ArrayView2<'_, f64>,
positions: &[i64],
) -> Result<Tier05SinkAtom, String> {
let config = Tier05SinkAtomConfig::position_zero();
fit_tier05_sink_atom(residual, positions, &[], &config)?.ok_or_else(|| {
"fit_position0_sink_atom: position-0 sink config unexpectedly disabled".to_string()
})
}
#[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>,
}
pub fn interference_subspace(
fit: &SparseDictFit,
rank: Option<usize>,
) -> Result<InterferenceSubspace, String> {
let decoder = fit.decoder.view();
let k = decoder.nrows();
let p = decoder.ncols();
if k == 0 || p == 0 {
return Err("interference_subspace: empty decoder".to_string());
}
let mut weight = vec![0.0f64; k];
for (idx_row, code_row) in fit.indices.rows().into_iter().zip(fit.codes.rows()) {
for (&atom_u32, &code) in idx_row.iter().zip(code_row.iter()) {
let atom = atom_u32 as usize;
if atom < k {
weight[atom] += (code as f64) * (code as f64);
}
}
}
let mut dw = Array2::<f64>::zeros((k, p));
for atom in 0..k {
let sw = weight[atom].max(0.0).sqrt();
if sw == 0.0 {
continue;
}
let src = decoder.row(atom);
let mut dst = dw.row_mut(atom);
for c in 0..p {
dst[c] = sw * (src[c] as f64);
}
}
let gram = dw.t().dot(&dw);
let (evals, evecs) = gram
.eigh(faer::Side::Lower)
.map_err(|err| format!("interference_subspace eigensolve failed: {err}"))?;
let total: f64 = evals.iter().map(|&e| e.max(0.0)).sum();
if total <= 0.0 {
return Err(
"interference_subspace: Tier-1 decoder carries no fired energy (all atoms dead)"
.to_string(),
);
}
let r = match rank {
Some(r) => r.min(p).max(1),
None => {
let mut acc = 0.0f64;
let mut chosen = 1usize;
for (taken, &e) in evals.iter().rev().enumerate() {
acc += e.max(0.0);
chosen = taken + 1;
if acc >= 0.99 * total {
break;
}
}
chosen.min(p).max(1)
}
};
let mut q = Array2::<f64>::zeros((p, r));
let mut scale = Array1::<f64>::zeros(r);
for j in 0..r {
let col = p - 1 - j; q.column_mut(j).assign(&evecs.column(col));
scale[j] = evals[col].max(0.0).sqrt();
}
let pr = p - r;
let mut q_perp = Array2::<f64>::zeros((p, pr));
for j in 0..pr {
q_perp.column_mut(j).assign(&evecs.column(j));
}
Ok(InterferenceSubspace { q, q_perp, scale })
}
#[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 }
}
pub fn explained_variance_vs_mean(
z: ArrayView2<'_, f64>,
recon: ArrayView2<'_, f64>,
mean: &Array1<f64>,
) -> f64 {
let mut rss = 0.0f64;
for (zr, rr) in z.rows().into_iter().zip(recon.rows()) {
for c in 0..z.ncols() {
let d = zr[c] - rr[c];
rss += d * d;
}
}
let baseline = &z - &mean.view().insert_axis(Axis(0));
let tss: f64 = baseline.iter().map(|&v| v * v).sum();
explained_variance_from_sums(rss, tss)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sparse_dict::DecoderSolveStats;
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 interference_subspace_q_and_qperp_are_orthonormal_complements() {
let decoder = array![[1.0f32, 0.0, 0.0], [0.0, 1.0, 0.0]];
let indices = array![[0u32, 1u32], [0u32, 1u32]];
let codes = array![[2.0f32, 1.0f32], [2.0f32, 1.0f32]];
let fit = SparseDictFit {
decoder,
indices,
codes,
explained_variance: 0.0,
epochs: 0,
convergence: crate::sparse_dict::SparseDictConvergence::trivially_converged(),
active: 2,
score_route_stats: Default::default(),
decoder_solve_stats: DecoderSolveStats::default(),
};
let sub = interference_subspace(&fit, Some(2)).expect("subspace");
assert_eq!(sub.q.dim(), (3, 2));
assert_eq!(sub.q_perp.dim(), (3, 1));
let gq = sub.q.t().dot(&sub.q);
assert!((gq[[0, 0]] - 1.0).abs() < 1e-9 && (gq[[1, 1]] - 1.0).abs() < 1e-9);
assert!(gq[[0, 1]].abs() < 1e-9);
let cross = sub.q.t().dot(&sub.q_perp);
assert!(cross.iter().all(|&v| v.abs() < 1e-9));
assert!(sub.q_perp[[2, 0]].abs() > 0.999);
assert!(sub.q_perp[[0, 0]].abs() < 1e-6 && sub.q_perp[[1, 0]].abs() < 1e-6);
assert!(sub.scale[0] >= sub.scale[1]);
}
#[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 qperp_weight_is_blind_to_in_plane_curvature() {
let decoder = array![[1.0f32, 0.0, 0.0, 0.0], [0.0, 1.0, 0.0, 0.0]];
let indices = array![[0u32, 1u32], [0u32, 1u32], [0u32, 1u32]];
let codes = array![[3.0f32, 2.0], [3.0, 2.0], [3.0, 2.0]];
let fit = SparseDictFit {
decoder,
indices,
codes,
explained_variance: 0.0,
epochs: 0,
convergence: crate::sparse_dict::SparseDictConvergence::trivially_converged(),
active: 2,
score_route_stats: Default::default(),
decoder_solve_stats: DecoderSolveStats::default(),
};
let sub = interference_subspace(&fit, Some(2)).expect("subspace");
for j in 0..2 {
let mut ej = Array1::<f64>::zeros(4);
ej[j] = 1.0;
let qte = sub.q.t().dot(&ej);
let proj_norm = qte.dot(&qte).sqrt();
assert!(
(proj_norm - 1.0).abs() < 1e-9,
"e{j} not fully in Q: {proj_norm}"
);
}
let curvature = array![0.6f64, -0.8, 0.0, 0.0];
let sig_rms = curvature.dot(&curvature).sqrt();
assert!(
sig_rms > 0.99,
"planted signal should be ~unit; got {sig_rms}"
);
let qperp_component = sub.q_perp.t().dot(&curvature);
let qperp_rms = qperp_component.dot(&qperp_component).sqrt();
assert!(
qperp_rms < 1e-9,
"Q⊥ weight crushes the in-plane curvature to noise ({qperp_rms}) — it is BLIND \
to what Tier-2 must model; fit the raw residual instead"
);
}
#[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);
}
#[test]
fn explained_variance_vs_mean_measures_about_the_mean_not_zero() {
let z = array![[10.0, 20.0], [12.0, 22.0], [14.0, 24.0]];
let mean = Array1::from(vec![12.0, 22.0]);
let recon = array![[11.0, 21.0], [12.0, 22.0], [15.0, 25.0]];
let ev = explained_variance_vs_mean(z.view(), recon.view(), &mean);
let expected = 1.0 - 4.0 / 16.0;
assert!(
(ev - expected).abs() < 1e-12,
"ev {ev} vs {expected}: baseline must be the mean, not zero"
);
assert!(ev < 0.99, "a zero baseline would inflate this to nearly 1");
}
#[test]
fn explained_variance_vs_mean_agrees_with_the_shared_policy() {
let z = array![[1.0, -2.0], [3.0, 4.0], [-5.0, 6.0]];
let recon = array![[0.5, -1.0], [2.0, 4.5], [-4.0, 5.0]];
let mean = Array1::from(vec![
z.column(0).sum() / 3.0,
z.column(1).sum() / 3.0,
]);
let mut rss = 0.0;
let mut tss = 0.0;
for r in 0..z.nrows() {
for c in 0..z.ncols() {
let d = z[[r, c]] - recon[[r, c]];
rss += d * d;
let m = z[[r, c]] - mean[c];
tss += m * m;
}
}
let direct = explained_variance_from_sums(rss, tss);
let wrapped = explained_variance_vs_mean(z.view(), recon.view(), &mean);
assert!((direct - wrapped).abs() < 1e-12, "{direct} vs {wrapped}");
}
}