use gam_terms::inference::structure_evidence::e_benjamini_hochberg;
use ndarray::Array2;
use super::block_chart::jacobi_eigh;
#[derive(Clone, Debug)]
pub struct FdrCertificate {
pub alpha: f64,
pub log_e: Vec<f64>,
pub rejected: Vec<usize>,
}
pub fn family_fdr_certificate(log_e: Vec<f64>, alpha: f64) -> FdrCertificate {
let rejected = e_benjamini_hochberg(&log_e, alpha);
FdrCertificate {
alpha,
log_e,
rejected,
}
}
pub fn crossfit_ui_log_evalue<A>(
n: usize,
folds: usize,
mut fit_alt: impl FnMut(&[usize]) -> Result<Option<A>, String>,
mut alt_loglik: impl FnMut(&A, &[usize]) -> Result<f64, String>,
mut null_sup_loglik: impl FnMut(&[usize]) -> Result<f64, String>,
) -> Result<f64, String> {
if n < 2 {
return Ok(f64::NEG_INFINITY);
}
let folds = folds.max(2).min(n);
let mut fold_log_e: Vec<f64> = Vec::with_capacity(folds);
for f in 0..folds {
let mut train = Vec::new();
let mut eval = Vec::new();
for i in 0..n {
if i % folds == f {
eval.push(i);
} else {
train.push(i);
}
}
if train.is_empty() || eval.is_empty() {
continue;
}
let log_e_f = match fit_alt(&train)? {
Some(alt) => {
let alt_ll = alt_loglik(&alt, &eval)?;
let null_ll = null_sup_loglik(&eval)?;
alt_ll - null_ll
}
None => f64::NEG_INFINITY,
};
fold_log_e.push(log_e_f);
}
if fold_log_e.is_empty() {
return Ok(f64::NEG_INFINITY);
}
Ok(logsumexp(&fold_log_e) - (fold_log_e.len() as f64).ln())
}
pub fn shell_vs_ring_log_evalue(
coords: &Array2<f64>,
folds: usize,
ridge: f64,
) -> Result<f64, String> {
let n = coords.nrows();
let q = coords.ncols();
if q < 2 {
return Ok(f64::NEG_INFINITY);
}
crossfit_ui_log_evalue(
n,
folds,
|train| {
if train.len() < 3 {
return Ok(None);
}
Ok(Some(fit_ring(coords, train, ridge)))
},
|ring, eval| Ok(ring_loglik(coords, eval, ring)),
|eval| {
let null = fit_ppca1(coords, eval, ridge)?;
Ok(ppca1_loglik(coords, eval, &null))
},
)
}
struct Ppca1 {
mean: Vec<f64>,
axis: Vec<f64>,
sigma2: f64,
lambda: f64,
logdet: f64,
}
fn fit_ppca1(coords: &Array2<f64>, rows: &[usize], ridge: f64) -> Result<Ppca1, String> {
let n = rows.len();
let q = coords.ncols();
if n == 0 || q == 0 {
return Err("ppca1 fit: empty rows or columns".to_string());
}
let mut mean = vec![0.0f64; q];
for &i in rows {
for j in 0..q {
mean[j] += coords[[i, j]];
}
}
for m in &mut mean {
*m /= n as f64;
}
let mut cov = vec![0.0f64; q * q];
for &i in rows {
for a in 0..q {
let va = coords[[i, a]] - mean[a];
for b in 0..q {
cov[a * q + b] += va * (coords[[i, b]] - mean[b]);
}
}
}
for v in &mut cov {
*v /= n as f64;
}
let (vals, vecs) = jacobi_eigh(cov, q)?;
let mut order: Vec<usize> = (0..q).collect();
order.sort_by(|&a, &b| {
vals[b]
.partial_cmp(&vals[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let trace: f64 = vals.iter().sum();
let floor = (ridge * trace / q as f64).max(1.0e-300);
let top = order[0];
let ell1 = vals[top].max(floor);
let axis: Vec<f64> = (0..q).map(|j| vecs[j * q + top]).collect();
let sigma2 = if q > 1 {
let rest: f64 = order.iter().skip(1).map(|&k| vals[k].max(0.0)).sum();
(rest / (q - 1) as f64).max(floor)
} else {
ell1
};
let lambda = (ell1 - sigma2).max(0.0);
let logdet = (q as f64 - 1.0) * sigma2.ln() + (sigma2 + lambda).ln();
Ok(Ppca1 {
mean,
axis,
sigma2,
lambda,
logdet,
})
}
fn ppca1_loglik(coords: &Array2<f64>, rows: &[usize], m: &Ppca1) -> f64 {
let q = coords.ncols();
let log_2pi = (2.0 * std::f64::consts::PI).ln();
let sm = m.lambda / (m.sigma2 + m.lambda);
let mut total = 0.0;
for &i in rows {
let mut norm2 = 0.0;
let mut proj = 0.0;
for j in 0..q {
let r = coords[[i, j]] - m.mean[j];
norm2 += r * r;
proj += r * m.axis[j];
}
let quad = (norm2 - sm * proj * proj) / m.sigma2;
total += -0.5 * (q as f64 * log_2pi + m.logdet + quad);
}
total
}
struct Ring {
mean: Vec<f64>,
e1: Vec<f64>,
e2: Vec<f64>,
radius: f64,
sigma2: f64,
sigma_perp2: Option<f64>,
}
fn fit_ring(coords: &Array2<f64>, rows: &[usize], ridge: f64) -> Ring {
let n = rows.len();
let q = coords.ncols();
let mut mean = vec![0.0f64; q];
for &i in rows {
for j in 0..q {
mean[j] += coords[[i, j]];
}
}
for m in &mut mean {
*m /= n as f64;
}
let mut cov = vec![0.0f64; q * q];
for &i in rows {
for a in 0..q {
let va = coords[[i, a]] - mean[a];
for b in 0..q {
cov[a * q + b] += va * (coords[[i, b]] - mean[b]);
}
}
}
for v in &mut cov {
*v /= n as f64;
}
let (vals, vecs) = match jacobi_eigh(cov, q) {
Ok(pair) => pair,
Err(_) => (vec![0.0; q], identity_flat(q)),
};
let mut order: Vec<usize> = (0..q).collect();
order.sort_by(|&a, &b| {
vals[b]
.partial_cmp(&vals[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let k1 = order[0];
let k2 = order[1];
let e1: Vec<f64> = (0..q).map(|j| vecs[j * q + k1]).collect();
let e2: Vec<f64> = (0..q).map(|j| vecs[j * q + k2]).collect();
let trace: f64 = vals.iter().sum();
let floor = (ridge * trace / q as f64).max(1.0e-12);
let mut radius = 0.0;
let mut radial_var = 0.0;
let mut perp_ss = 0.0;
let mut rhos = Vec::with_capacity(n);
for &i in rows {
let mut z1 = 0.0;
let mut z2 = 0.0;
let mut norm2 = 0.0;
for j in 0..q {
let r = coords[[i, j]] - mean[j];
z1 += r * e1[j];
z2 += r * e2[j];
norm2 += r * r;
}
let rho = (z1 * z1 + z2 * z2).sqrt();
rhos.push(rho);
radius += rho;
perp_ss += (norm2 - z1 * z1 - z2 * z2).max(0.0);
}
radius /= n as f64;
for &rho in &rhos {
radial_var += (rho - radius) * (rho - radius);
}
let sigma2 = (radial_var / n as f64).max(floor);
let sigma_perp2 = if q > 2 {
Some((perp_ss / (n as f64 * (q - 2) as f64)).max(floor))
} else {
None
};
Ring {
mean,
e1,
e2,
radius,
sigma2,
sigma_perp2,
}
}
fn ring_loglik(coords: &Array2<f64>, rows: &[usize], m: &Ring) -> f64 {
let q = coords.ncols();
let log_2pi = (2.0 * std::f64::consts::PI).ln();
let mut total = 0.0;
for &i in rows {
let mut z1 = 0.0;
let mut z2 = 0.0;
let mut norm2 = 0.0;
for j in 0..q {
let r = coords[[i, j]] - m.mean[j];
z1 += r * m.e1[j];
z2 += r * m.e2[j];
norm2 += r * r;
}
let rho = (z1 * z1 + z2 * z2).sqrt();
let plane = -(log_2pi + m.sigma2.ln())
- (rho * rho + m.radius * m.radius) / (2.0 * m.sigma2)
+ ln_i0(m.radius * rho / m.sigma2);
total += plane;
if let Some(sp) = m.sigma_perp2 {
let perp_ss = (norm2 - z1 * z1 - z2 * z2).max(0.0);
let k = (q - 2) as f64;
total += -0.5 * k * (log_2pi + sp.ln()) - perp_ss / (2.0 * sp);
}
}
total
}
fn identity_flat(q: usize) -> Vec<f64> {
let mut v = vec![0.0f64; q * q];
for i in 0..q {
v[i * q + i] = 1.0;
}
v
}
fn logsumexp(xs: &[f64]) -> f64 {
let m = xs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
if !m.is_finite() {
return m;
}
let s: f64 = xs.iter().map(|&x| (x - m).exp()).sum();
m + s.ln()
}
fn ln_i0(x: f64) -> f64 {
let ax = x.abs();
if ax < 3.75 {
let t = ax / 3.75;
let t2 = t * t;
(1.0 + t2
* (3.5156229
+ t2 * (3.0899424
+ t2 * (1.2067492 + t2 * (0.2659732 + t2 * (0.0360768 + t2 * 0.0045813))))))
.ln()
} else {
let y = 3.75 / ax;
let poly = 0.39894228
+ y * (0.01328592
+ y * (0.00225319
+ y * (-0.00157565
+ y * (0.00916281
+ y * (-0.02057706
+ y * (0.02635537 + y * (-0.01647633 + y * 0.00392377)))))));
ax - 0.5 * ax.ln() + poly.ln()
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array2;
struct Rng(u64);
impl Rng {
fn u01(&mut self) -> f64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
((x >> 11) as f64) / ((1u64 << 53) as f64)
}
fn normal(&mut self) -> f64 {
let u1 = self.u01().max(1.0e-300);
let u2 = self.u01();
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
}
}
fn draw_line(rng: &mut Rng, n: usize, q: usize, spread: f64, noise: f64) -> Array2<f64> {
let axis: Vec<f64> = {
let mut a: Vec<f64> = (0..q).map(|_| rng.normal()).collect();
let nrm = (a.iter().map(|x| x * x).sum::<f64>()).sqrt().max(1.0e-12);
for x in &mut a {
*x /= nrm;
}
a
};
let mut out = Array2::<f64>::zeros((n, q));
for i in 0..n {
let t = spread * rng.normal();
for j in 0..q {
out[[i, j]] = t * axis[j] + noise * rng.normal();
}
}
out
}
fn draw_ring(rng: &mut Rng, n: usize, q: usize, radius: f64, noise: f64) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((n, q));
for i in 0..n {
let theta = 2.0 * std::f64::consts::PI * rng.u01();
out[[i, 0]] = radius * theta.cos() + noise * rng.normal();
out[[i, 1]] = radius * theta.sin() + noise * rng.normal();
for j in 2..q {
out[[i, j]] = noise * rng.normal();
}
}
out
}
#[test]
fn evalue_null_expectation_leq_one() {
let mut rng = Rng(0x1234_5678_9abc_def0);
let trials = 400;
let mut sum_e = 0.0;
for _ in 0..trials {
let coords = draw_line(&mut rng, 80, 3, 2.0, 0.3);
let log_e = shell_vs_ring_log_evalue(&coords, 2, 1.0e-6).unwrap();
sum_e += log_e.exp();
}
let mean_e = sum_e / trials as f64;
assert!(
mean_e <= 1.25,
"UI e-value must satisfy E_H0[E] <= 1 (MC margin); got mean E = {mean_e}"
);
}
#[test]
fn null_battery_rank1_shell_fdr_controlled() {
let alpha = 0.1;
let mut rng = Rng(0xdead_beef_0bad_f00d);
let sims = 300;
let family_size = 12;
let mut false_discovery_sims = 0;
for _ in 0..sims {
let log_e: Vec<f64> = (0..family_size)
.map(|_| {
let coords = draw_line(&mut rng, 60, 3, 2.0, 0.3);
shell_vs_ring_log_evalue(&coords, 2, 1.0e-6).unwrap()
})
.collect();
let cert = family_fdr_certificate(log_e, alpha);
if !cert.rejected.is_empty() {
false_discovery_sims += 1;
}
}
let empirical_fdr = false_discovery_sims as f64 / sims as f64;
assert!(
empirical_fdr <= alpha + 0.03,
"empirical FDR {empirical_fdr} must be <= alpha {alpha} (MC margin)"
);
}
#[test]
fn power_under_curved_signal() {
let alpha = 0.1;
let mut rng = Rng(0x0123_4567_89ab_cdef);
let sims = 60;
let mut discovered = 0;
for _ in 0..sims {
let mut log_e = Vec::new();
let ring = draw_ring(&mut rng, 240, 3, 3.0, 0.25);
log_e.push(shell_vs_ring_log_evalue(&ring, 2, 1.0e-6).unwrap());
for _ in 0..6 {
let line = draw_line(&mut rng, 240, 3, 2.0, 0.3);
log_e.push(shell_vs_ring_log_evalue(&line, 2, 1.0e-6).unwrap());
}
let cert = family_fdr_certificate(log_e, alpha);
if cert.rejected.contains(&0) {
discovered += 1;
}
}
let power = discovered as f64 / sims as f64;
assert!(
power > 0.5,
"the genuine ring should be discovered with real power; got {power}"
);
}
#[test]
fn degenerate_candidate_never_rejected() {
let coords =
Array2::<f64>::from_shape_vec((10, 1), (0..10).map(|x| x as f64).collect()).unwrap();
let log_e = shell_vs_ring_log_evalue(&coords, 2, 1.0e-6).unwrap();
assert!(
log_e.is_infinite() && log_e < 0.0,
"q<2 must give -inf log_e"
);
let cert = family_fdr_certificate(vec![log_e], 0.1);
assert!(cert.rejected.is_empty());
}
}