use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
#[derive(Clone, Debug)]
pub struct SphereTangentEmbedding {
basepoint: Array1<f64>,
tangent_basis: Array2<f64>,
}
impl SphereTangentEmbedding {
pub fn fit(prob_rows: ArrayView2<'_, f64>) -> Result<(Self, Array2<f64>), String> {
let (n, v) = prob_rows.dim();
if n == 0 || v < 2 {
return Err(format!(
"SphereTangentEmbedding::fit: need n ≥ 1 rows and V ≥ 2 tokens; got ({n}, {v})"
));
}
let mut q = Array2::<f64>::zeros((n, v));
let mut mean = Array1::<f64>::zeros(v);
for i in 0..n {
let row = prob_rows.row(i);
let mut sum = 0.0_f64;
for &value in row.iter() {
if !(value.is_finite() && value >= 0.0) {
return Err(format!(
"SphereTangentEmbedding::fit: row {i} has a non-finite or negative \
probability entry ({value})"
));
}
sum += value;
}
if !(sum > 0.0) {
return Err(format!(
"SphereTangentEmbedding::fit: row {i} sums to {sum}; a behavioral summary \
must have positive mass"
));
}
let inv_sqrt_sum = 1.0 / sum.sqrt();
let mut q_row = q.row_mut(i);
for j in 0..v {
let qij = prob_rows[[i, j]].sqrt() * inv_sqrt_sum;
q_row[j] = qij;
mean[j] += qij;
}
}
let mean_norm = mean.dot(&mean).sqrt();
if !(mean_norm > 0.0) {
return Err(
"SphereTangentEmbedding::fit: the extrinsic mean of the half-densities is the \
zero vector (antipodally balanced behavior); no basepoint is defined"
.to_string(),
);
}
let basepoint = &mean / mean_norm;
let tangent_basis = tangent_basis_orthogonal_to(basepoint.view())?;
let root_two = std::f64::consts::SQRT_2;
let mut target = q.dot(&tangent_basis);
target.mapv_inplace(|value| root_two * value);
Ok((
Self {
basepoint,
tangent_basis,
},
target,
))
}
pub fn vocab(&self) -> usize {
self.basepoint.len()
}
pub fn behavior_dim(&self) -> usize {
self.tangent_basis.ncols()
}
pub fn decode_sphere(&self, y: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
let py = self.behavior_dim();
if y.len() != py {
return Err(format!(
"SphereTangentEmbedding::decode_sphere: coordinate has length {}; chart tangent \
dim is {py}",
y.len()
));
}
let inv_root_two = std::f64::consts::FRAC_1_SQRT_2;
let c = &y.to_owned() * inv_root_two;
let tangent = self.tangent_basis.dot(&c);
let radial_sq = 1.0 - c.dot(&c);
let radial = if radial_sq > 0.0 {
radial_sq.sqrt()
} else {
0.0
};
let mut q = &tangent + &(&self.basepoint * radial);
let norm = q.dot(&q).sqrt();
if norm > 0.0 {
q.mapv_inplace(|value| value / norm);
}
Ok(q)
}
pub fn decode(&self, y: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
let q = self.decode_sphere(y)?;
Ok(q.mapv(|value| value * value))
}
pub fn decode_rows(&self, y: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
if y.ncols() != self.behavior_dim() {
return Err(format!(
"SphereTangentEmbedding::decode_rows: coordinates have {} columns; chart tangent dim is {}",
y.ncols(),
self.behavior_dim()
));
}
let mut probabilities = Array2::<f64>::zeros((y.nrows(), self.vocab()));
for row in 0..y.nrows() {
let decoded = self.decode(y.row(row))?;
probabilities.row_mut(row).assign(&decoded);
}
Ok(probabilities)
}
pub fn predicted_nats(delta_y: ArrayView1<'_, f64>) -> f64 {
delta_y.dot(&delta_y)
}
pub fn exact_kl(p_a: ArrayView1<'_, f64>, p_b: ArrayView1<'_, f64>) -> Result<f64, String> {
if p_a.len() != p_b.len() {
return Err(format!(
"SphereTangentEmbedding::exact_kl: length mismatch {} vs {}",
p_a.len(),
p_b.len()
));
}
let mut kl = 0.0_f64;
for (&a, &b) in p_a.iter().zip(p_b.iter()) {
if a > 0.0 {
kl += a * (a / b).ln();
}
}
Ok(kl)
}
}
fn tangent_basis_orthogonal_to(axis: ArrayView1<'_, f64>) -> Result<Array2<f64>, String> {
let v = axis.len();
if v < 2 {
return Err(format!("tangent_basis_orthogonal_to: need V ≥ 2; got {v}"));
}
let mut pivot = 0usize;
let mut best = axis[0].abs();
for j in 1..v {
let a = axis[j].abs();
if a > best {
best = a;
pivot = j;
}
}
let mut w = axis.to_owned();
w.mapv_inplace(|value| -value);
w[pivot] += 1.0;
let w_norm = f64::sqrt(w.dot(&w));
if !(w_norm > 0.0) {
w.fill(0.0);
} else {
w.mapv_inplace(|value| value / w_norm);
}
let mut basis = Array2::<f64>::zeros((v, v - 1));
let mut col = 0usize;
for j in 0..v {
if j == pivot {
continue;
}
let two_wj = 2.0 * w[j];
for i in 0..v {
let e_ij = if i == j { 1.0 } else { 0.0 };
basis[[i, col]] = e_ij - two_wj * w[i];
}
col += 1;
}
Ok(basis)
}
#[derive(Clone, Debug)]
pub struct BehaviorBlock {
pub embedding: SphereTangentEmbedding,
pub target: Array2<f64>,
pub activation_dim: usize,
log_lambda_y: f64,
lambda_y: f64,
sqrt_lambda_y: f64,
}
impl BehaviorBlock {
pub fn fit(
prob_rows: ArrayView2<'_, f64>,
activation_dim: usize,
log_lambda_y: f64,
) -> Result<Self, String> {
if activation_dim == 0 {
return Err("BehaviorBlock::fit: activation_dim must be positive".into());
}
let lambda_y = gam_problem::checked_exp_log_strength(log_lambda_y)
.map_err(|error| format!("BehaviorBlock::fit: {error}"))?;
let sqrt_lambda_y = gam_problem::checked_exp_log_strength(0.5 * log_lambda_y)
.map_err(|error| format!("BehaviorBlock::fit square-root strength: {error}"))?;
let (embedding, target) = SphereTangentEmbedding::fit(prob_rows)?;
Ok(Self {
embedding,
target,
activation_dim,
log_lambda_y,
lambda_y,
sqrt_lambda_y,
})
}
pub fn behavior_dim(&self) -> usize {
self.embedding.behavior_dim()
}
pub fn augmented_dim(&self) -> usize {
self.activation_dim + self.behavior_dim()
}
pub fn lambda_y(&self) -> f64 {
self.lambda_y
}
pub fn sqrt_lambda_y(&self) -> f64 {
self.sqrt_lambda_y
}
pub fn log_lambda_y(&self) -> f64 {
self.log_lambda_y
}
pub fn augmented_target(&self, activation: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
let (n, px) = activation.dim();
if px != self.activation_dim {
return Err(format!(
"BehaviorBlock::augmented_target: activation has {px} columns; block activation_dim \
is {}",
self.activation_dim
));
}
if self.target.nrows() != n {
return Err(format!(
"BehaviorBlock::augmented_target: activation has {n} rows but behavior target has {}",
self.target.nrows()
));
}
let py = self.behavior_dim();
let sqrt_lambda = self.sqrt_lambda_y();
let mut augmented = Array2::<f64>::zeros((n, px + py));
for i in 0..n {
for j in 0..px {
augmented[[i, j]] = activation[[i, j]];
}
for j in 0..py {
augmented[[i, px + j]] = sqrt_lambda * self.target[[i, j]];
}
}
Ok(augmented)
}
pub fn split_decoder(
&self,
augmented_decoder: ArrayView2<'_, f64>,
) -> Result<(Array2<f64>, Array2<f64>), String> {
let px = self.activation_dim;
let py = self.behavior_dim();
let (m, p_tot) = augmented_decoder.dim();
if p_tot != px + py {
return Err(format!(
"BehaviorBlock::split_decoder: decoder has {p_tot} output columns; expected \
p_x + p_y = {px} + {py} = {}",
px + py
));
}
let inv_sqrt_lambda = 1.0 / self.sqrt_lambda_y();
let mut b = Array2::<f64>::zeros((m, px));
let mut c = Array2::<f64>::zeros((m, py));
for row in 0..m {
for j in 0..px {
b[[row, j]] = augmented_decoder[[row, j]];
}
for j in 0..py {
c[[row, j]] = inv_sqrt_lambda * augmented_decoder[[row, px + j]];
}
}
Ok((b, c))
}
pub fn with_log_lambda_y(&self, log_lambda_y: f64) -> Result<Self, String> {
let lambda_y = gam_problem::checked_exp_log_strength(log_lambda_y)
.map_err(|error| format!("BehaviorBlock::with_log_lambda_y: {error}"))?;
let sqrt_lambda_y =
gam_problem::checked_exp_log_strength(0.5 * log_lambda_y).map_err(|error| {
format!("BehaviorBlock::with_log_lambda_y square-root strength: {error}")
})?;
let mut next = self.clone();
next.log_lambda_y = log_lambda_y;
next.lambda_y = lambda_y;
next.sqrt_lambda_y = sqrt_lambda_y;
Ok(next)
}
}
#[derive(Clone, Debug)]
pub struct OutputBlock {
pub label: String,
pub target: Array2<f64>,
log_lambda: f64,
lambda: f64,
sqrt_lambda: f64,
}
impl OutputBlock {
pub fn new(
label: impl Into<String>,
target: Array2<f64>,
log_lambda: f64,
) -> Result<Self, String> {
let (n, p) = target.dim();
if n == 0 || p == 0 {
return Err(format!(
"OutputBlock::new: target must be a non-empty (n × p_ℓ) matrix; got ({n}, {p})"
));
}
let lambda = gam_problem::checked_exp_log_strength(log_lambda)
.map_err(|error| format!("OutputBlock::new: {error}"))?;
let sqrt_lambda = gam_problem::checked_exp_log_strength(0.5 * log_lambda)
.map_err(|error| format!("OutputBlock::new square-root strength: {error}"))?;
Ok(Self {
label: label.into(),
target,
log_lambda,
lambda,
sqrt_lambda,
})
}
pub fn block_dim(&self) -> usize {
self.target.ncols()
}
pub fn lambda(&self) -> f64 {
self.lambda
}
pub fn sqrt_lambda(&self) -> f64 {
self.sqrt_lambda
}
pub fn log_lambda(&self) -> f64 {
self.log_lambda
}
pub fn with_log_lambda(&self, log_lambda: f64) -> Result<Self, String> {
let lambda = gam_problem::checked_exp_log_strength(log_lambda)
.map_err(|error| format!("OutputBlock::with_log_lambda: {error}"))?;
let sqrt_lambda =
gam_problem::checked_exp_log_strength(0.5 * log_lambda).map_err(|error| {
format!("OutputBlock::with_log_lambda square-root strength: {error}")
})?;
let mut next = self.clone();
next.log_lambda = log_lambda;
next.lambda = lambda;
next.sqrt_lambda = sqrt_lambda;
Ok(next)
}
pub fn split_honest_decoder(&self, scaled_decoder: ArrayView2<'_, f64>) -> Array2<f64> {
let inv = 1.0 / self.sqrt_lambda();
scaled_decoder.mapv(|value| inv * value)
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct CrosscoderLayout {
p_x: usize,
block_dims: Vec<usize>,
labels: Vec<String>,
block_log_lambda: Vec<f64>,
block_sqrt_lambda: Vec<f64>,
}
impl CrosscoderLayout {
pub fn new(
p_x: usize,
block_dims: Vec<usize>,
labels: Vec<String>,
block_log_lambda: Vec<f64>,
) -> Result<Self, String> {
if p_x == 0 {
return Err("CrosscoderLayout::new: anchor width p_x must be non-zero".to_string());
}
if block_dims.len() != labels.len() || block_dims.len() != block_log_lambda.len() {
return Err(format!(
"CrosscoderLayout::new: block_dims ({}), labels ({}), and block_log_lambda ({}) \
must have equal length",
block_dims.len(),
labels.len(),
block_log_lambda.len()
));
}
for (l, &dim) in block_dims.iter().enumerate() {
if dim == 0 {
return Err(format!(
"CrosscoderLayout::new: block {l} ('{}') has width 0",
labels[l]
));
}
}
gam_problem::validate_log_strengths(block_log_lambda.iter().copied()).map_err(|error| {
format!(
"CrosscoderLayout::new: block {} ('{}') has invalid log λ: {error}",
error.coordinate, labels[error.coordinate]
)
})?;
let block_sqrt_lambda = block_log_lambda
.iter()
.copied()
.map(|log_lambda| {
gam_problem::checked_exp_log_strength(0.5 * log_lambda)
.expect("half of a validated log strength remains canonical")
})
.collect();
Ok(Self {
p_x,
block_dims,
labels,
block_log_lambda,
block_sqrt_lambda,
})
}
pub fn from_blocks(p_x: usize, blocks: &[OutputBlock]) -> Self {
Self {
p_x,
block_dims: blocks.iter().map(|b| b.block_dim()).collect(),
labels: blocks.iter().map(|b| b.label.clone()).collect(),
block_log_lambda: blocks.iter().map(OutputBlock::log_lambda).collect(),
block_sqrt_lambda: blocks.iter().map(OutputBlock::sqrt_lambda).collect(),
}
}
pub fn anchor_dim(&self) -> usize {
self.p_x
}
pub fn num_blocks(&self) -> usize {
self.block_dims.len()
}
pub fn block_dims(&self) -> &[usize] {
&self.block_dims
}
pub fn labels(&self) -> &[String] {
&self.labels
}
pub fn block_log_lambda(&self) -> &[f64] {
&self.block_log_lambda
}
pub fn total_dim(&self) -> usize {
self.p_x + self.block_dims.iter().sum::<usize>()
}
pub fn block_range(&self, l: usize) -> std::ops::Range<usize> {
assert!(
l < self.block_dims.len(),
"CrosscoderLayout::block_range: block {l} out of range (L-1 = {})",
self.block_dims.len()
);
let start = self.p_x + self.block_dims[..l].iter().sum::<usize>();
start..start + self.block_dims[l]
}
pub fn log_lambda(&self, l: usize) -> f64 {
self.block_log_lambda[l]
}
pub fn sqrt_lambda(&self, l: usize) -> f64 {
self.block_sqrt_lambda[l]
}
}
pub fn stack_augmented_target(
anchor: ArrayView2<'_, f64>,
blocks: &[OutputBlock],
) -> Result<Array2<f64>, String> {
let (n, px) = anchor.dim();
if n == 0 || px == 0 {
return Err(format!(
"stack_augmented_target: anchor must be a non-empty (n × p_x) matrix; got ({n}, {px})"
));
}
for block in blocks {
if block.target.nrows() != n {
return Err(format!(
"stack_augmented_target: block '{}' has {} rows but anchor has {n}",
block.label,
block.target.nrows()
));
}
}
let layout = CrosscoderLayout::from_blocks(px, blocks);
let mut augmented = Array2::<f64>::zeros((n, layout.total_dim()));
for i in 0..n {
for j in 0..px {
augmented[[i, j]] = anchor[[i, j]];
}
for (l, block) in blocks.iter().enumerate() {
let sqrt_lambda = layout.sqrt_lambda(l);
for (jj, col) in layout.block_range(l).enumerate() {
augmented[[i, col]] = sqrt_lambda * block.target[[i, jj]];
}
}
}
Ok(augmented)
}
pub fn profiled_penalized_quasi_laplace_criterion(
n_obs: usize,
p_x: usize,
rss_x: f64,
block_rss_unscaled: &[f64],
block_dims: &[usize],
block_log_lambda: &[f64],
penalty_energy: f64,
) -> Result<f64, String> {
let lambdas = gam_problem::checked_exp_log_strengths(block_log_lambda.iter().copied())
.map_err(|error| format!("profiled block criterion: {error}"))?;
let n = n_obs as f64;
let mut p_tilde = p_x as f64;
let mut pooled = rss_x;
let mut jac = 0.0_f64;
for (((&rss, &dim), &log_lambda), &lambda) in block_rss_unscaled
.iter()
.zip(block_dims.iter())
.zip(block_log_lambda.iter())
.zip(lambdas.iter())
{
pooled += lambda * rss;
p_tilde += dim as f64;
jac += (dim as f64) * log_lambda;
}
pooled += penalty_energy;
if !(pooled > 0.0) {
return Ok(f64::INFINITY);
}
Ok(0.5 * n * p_tilde * (pooled / (n * p_tilde)).ln() - 0.5 * n * jac)
}
pub fn profiled_penalized_quasi_laplace_block_efs_log_lambda_steps(
p_x: usize,
rss_x: f64,
block_rss_unscaled: &[f64],
block_dims: &[usize],
block_log_lambda: &[f64],
penalty_energy: f64,
) -> Vec<f64> {
let var_x = (rss_x + penalty_energy) / p_x as f64;
block_rss_unscaled
.iter()
.zip(block_dims.iter())
.zip(block_log_lambda.iter())
.map(|((&rss, &dim), &log_lambda)| {
if var_x > 0.0 && rss > 0.0 {
let var_y = rss / dim as f64;
(var_x / var_y).ln() - log_lambda
} else {
0.0
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::{Array1, Array2};
#[test]
fn tangent_basis_is_orthonormal_and_orthogonal_to_axis() {
let mut axis = Array1::<f64>::from(vec![0.3, -0.5, 0.2, 0.7, -0.34]);
let norm = axis.dot(&axis).sqrt();
axis.mapv_inplace(|v| v / norm);
let e = tangent_basis_orthogonal_to(axis.view()).unwrap();
assert_eq!(e.dim(), (5, 4));
for col in 0..e.ncols() {
let dot = e.column(col).dot(&axis);
assert!(dot.abs() < 1e-12, "column {col} not ⟂ axis: {dot}");
}
let gram = e.t().dot(&e);
for i in 0..4 {
for j in 0..4 {
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
(gram[[i, j]] - expected).abs() < 1e-12,
"EᵀE[{i},{j}] = {} != {expected}",
gram[[i, j]]
);
}
}
}
#[test]
fn embed_decode_round_trips_distributions() {
let rows = vec![
vec![0.4, 0.2, 0.1, 0.1, 0.1, 0.1],
vec![0.1, 0.5, 0.1, 0.1, 0.1, 0.1],
vec![0.2, 0.2, 0.2, 0.2, 0.1, 0.1],
vec![0.05, 0.05, 0.6, 0.1, 0.1, 0.1],
];
let n = rows.len();
let v = rows[0].len();
let mut p = Array2::<f64>::zeros((n, v));
for (i, row) in rows.iter().enumerate() {
for (j, &value) in row.iter().enumerate() {
p[[i, j]] = value;
}
}
let (chart, y) = SphereTangentEmbedding::fit(p.view()).unwrap();
assert_eq!(chart.behavior_dim(), v - 1);
for i in 0..n {
let decoded = chart.decode(y.row(i)).unwrap();
for j in 0..v {
assert!(
(decoded[j] - p[[i, j]]).abs() < 1e-10,
"row {i} token {j}: decoded {} != original {}",
decoded[j],
p[[i, j]]
);
}
}
}
#[test]
fn predicted_nats_matches_exact_kl_to_second_order() {
let base = Array1::from(vec![0.25, 0.25, 0.2, 0.15, 0.15]);
let v = base.len();
let dir = Array1::from(vec![0.1, -0.05, -0.02, -0.02, -0.01]);
let make = |eps: f64| -> Array2<f64> {
let mut p = Array2::<f64>::zeros((2, v));
for j in 0..v {
p[[0, j]] = base[j];
p[[1, j]] = base[j] + eps * dir[j];
}
p
};
let mut prev_rel: Option<f64> = None;
for &eps in &[0.2_f64, 0.1, 0.05, 0.025] {
let p = make(eps);
let (chart, y) = SphereTangentEmbedding::fit(p.view()).unwrap();
let delta_y = &y.row(1).to_owned() - &y.row(0).to_owned();
let predicted = SphereTangentEmbedding::predicted_nats(delta_y.view());
let p0 = chart.decode(y.row(0)).unwrap();
let p1 = chart.decode(y.row(1)).unwrap();
let kl = SphereTangentEmbedding::exact_kl(p1.view(), p0.view()).unwrap();
let rel = (predicted - kl).abs() / kl.max(1e-12);
if let Some(prev) = prev_rel {
assert!(
rel < prev * 0.6,
"relative KL error did not fall second-order: {prev} → {rel} at ε={eps}"
);
}
prev_rel = Some(rel);
}
}
#[test]
fn constant_behavior_has_zero_tangent_target() {
let base = vec![0.3, 0.3, 0.2, 0.2];
let n = 5;
let v = base.len();
let mut p = Array2::<f64>::zeros((n, v));
for i in 0..n {
for j in 0..v {
p[[i, j]] = base[j];
}
}
let (_chart, y) = SphereTangentEmbedding::fit(p.view()).unwrap();
for value in y.iter() {
assert!(
value.abs() < 1e-12,
"constant behavior gave nonzero target {value}"
);
}
}
}