use super::*;
use crate::chart_coordinate_solve::{ChartBasisKind, PeriodicCurveExtrema};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CrosscoderLayer {
Anchor,
Block(usize),
}
#[derive(Clone, Debug)]
pub struct AtomTransportReport {
pub atom: usize,
pub source: CrosscoderLayer,
pub target: CrosscoderLayer,
pub grid_resolution: usize,
pub n_harmonics: usize,
pub phase_shift: (f64, f64),
pub phase_r2: f64,
pub smooth_r2: f64,
pub drift: f64,
pub principal_angles: Vec<f64>,
pub transport_grid: Vec<(f64, f64)>,
}
impl AtomTransportReport {
pub fn law_gap(&self) -> f64 {
self.smooth_r2 - self.phase_r2
}
pub fn law_holds(&self, gap_tol: f64) -> bool {
self.phase_r2.is_finite() && self.smooth_r2.is_finite() && self.law_gap() <= gap_tol
}
pub fn deviation_locus(&self) -> Option<f64> {
let (s, phi) = self.phase_shift;
let two_pi = std::f64::consts::TAU;
self.transport_grid
.iter()
.map(|&(t, tp)| {
let resid = 1.0 - (two_pi * (tp - s * t - phi)).cos();
(t, resid)
})
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(t, _)| t)
}
}
pub fn measure_atom_transport(
term: &SaeManifoldTerm,
layout: &CrosscoderLayout,
atom: usize,
grid_resolution: usize,
) -> Result<AtomTransportReport, String> {
measure_atom_transport_between(
term,
layout,
atom,
CrosscoderLayer::Anchor,
CrosscoderLayer::Block(0),
grid_resolution,
)
}
pub fn measure_atom_transport_between(
term: &SaeManifoldTerm,
layout: &CrosscoderLayout,
atom: usize,
source: CrosscoderLayer,
target: CrosscoderLayer,
grid_resolution: usize,
) -> Result<AtomTransportReport, String> {
if atom >= term.atoms.len() {
return Err(format!(
"measure_atom_transport: atom index {atom} out of range (K = {})",
term.atoms.len()
));
}
if layout.total_dim() != term.output_dim() {
return Err(format!(
"measure_atom_transport: layout total width {} != term output_dim {} (the layout \
must describe this term's augmented columns)",
layout.total_dim(),
term.output_dim()
));
}
let atom_ref = &term.atoms[atom];
if atom_ref.latent_dim != 1 {
return Err(format!(
"measure_atom_transport: the phase-shift law is defined for a 1-D circle atom; atom \
{atom} has latent_dim {}",
atom_ref.latent_dim
));
}
if atom_ref.basis_kind != SaeAtomBasisKind::Periodic {
return Err(format!(
"measure_atom_transport: atom {atom} must use the standard periodic harmonic basis, got {:?}",
atom_ref.basis_kind
));
}
if atom_ref.homotopy_eta != 1.0 {
return Err(format!(
"measure_atom_transport: atom {atom} is at homotopy eta {}, not the fitted eta = 1 endpoint",
atom_ref.homotopy_eta
));
}
let physical_decoder = atom_ref.full_width_decoder();
let b_src = honest_layer_decoder(&physical_decoder, layout, source)?;
let b_tgt = honest_layer_decoder(&physical_decoder, layout, target)?;
if b_src.ncols() != b_tgt.ncols() {
return Err(format!(
"measure_atom_transport: source ambient width {} != target ambient width {} — the \
nearest-point transport needs both layer images in one ambient space (a crosscoder \
shares the residual-stream dimension across layers)",
b_src.ncols(),
b_tgt.ncols()
));
}
let m = physical_decoder.nrows();
let n_harmonics = m.saturating_sub(1) / 2;
if grid_resolution == 0 {
return Err("measure_atom_transport: grid_resolution must be positive".to_string());
}
let alt_harmonics = (2 * n_harmonics + 1).max(grid_resolution / 16);
let smooth_coefficient_count = 2 * alt_harmonics + 1;
let required_fit_samples = smooth_coefficient_count + 1;
let diagnostic_multiplier = required_fit_samples.div_ceil(grid_resolution).max(1);
let diagnostic_resolution = grid_resolution
.checked_mul(diagnostic_multiplier)
.ok_or_else(|| "measure_atom_transport: diagnostic grid size overflow".to_string())?;
let basis = ChartBasisKind::Periodic { n_harmonics };
let grid = Array2::<f64>::from_shape_fn((diagnostic_resolution, 1), |(g, _)| {
g as f64 / diagnostic_resolution as f64
});
if basis.width() != m {
return Err(format!(
"measure_atom_transport: periodic basis width {} != physical decoder width {m}",
basis.width()
));
}
let mut phi_grid = Array2::<f64>::zeros((diagnostic_resolution, m));
let mut phi = vec![0.0; m];
for g in 0..diagnostic_resolution {
basis.eval_into(grid[[g, 0]], &mut phi);
for column in 0..m {
phi_grid[[g, column]] = phi[column];
}
}
let source_image = phi_grid.dot(&b_src); let target_gram = b_tgt.dot(&b_tgt.t());
let target_extrema = PeriodicCurveExtrema::from_gram(target_gram.view())?;
let linear_all = source_image.dot(&b_tgt.t()); use rayon::prelude::*;
let tprime: Vec<f64> = (0..diagnostic_resolution)
.into_par_iter()
.map(|g| {
let linear = linear_all.row(g);
let projection = target_extrema
.minimize_squared_distance(linear.as_slice().ok_or_else(|| {
"measure_atom_transport: target linear coefficients are not contiguous"
.to_string()
})?)
.map_err(|error| {
format!("measure_atom_transport: source sample {g} target projection: {error}")
})?;
Ok(projection.coordinate)
})
.collect::<Result<Vec<f64>, String>>()?;
let t_arr: Vec<f64> = (0..diagnostic_resolution)
.map(|g| g as f64 / diagnostic_resolution as f64)
.collect();
let transport_grid: Vec<(f64, f64)> = (0..grid_resolution)
.map(|g| {
let diagnostic_index = g * diagnostic_multiplier;
(t_arr[diagnostic_index], tprime[diagnostic_index])
})
.collect();
let (phase_shift, phase_r2, smooth_r2) = fit_transport_law(&t_arr, &tprime, alt_harmonics);
let drift = decoder_drift(&b_src, &b_tgt);
let principal_angles = principal_angles_between_images(&b_src, &b_tgt)?;
Ok(AtomTransportReport {
atom,
source,
target,
grid_resolution,
n_harmonics,
phase_shift,
phase_r2,
smooth_r2,
drift,
principal_angles,
transport_grid,
})
}
pub(crate) fn honest_layer_decoder(
decoder: &Array2<f64>,
layout: &CrosscoderLayout,
layer: CrosscoderLayer,
) -> Result<Array2<f64>, String> {
match layer {
CrosscoderLayer::Anchor => Ok(decoder.slice(s![.., 0..layout.anchor_dim()]).to_owned()),
CrosscoderLayer::Block(l) => {
if l >= layout.num_blocks() {
return Err(format!(
"measure_atom_transport: block index ℓ={l} out of range (L−1 = {})",
layout.num_blocks()
));
}
let inv = 1.0 / layout.sqrt_lambda(l);
Ok(decoder
.slice(s![.., layout.block_range(l)])
.mapv(|v| inv * v))
}
}
}
fn fit_transport_law(t: &[f64], tprime: &[f64], n_harmonics: usize) -> ((f64, f64), f64, f64) {
let two_pi = std::f64::consts::TAU;
let g = t.len();
let gf = g as f64;
let (sum_sin, sum_cos) = tprime.iter().fold((0.0, 0.0), |(s, c), &v| {
(s + (two_pi * v).sin(), c + (two_pi * v).cos())
});
let r_tot = (sum_sin * sum_sin + sum_cos * sum_cos).sqrt();
let ss_tot = gf - r_tot;
let mut best_s = 1.0_f64;
let mut best_phi = 0.0_f64;
let mut best_ss_res = f64::INFINITY;
for &s in &[1.0_f64, -1.0_f64] {
let (su, cu) = t
.iter()
.zip(tprime.iter())
.fold((0.0, 0.0), |(a, b), (&ti, &tpi)| {
let u = tpi - s * ti;
(a + (two_pi * u).sin(), b + (two_pi * u).cos())
});
let r_u = (su * su + cu * cu).sqrt();
let ss_res = gf - r_u;
if ss_res < best_ss_res {
best_ss_res = ss_res;
best_s = s;
best_phi = wrap_half(su.atan2(cu) / two_pi);
}
}
let phase_r2 = circular_r2(ss_tot, best_ss_res);
let smooth_r2 = fit_smooth_alternative(t, tprime, best_s, best_phi, n_harmonics, ss_tot)
.unwrap_or(phase_r2);
((best_s, best_phi), phase_r2, smooth_r2)
}
fn fit_smooth_alternative(
t: &[f64],
tprime: &[f64],
s: f64,
phi: f64,
n_harmonics: usize,
ss_tot: f64,
) -> Option<f64> {
let two_pi = std::f64::consts::TAU;
let k = 2 * n_harmonics + 1;
let design = |ti: f64| -> Vec<f64> {
let mut row = Vec::with_capacity(k);
row.push(1.0);
for h in 1..=n_harmonics {
let a = two_pi * h as f64 * ti;
row.push(a.sin());
row.push(a.cos());
}
row
};
let mut dtd = Array2::<f64>::zeros((k, k));
let mut dtb = Array1::<f64>::zeros(k);
for (&ti, &tpi) in t.iter().zip(tprime.iter()) {
let row = design(ti);
let d = wrap_half(tpi - s * ti - phi);
for i in 0..k {
dtb[i] += row[i] * d;
for j in 0..k {
dtd[[i, j]] += row[i] * row[j];
}
}
}
let coeffs = solve_spd(&dtd, &dtb)?;
let mut ss_res = 0.0_f64;
for (&ti, &tpi) in t.iter().zip(tprime.iter()) {
let row = design(ti);
let f: f64 = row.iter().zip(coeffs.iter()).map(|(&r, &c)| r * c).sum();
let pred = s * ti + phi + f;
ss_res += 1.0 - (two_pi * (tpi - pred)).cos();
}
Some(circular_r2(ss_tot, ss_res))
}
fn circular_r2(ss_tot: f64, ss_res: f64) -> f64 {
if ss_tot > 0.0 {
1.0 - ss_res / ss_tot
} else {
f64::NAN
}
}
fn wrap_half(x: f64) -> f64 {
let r = x.rem_euclid(1.0);
if r >= 0.5 { r - 1.0 } else { r }
}
pub(crate) fn decoder_drift(b_src: &Array2<f64>, b_tgt: &Array2<f64>) -> f64 {
let fro = |a: &Array2<f64>| a.iter().map(|&v| v * v).sum::<f64>().sqrt();
let ns = fro(b_src);
let nt = fro(b_tgt);
let dead = 1e-12 * ns.max(nt);
if ns > dead && nt > dead {
let diff: f64 = b_src
.iter()
.zip(b_tgt.iter())
.map(|(&a, &b)| (a - b) * (a - b))
.sum::<f64>()
.sqrt();
diff / (ns * nt).sqrt()
} else {
f64::NAN
}
}
pub(crate) fn principal_angles_between_images(
b_src: &Array2<f64>,
b_tgt: &Array2<f64>,
) -> Result<Vec<f64>, String> {
let q_src = orthonormal_row_basis(b_src)?; let q_tgt = orthonormal_row_basis(b_tgt)?; let r_src = q_src.nrows();
let r_tgt = q_tgt.nrows();
if r_src == 0 || r_tgt == 0 {
return Ok(vec![std::f64::consts::FRAC_PI_2; r_src.max(r_tgt)]);
}
let cross = q_src.dot(&q_tgt.t()); let (_u, svals, _vt) = cross
.svd(false, false)
.map_err(|e| format!("principal_angles_between_images: SVD failed: {e}"))?;
let mut angles = svals
.iter()
.map(|&sv| sv.clamp(0.0, 1.0).acos())
.collect::<Vec<f64>>();
angles.extend(std::iter::repeat(std::f64::consts::FRAC_PI_2).take(r_src.abs_diff(r_tgt)));
Ok(angles)
}
fn orthonormal_row_basis(b: &Array2<f64>) -> Result<Array2<f64>, String> {
let (_u, svals, vt) = b
.svd(false, true)
.map_err(|e| format!("orthonormal_row_basis: SVD failed: {e}"))?;
let vt = vt.ok_or_else(|| "orthonormal_row_basis: SVD returned no right factor".to_string())?;
let smax = svals.iter().cloned().fold(0.0_f64, f64::max);
let tol = smax * (b.nrows().max(b.ncols()) as f64) * f64::EPSILON;
let rank = svals.iter().filter(|&&s| s > tol).count();
Ok(vt.slice(s![0..rank, ..]).to_owned())
}
fn solve_spd(a: &Array2<f64>, b: &Array1<f64>) -> Option<Array1<f64>> {
let k = a.nrows();
let mut l = Array2::<f64>::zeros((k, k));
for i in 0..k {
for j in 0..=i {
let mut sum = a[[i, j]];
for p in 0..j {
sum -= l[[i, p]] * l[[j, p]];
}
if i == j {
if !(sum > 0.0) {
return None;
}
l[[i, j]] = sum.sqrt();
} else {
l[[i, j]] = sum / l[[j, j]];
}
}
}
let mut y = Array1::<f64>::zeros(k);
for i in 0..k {
let mut sum = b[i];
for p in 0..i {
sum -= l[[i, p]] * y[p];
}
y[i] = sum / l[[i, i]];
}
let mut c = Array1::<f64>::zeros(k);
for i in (0..k).rev() {
let mut sum = y[i];
for p in (i + 1)..k {
sum -= l[[p, i]] * c[p];
}
c[i] = sum / l[[i, i]];
}
Some(c)
}