use gam_math::probability::normal_logsf;
use gam_terms::inference::structure_evidence::{ClaimKind, StructureLedger, e_benjamini_hochberg};
use ndarray::{Array2, ArrayView2, ArrayView4};
use super::wbic_audit::{ReconSpectrum, recon_spectrum};
use crate::inference::layer_transport::{ChartTopology, LayerTransportReport, fit_layer_transport};
const LAMBDA_JUMP_SE_ABS_FLOOR: f64 = 1.0e-9;
const LAMBDA_JUMP_SE_REL_FLOOR: f64 = 1.0e-6;
pub struct WbicDynamicsInput<'a> {
pub decoder_grid: ArrayView4<'a, f64>,
pub checkpoint_ids: &'a [String],
pub atom_names: &'a [String],
pub r_floor: f64,
pub birth_alpha: f64,
}
#[derive(Clone, Debug)]
pub struct LambdaJump {
pub step: usize,
pub from_ckpt: String,
pub to_ckpt: String,
pub delta_lambda: f64,
pub se: f64,
pub z: f64,
pub log_e: f64,
pub born: bool,
}
#[derive(Clone, Debug)]
pub struct AtomLambdaTrajectory {
pub atom_name: String,
pub lambda: Vec<f64>,
pub rank_soft: Vec<f64>,
pub mp_reconstruction_rank: Vec<usize>,
pub production_chargeable_rank: Vec<usize>,
pub jumps: Vec<LambdaJump>,
pub birth_evidence: StructureLedger,
pub transports: Vec<Option<LayerTransportReport>>,
}
impl AtomLambdaTrajectory {
pub fn peak_birth_log_e(&self) -> Result<f64, String> {
let n_steps = self.jumps.len();
if n_steps == 0 {
return Ok(f64::NEG_INFINITY);
}
let mut positive = Vec::new();
for jump in self.jumps.iter().filter(|jump| jump.delta_lambda > 0.0) {
if jump.log_e.is_nan() || jump.log_e == f64::INFINITY {
return Err(format!(
"atom {} has invalid positive-jump log-e value {} at step {}→{}",
self.atom_name, jump.log_e, jump.from_ckpt, jump.to_ckpt
));
}
positive.push(jump.log_e);
}
let max = positive.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
if max == f64::NEG_INFINITY {
return Ok(f64::NEG_INFINITY);
}
let sumexp: f64 = positive.iter().map(|&l| (l - max).exp()).sum();
let peak = max + sumexp.ln() - (n_steps as f64).ln();
if peak.is_finite() {
Ok(peak)
} else {
Err(format!(
"atom {} peak birth log-e mixture is non-finite: {peak}",
self.atom_name
))
}
}
}
#[derive(Clone, Debug)]
pub struct WbicDynamicsReport {
pub atoms: Vec<AtomLambdaTrajectory>,
pub cross_atom_born: Vec<usize>,
}
pub fn wbic_lambda_dynamics(input: &WbicDynamicsInput<'_>) -> Result<WbicDynamicsReport, String> {
let shape = input.decoder_grid.shape();
let (n_checkpoints, n_atoms, n_grid, ambient_dim) = (shape[0], shape[1], shape[2], shape[3]);
if n_checkpoints < 2 {
return Err(format!(
"wbic lambda dynamics needs at least two checkpoints, got {n_checkpoints}"
));
}
if input.checkpoint_ids.len() != n_checkpoints {
return Err(format!(
"checkpoint_ids length {} disagrees with decoder grid checkpoint axis {n_checkpoints}",
input.checkpoint_ids.len()
));
}
if input.atom_names.len() != n_atoms {
return Err(format!(
"atom_names length {} disagrees with decoder grid atom axis {n_atoms}",
input.atom_names.len()
));
}
if n_grid < 2 || ambient_dim == 0 {
return Err(format!(
"wbic lambda dynamics needs a non-trivial grid ({n_grid}) and ambient dim ({ambient_dim})"
));
}
if !(input.r_floor > 0.0) {
return Err(format!(
"wbic lambda dynamics needs a positive noise floor r_floor, got {}",
input.r_floor
));
}
if !(input.birth_alpha > 0.0 && input.birth_alpha < 1.0) {
return Err(format!(
"wbic lambda dynamics needs birth_alpha in (0,1), got {}",
input.birth_alpha
));
}
if input.decoder_grid.iter().any(|v| !v.is_finite()) {
return Err("wbic lambda dynamics decoder grid must be finite".to_string());
}
let mut atoms = Vec::with_capacity(n_atoms);
for atom in 0..n_atoms {
let atom_name = input.atom_names[atom].clone();
let mut lambda = Vec::with_capacity(n_checkpoints);
let mut rank_soft = Vec::with_capacity(n_checkpoints);
let mut mp_reconstruction_rank = Vec::with_capacity(n_checkpoints);
let mut production_chargeable_rank = Vec::with_capacity(n_checkpoints);
let mut lambda_se = Vec::with_capacity(n_checkpoints);
for c in 0..n_checkpoints {
let curve = input.decoder_grid.slice(ndarray::s![c, atom, .., ..]);
let spec = atom_learning_spectrum(curve, input.r_floor).map_err(|e| {
format!(
"wbic spectrum for atom '{atom_name}' checkpoint {} failed: {e}",
input.checkpoint_ids[c]
)
})?;
lambda.push(spec.learning_coefficient());
rank_soft.push(spec.rank_soft());
mp_reconstruction_rank.push(spec.mp_reconstruction_rank());
production_chargeable_rank.push(spec.production_chargeable_rank());
lambda_se.push(lambda_jackknife_se(curve, input.r_floor));
}
let mut jumps = Vec::with_capacity(n_checkpoints - 1);
let mut transports = Vec::with_capacity(n_checkpoints - 1);
let mut birth_evidence = StructureLedger::new();
for step in 0..n_checkpoints - 1 {
let c0 = step;
let c1 = step + 1;
let delta_lambda = lambda[c1] - lambda[c0];
let raw_se = (lambda_se[c0] * lambda_se[c0] + lambda_se[c1] * lambda_se[c1]).sqrt();
let scale = 0.5 * (lambda[c0].abs() + lambda[c1].abs());
let se = raw_se
.max(LAMBDA_JUMP_SE_ABS_FLOOR)
.max(LAMBDA_JUMP_SE_REL_FLOOR * scale);
let z = delta_lambda / se;
let log_e = no_jump_log_e_value(z)?;
let claim = birth_evidence.register(ClaimKind::Custom {
label: format!(
"atom '{atom_name}' λ jumped from checkpoint {} to {}",
input.checkpoint_ids[c0], input.checkpoint_ids[c1]
),
});
birth_evidence.absorb_log(claim, log_e)?;
transports.push(best_effort_transport(input.decoder_grid, atom, c0, c1));
jumps.push(LambdaJump {
step,
from_ckpt: input.checkpoint_ids[c0].clone(),
to_ckpt: input.checkpoint_ids[c1].clone(),
delta_lambda,
se,
z,
log_e,
born: false, });
}
let cert = birth_evidence
.certify(input.birth_alpha)
.map_err(|error| error.to_string())?;
for (jump, entry) in jumps.iter_mut().zip(cert.entries.iter()) {
jump.born = entry.confirmed && jump.delta_lambda > 0.0;
}
atoms.push(AtomLambdaTrajectory {
atom_name,
lambda,
rank_soft,
mp_reconstruction_rank,
production_chargeable_rank,
jumps,
birth_evidence,
transports,
});
}
let peak_log_e: Vec<f64> = atoms
.iter()
.map(AtomLambdaTrajectory::peak_birth_log_e)
.collect::<Result<Vec<_>, _>>()?;
let cross_atom_born =
e_benjamini_hochberg(&peak_log_e, input.birth_alpha).map_err(|error| error.to_string())?;
Ok(WbicDynamicsReport {
atoms,
cross_atom_born,
})
}
fn atom_learning_spectrum(
curve: ArrayView2<'_, f64>,
r_floor: f64,
) -> Result<ReconSpectrum, String> {
let (n_grid, ambient) = curve.dim();
let gram = Array2::<f64>::eye(n_grid);
let decoder = curve.to_owned();
recon_spectrum(
&gram,
&decoder,
n_grid as f64,
ambient as f64,
r_floor,
0.0,
None,
)?
.with_audit_basis_edf(1.0)
}
fn lambda_jackknife_se(curve: ArrayView2<'_, f64>, r_floor: f64) -> f64 {
let (n_grid, ambient) = curve.dim();
if n_grid < 3 {
return 0.0;
}
let mut leave_one = Vec::with_capacity(n_grid);
for drop in 0..n_grid {
let mut sub = Array2::<f64>::zeros((n_grid - 1, ambient));
let mut r = 0usize;
for i in 0..n_grid {
if i == drop {
continue;
}
for j in 0..ambient {
sub[[r, j]] = curve[[i, j]];
}
r += 1;
}
match atom_learning_spectrum(sub.view(), r_floor) {
Ok(spec) => leave_one.push(spec.learning_coefficient()),
Err(_) => return 0.0,
}
}
let m = leave_one.len() as f64;
let mean = leave_one.iter().sum::<f64>() / m;
let ss = leave_one
.iter()
.map(|&l| (l - mean) * (l - mean))
.sum::<f64>();
((m - 1.0) / m * ss).max(0.0).sqrt()
}
fn best_effort_transport(
grid: ArrayView4<'_, f64>,
atom: usize,
c0: usize,
c1: usize,
) -> Option<LayerTransportReport> {
let coords_from = grid.slice(ndarray::s![c0, atom, .., 0]).to_owned();
let coords_to = grid.slice(ndarray::s![c1, atom, .., 0]).to_owned();
let mut lo = f64::INFINITY;
let mut hi = f64::NEG_INFINITY;
for &v in coords_from.iter().chain(coords_to.iter()) {
lo = lo.min(v);
hi = hi.max(v);
}
if !(lo.is_finite() && hi.is_finite()) {
return None;
}
let (lo, hi) = if hi <= lo {
(lo - 0.5, lo + 0.5)
} else {
let pad = (hi - lo) * 1e-6;
(lo - pad, hi + pad)
};
let topology = ChartTopology::Interval { lo, hi };
fit_layer_transport(
c0,
c1,
coords_from.view(),
coords_to.view(),
topology,
topology,
)
.ok()
}
fn no_jump_log_e_value(z: f64) -> Result<f64, String> {
if !z.is_finite() {
return Err(format!(
"no-jump evidence requires a finite studentized jump; got {z}"
));
}
let log_p = std::f64::consts::LN_2 + normal_logsf(z.abs());
if !(log_p.is_finite() && log_p <= 0.0) {
return Err(format!(
"two-sided normal log-tail is not representable for studentized jump {z}: log_p={log_p}"
));
}
let log_e = -std::f64::consts::LN_2 - 0.5 * log_p;
if log_e.is_finite() {
Ok(log_e)
} else {
Err(format!(
"no-jump calibrated log-e is not representable for studentized jump {z}: {log_e}"
))
}
}
pub fn render_lambda_dynamics(report: &WbicDynamicsReport, checkpoint_ids: &[String]) -> String {
let mut out = String::new();
out.push_str("atom ");
for id in checkpoint_ids {
out.push_str(&format!("{id:>8}"));
}
out.push_str(" birth@step\n");
for (idx, atom) in report.atoms.iter().enumerate() {
out.push_str(&format!("{:<23}", atom.atom_name));
for l in &atom.lambda {
out.push_str(&format!("{l:>8.3}"));
}
let birth = atom
.jumps
.iter()
.find(|j| j.born)
.map(|j| format!(" {}→{}", j.from_ckpt, j.to_ckpt))
.unwrap_or_else(|| " —".to_string());
let cross = if report.cross_atom_born.contains(&idx) {
" [born]"
} else {
""
};
out.push_str(&format!("{birth}{cross}\n"));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array4;
#[test]
fn no_jump_log_e_remains_representable_beyond_probability_underflow() {
let moderate = no_jump_log_e_value(8.0).unwrap();
let deep = no_jump_log_e_value(40.0).unwrap();
assert!(moderate.is_finite() && deep.is_finite() && deep > moderate);
assert_eq!(deep, no_jump_log_e_value(-40.0).unwrap());
assert!(no_jump_log_e_value(f64::NAN).is_err());
assert!(no_jump_log_e_value(f64::INFINITY).is_err());
}
fn lcg(s: &mut u64) -> f64 {
*s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*s >> 11) as f64) / ((1u64 << 53) as f64)
}
fn lcg_normal(s: &mut u64) -> f64 {
let u1 = lcg(s).max(1e-12);
let u2 = lcg(s);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn olmo_birth_fixture(
n_ckpt: usize,
n_grid: usize,
ambient: usize,
birth_ckpt: usize,
) -> Array4<f64> {
let mut s = 0xB1D_7000_u64;
let mut grid = Array4::<f64>::zeros((n_ckpt, 3, n_grid, ambient));
let noise = 0.02_f64;
let signal = 1.0_f64;
for c in 0..n_ckpt {
for g in 0..n_grid {
let t = g as f64 / (n_grid - 1) as f64;
let a = std::f64::consts::TAU * t;
for comp in 0..ambient {
grid[[c, 0, g, comp]] = if comp == 0 {
signal * a.cos()
} else {
noise * lcg_normal(&mut s)
};
grid[[c, 1, g, comp]] = if c >= birth_ckpt && comp == 0 {
signal * a.cos()
} else {
noise * lcg_normal(&mut s)
};
grid[[c, 2, g, comp]] = noise * lcg_normal(&mut s);
}
}
}
grid
}
fn run(grid: &Array4<f64>) -> (WbicDynamicsReport, Vec<String>) {
let (n_ckpt, _n_atoms, _n_grid, _amb) = grid.dim();
let ckpt_ids: Vec<String> = (0..n_ckpt).map(|c| format!("step{c}")).collect();
let atom_names = vec![
"stable-strong".to_string(),
"born".to_string(),
"null".to_string(),
];
let input = WbicDynamicsInput {
decoder_grid: grid.view(),
checkpoint_ids: &ckpt_ids,
atom_names: &atom_names,
r_floor: 0.02_f64 * 0.02,
birth_alpha: 0.05,
};
(wbic_lambda_dynamics(&input).expect("dynamics"), ckpt_ids)
}
#[test]
fn detects_feature_birth_as_lambda_jump() {
let n_ckpt = 6;
let n_grid = 33; let ambient = 8;
let birth_ckpt = 3;
let grid = olmo_birth_fixture(n_ckpt, n_grid, ambient, birth_ckpt);
let (report, ckpt_ids) = run(&grid);
eprintln!("\n{}", render_lambda_dynamics(&report, &ckpt_ids));
assert_eq!(report.atoms.len(), 3);
let strong = &report.atoms[0];
let born = &report.atoms[1];
let null = &report.atoms[2];
for a in &report.atoms {
assert_eq!(a.lambda.len(), n_ckpt);
assert_eq!(a.jumps.len(), n_ckpt - 1);
}
let birth_jump = born.jumps[birth_ckpt - 1].delta_lambda;
assert!(
birth_jump > 0.25,
"born atom must show a large positive λ jump at birth, got {birth_jump}"
);
let max_other_jump = born
.jumps
.iter()
.enumerate()
.filter(|(i, _)| *i != birth_ckpt - 1)
.map(|(_, j)| j.delta_lambda.abs())
.fold(0.0_f64, f64::max);
assert!(
birth_jump > max_other_jump + 0.1,
"the birth jump ({birth_jump}) must dominate every other step jump ({max_other_jump})"
);
let born_steps: Vec<&LambdaJump> = born.jumps.iter().filter(|j| j.born).collect();
assert_eq!(
born_steps.len(),
1,
"born atom must have exactly one confirmed birth step, got {}",
born_steps.len()
);
assert_eq!(
born_steps[0].step,
birth_ckpt - 1,
"birth must be detected at the birth checkpoint transition"
);
assert!(
born_steps[0].delta_lambda > 0.0,
"a birth is a positive λ jump"
);
assert!(
strong.jumps.iter().all(|j| !j.born),
"stable-strong atom must never be flagged born (λ flat): {:?}",
strong.lambda
);
assert!(
null.jumps.iter().all(|j| !j.born),
"null atom must never be flagged born (λ ~ 0): {:?}",
null.lambda
);
assert!(
report.cross_atom_born.contains(&1),
"cross-atom certificate must confirm the born atom, got {:?}",
report.cross_atom_born
);
assert!(
!report.cross_atom_born.contains(&0) && !report.cross_atom_born.contains(&2),
"cross-atom certificate must not confirm the stable or null atoms, got {:?}",
report.cross_atom_born
);
}
#[test]
fn lambda_orders_strong_above_null_and_born_crosses() {
let grid = olmo_birth_fixture(6, 33, 8, 3);
let (report, _ids) = run(&grid);
let strong = &report.atoms[0];
let born = &report.atoms[1];
let null = &report.atoms[2];
for c in 0..strong.lambda.len() {
assert!(
strong.lambda[c] > null.lambda[c] + 0.15,
"strong λ ({}) must exceed null λ ({}) at checkpoint {c}",
strong.lambda[c],
null.lambda[c]
);
}
assert!(born.lambda[0] < strong.lambda[0]);
assert!(*born.lambda.last().unwrap() > *null.lambda.last().unwrap() + 0.15);
}
#[test]
fn born_out_evidences_null() {
let grid = olmo_birth_fixture(6, 33, 8, 3);
let (report, _ids) = run(&grid);
let total_log_e = |a: &AtomLambdaTrajectory| -> f64 {
a.birth_evidence
.certify(0.05)
.unwrap()
.entries
.iter()
.map(|e| e.log_e)
.sum()
};
let born_log_e = total_log_e(&report.atoms[1]);
let null_log_e = total_log_e(&report.atoms[2]);
assert!(
born_log_e > null_log_e,
"born change-evidence {born_log_e} must exceed null {null_log_e}"
);
assert!(
born_log_e > 0.0,
"a real birth must accumulate positive log-evidence, got {born_log_e}"
);
}
#[test]
fn rejects_single_checkpoint_and_bad_alpha() {
let grid = Array4::<f64>::zeros((1, 2, 8, 3));
let ids = vec!["only".to_string()];
let names = vec!["a".to_string(), "b".to_string()];
let bad = WbicDynamicsInput {
decoder_grid: grid.view(),
checkpoint_ids: &ids,
atom_names: &names,
r_floor: 0.01,
birth_alpha: 0.05,
};
assert!(wbic_lambda_dynamics(&bad).is_err());
let grid2 = Array4::<f64>::zeros((3, 2, 8, 3));
let ids2 = vec!["a".to_string(), "b".to_string(), "c".to_string()];
let bad_alpha = WbicDynamicsInput {
decoder_grid: grid2.view(),
checkpoint_ids: &ids2,
atom_names: &names,
r_floor: 0.01,
birth_alpha: 1.5,
};
assert!(wbic_lambda_dynamics(&bad_alpha).is_err());
}
}