use rand::RngExt;
use rand::SeedableRng;
use rand::rngs::StdRng;
use std::collections::BTreeMap;
use crate::atom_codes::SparseAtomCodes;
pub const NULL_REPLICATES: usize = 200;
pub struct CurveballSampler {
rows: Vec<Vec<usize>>,
rng: StdRng,
}
impl CurveballSampler {
pub fn from_codes(codes: &SparseAtomCodes) -> Self {
let n_atoms = codes.k_atoms();
let mut rows: Vec<Vec<usize>> = Vec::with_capacity(codes.n_obs());
let mut seed = gam_linalg::utils::splitmix64_hash(
(codes.n_obs() as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ (n_atoms as u64),
);
for code in codes.iter() {
let active: Vec<usize> = code.active_mask.iter_ones().collect();
for &a in &active {
seed = gam_linalg::utils::splitmix64_hash(seed ^ (a as u64).wrapping_add(1));
}
seed = gam_linalg::utils::splitmix64_hash(seed ^ 0xD1B5_4A32_D192_ED03);
rows.push(active);
}
Self {
rows,
rng: StdRng::seed_from_u64(seed),
}
}
pub fn n_ones(&self) -> usize {
self.rows.iter().map(|r| r.len()).sum()
}
pub fn n_rows(&self) -> usize {
self.rows.len()
}
pub fn trade(&mut self) {
let n = self.rows.len();
if n < 2 {
return;
}
let i = self.rng.random_range(0..n);
let mut j = self.rng.random_range(0..n - 1);
if j >= i {
j += 1;
}
let (a, b) = (&self.rows[i], &self.rows[j]);
let mut shared_i: Vec<usize> = Vec::new();
let mut pool: Vec<usize> = Vec::new();
let (mut p, mut q) = (0usize, 0usize);
let mut n_from_i = 0usize;
while p < a.len() && q < b.len() {
match a[p].cmp(&b[q]) {
std::cmp::Ordering::Equal => {
shared_i.push(a[p]);
p += 1;
q += 1;
}
std::cmp::Ordering::Less => {
pool.push(a[p]);
n_from_i += 1;
p += 1;
}
std::cmp::Ordering::Greater => {
pool.push(b[q]);
q += 1;
}
}
}
while p < a.len() {
pool.push(a[p]);
n_from_i += 1;
p += 1;
}
while q < b.len() {
pool.push(b[q]);
q += 1;
}
if pool.is_empty() || n_from_i == 0 || n_from_i == pool.len() {
return;
}
let m = pool.len();
for t in 0..n_from_i {
let swap = t + self.rng.random_range(0..(m - t));
pool.swap(t, swap);
}
let build = |shared: &[usize], extra: &[usize]| -> Vec<usize> {
let mut v: Vec<usize> = Vec::with_capacity(shared.len() + extra.len());
v.extend_from_slice(shared);
v.extend_from_slice(extra);
v.sort_unstable();
v
};
let new_i = build(&shared_i, &pool[..n_from_i]);
let new_j = build(&shared_i, &pool[n_from_i..]);
self.rows[i] = new_i;
self.rows[j] = new_j;
}
pub fn mix(&mut self, trades: usize) {
for _ in 0..trades {
self.trade();
}
}
fn accumulate_selected(
&self,
pair_to_pos: &BTreeMap<(usize, usize), usize>,
joint: &mut [f64],
) {
for row in &self.rows {
for (idx, &u) in row.iter().enumerate() {
for &v in &row[idx + 1..] {
let key = if u < v { (u, v) } else { (v, u) };
if let Some(&pos) = pair_to_pos.get(&key) {
joint[pos] += 1.0;
}
}
}
}
}
}
#[derive(Clone, Debug)]
pub struct CoactivationExceedance {
n_obs: usize,
}
impl CoactivationExceedance {
pub fn n_obs(&self) -> usize {
self.n_obs
}
}
const NULL_SD_FLOOR: f64 = 1e-9;
pub fn coactivation_exceedance_for_pairs(
codes: &SparseAtomCodes,
pairs: &[(usize, usize)],
replicates: usize,
) -> Vec<f64> {
let g = codes.k_atoms();
let n_obs = codes.n_obs();
let mut pair_to_pos = BTreeMap::new();
let mut canonical = Vec::with_capacity(pairs.len());
for &(a, b) in pairs {
if a == b || a >= g || b >= g {
canonical.push(None);
continue;
}
let key = if a < b { (a, b) } else { (b, a) };
let next = pair_to_pos.len();
let pos = *pair_to_pos.entry(key).or_insert(next);
canonical.push(Some(pos));
}
let m = pair_to_pos.len();
if m == 0 {
return vec![0.0; pairs.len()];
}
let mut obs = vec![0.0_f64; m];
let sampler = CurveballSampler::from_codes(codes);
sampler.accumulate_selected(&pair_to_pos, &mut obs);
if g < 2 || n_obs < 2 || replicates == 0 {
return vec![0.0; pairs.len()];
}
let mut sampler = CurveballSampler::from_codes(codes);
let sweep = sampler.n_ones().max(sampler.n_rows());
sampler.mix(sweep);
let mut mean = vec![0.0_f64; m];
let mut m2 = vec![0.0_f64; m];
let mut scratch = vec![0.0_f64; m];
for r in 0..replicates {
sampler.mix(sweep);
for value in scratch.iter_mut() {
*value = 0.0;
}
sampler.accumulate_selected(&pair_to_pos, &mut scratch);
let count = (r + 1) as f64;
for pos in 0..m {
let x = scratch[pos];
let delta = x - mean[pos];
mean[pos] += delta / count;
m2[pos] += delta * (x - mean[pos]);
}
}
let denom = (replicates.saturating_sub(1)).max(1) as f64;
let mut sparse_z = vec![0.0_f64; m];
for pos in 0..m {
let var = m2[pos] / denom;
let sd = var.max(0.0).sqrt();
sparse_z[pos] = if sd > NULL_SD_FLOOR {
(obs[pos] - mean[pos]) / sd
} else {
0.0
};
}
canonical
.into_iter()
.map(|pos| pos.map_or(0.0, |idx| sparse_z[idx]))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn curveball_preserves_both_margins() {
let n = 60usize;
let g = 12usize;
let mut codes = SparseAtomCodes::empty(n, g);
for row in 0..n {
let start = (row * 5) % g;
for off in 0..3 {
codes.row_mut(row).assign((start + off) % g, 1.0);
}
}
let row_sums: Vec<usize> = (0..n).map(|r| codes.row(r).n_active()).collect();
let mut col_sums = vec![0usize; g];
for r in 0..n {
for c in codes.row(r).active_mask.iter_ones() {
col_sums[c] += 1;
}
}
let mut s = CurveballSampler::from_codes(&codes);
s.mix(2000);
for (r, &want) in row_sums.iter().enumerate() {
assert_eq!(s.rows[r].len(), want, "row {r} sum changed");
for w in s.rows[r].windows(2) {
assert!(w[0] < w[1], "row {r} not sorted/unique");
}
}
let mut got_col = vec![0usize; g];
for row in &s.rows {
for &c in row {
got_col[c] += 1;
}
}
assert_eq!(got_col, col_sums, "column sums changed");
}
}
#[derive(Clone)]
pub struct AuditSparseRoute {
pub indices: ndarray::Array2<u32>,
pub values: ndarray::Array3<f32>,
pub n_units: usize,
pub block_size: usize,
}
impl AuditSparseRoute {
pub fn new(
indices: ndarray::Array2<u32>,
values: ndarray::Array3<f32>,
n_units: usize,
block_size: usize,
label: &str,
) -> Result<Self, String> {
if block_size == 0 {
return Err("audit_sae block_size must be >= 1".to_string());
}
if n_units == 0 {
return Err(format!(
"audit_sae {label} requires at least one routing unit"
));
}
let (n_rows, width) = indices.dim();
if n_rows == 0 || width == 0 {
return Err(format!(
"audit_sae {label} must be a non-empty N×s route; got {:?}",
indices.dim()
));
}
if values.shape() != [n_rows, width, block_size] {
return Err(format!(
"audit_sae {label} values shape {:?} does not match indices {:?} and block_size {block_size}",
values.shape(),
indices.dim()
));
}
for row in 0..n_rows {
let mut live = std::collections::HashSet::with_capacity(width);
for slot in 0..width {
let unit = indices[[row, slot]] as usize;
if unit >= n_units {
return Err(format!(
"audit_sae {label} index {unit} at row {row}, slot {slot} is outside 0..{n_units}"
));
}
let mut norm2 = 0.0_f64;
for offset in 0..block_size {
let value = values[[row, slot, offset]] as f64;
if !value.is_finite() {
return Err(format!(
"audit_sae {label} value at row {row}, slot {slot}, offset {offset} is not finite"
));
}
norm2 += value * value;
}
if norm2 > 0.0 && !live.insert(unit) {
return Err(format!(
"audit_sae {label} repeats live unit {unit} in row {row}"
));
}
}
}
Ok(Self {
indices,
values,
n_units,
block_size,
})
}
pub fn nrows(&self) -> usize {
self.indices.nrows()
}
pub fn width(&self) -> usize {
self.indices.ncols()
}
pub fn gate(&self, row: usize, slot: usize) -> f64 {
let mut norm2 = 0.0_f64;
for offset in 0..self.block_size {
let value = self.values[[row, slot, offset]] as f64;
norm2 += value * value;
}
norm2.sqrt()
}
pub fn reconstruct(
&self,
decoder: ndarray::ArrayView2<'_, f32>,
) -> Result<ndarray::Array2<f32>, String> {
if self.block_size == 1 {
crate::sparse_dict::reconstruct_sparse_rows(
decoder,
self.indices.view(),
self.values.index_axis(ndarray::Axis(2), 0),
)
} else {
crate::sparse_dict::reconstruct_block_sparse_rows(
decoder,
self.indices.view(),
self.values.view(),
self.block_size,
)
}
}
}
#[derive(Clone, Copy, Default)]
pub struct LiveAmplitudeMoments {
pub count: usize,
pub sum: f64,
pub sum2: f64,
}
impl LiveAmplitudeMoments {
pub fn mean(self) -> f64 {
if self.count == 0 {
0.0
} else {
self.sum / self.count as f64
}
}
pub fn sd(self) -> f64 {
if self.count < 2 {
0.0
} else {
let n = self.count as f64;
((self.sum2 - self.sum * self.sum / n) / (n - 1.0))
.max(0.0)
.sqrt()
}
}
}
pub fn live_amplitude_moments(route: &AuditSparseRoute) -> Vec<LiveAmplitudeMoments> {
let mut moments = vec![LiveAmplitudeMoments::default(); route.n_units];
for row in 0..route.nrows() {
for slot in 0..route.width() {
let gate = route.gate(row, slot);
if gate > 0.0 {
let unit = route.indices[[row, slot]] as usize;
moments[unit].count += 1;
moments[unit].sum += gate;
moments[unit].sum2 += gate * gate;
}
}
}
moments
}
pub fn resample_sparse_architecture_null<R: rand::Rng + ?Sized>(
observed: &AuditSparseRoute,
donor: &AuditSparseRoute,
rng: &mut R,
) -> Result<AuditSparseRoute, String> {
use rand::RngExt;
let observed_moments = live_amplitude_moments(observed);
let donor_moments = live_amplitude_moments(donor);
let mut indices = ndarray::Array2::<u32>::zeros((observed.nrows(), donor.width()));
let mut values =
ndarray::Array3::<f32>::zeros((observed.nrows(), donor.width(), donor.block_size));
for row in 0..observed.nrows() {
let source = rng.random_range(0..donor.nrows());
for slot in 0..donor.width() {
let unit = donor.indices[[source, slot]] as usize;
indices[[row, slot]] = unit as u32;
let gate = donor.gate(source, slot);
if gate == 0.0 {
continue;
}
let observed_moment = observed_moments[unit];
let donor_moment = donor_moments[unit];
if observed_moment.count == 0 {
continue;
}
let donor_sd = donor_moment.sd();
let target_gate = if donor_sd > 0.0 {
(observed_moment.mean()
+ (gate - donor_moment.mean()) * observed_moment.sd() / donor_sd)
.max(0.0)
} else {
observed_moment.mean()
};
if target_gate == 0.0 {
continue;
}
let scale = target_gate / gate;
for offset in 0..donor.block_size {
values[[row, slot, offset]] =
(donor.values[[source, slot, offset]] as f64 * scale) as f32;
}
}
}
AuditSparseRoute::new(
indices,
values,
observed.n_units,
observed.block_size,
"architecture-matched null route",
)
}