use ndarray::{Array2, ArrayView2};
use crate::routability::{RoutabilityFloor, minimum_routable_energy};
#[derive(Clone, Debug)]
pub struct StratumLocalPick {
pub rows: Vec<usize>,
pub masked_residual: Array2<f64>,
pub local_fraction: f64,
pub pooled_fraction: f64,
pub min_routable_energy: f64,
pub stratum_rank: usize,
pub effective_sample_size: f64,
pub min_effective_rows: f64,
}
pub fn dominant_energy_fraction(residual: ArrayView2<'_, f64>, rows: &[usize]) -> f64 {
let p = residual.ncols();
if p == 0 || rows.is_empty() {
return 0.0;
}
let mut total = 0.0_f64;
for &i in rows {
for &v in residual.row(i) {
total += v * v;
}
}
if total <= 0.0 {
return 0.0;
}
let lambda_max = power_iteration_top_eigenvalue(residual, rows, p);
(lambda_max / total).clamp(0.0, 1.0)
}
fn power_iteration_top_eigenvalue(residual: ArrayView2<'_, f64>, rows: &[usize], p: usize) -> f64 {
let mut v = vec![1.0_f64 / (p as f64).sqrt(); p];
let mut lambda = 0.0_f64;
for _ in 0..24 {
let mut u = vec![0.0_f64; p];
for &i in rows {
let row = residual.row(i);
let mut wi = 0.0_f64;
for j in 0..p {
wi += row[j] * v[j];
}
for j in 0..p {
u[j] += wi * row[j];
}
}
let norm = u.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm <= 0.0 {
return 0.0;
}
lambda = norm;
for j in 0..p {
v[j] = u[j] / norm;
}
}
lambda
}
pub fn min_effective_rows_for_birth(p: usize, min_routable: f64) -> f64 {
if p == 0 || !(min_routable > 0.0) {
return f64::INFINITY;
}
let pf = p as f64;
let rho_p = min_routable * pf;
if rho_p <= 1.0 {
return f64::INFINITY;
}
pf * (1.0 + rho_p.sqrt()).powi(2) / (rho_p - 1.0).powi(2)
}
fn effective_sample_size(residual: ArrayView2<'_, f64>, rows: &[usize]) -> f64 {
let mut sum = 0.0_f64;
let mut sumsq = 0.0_f64;
for &i in rows {
let e: f64 = residual.row(i).iter().map(|&v| v * v).sum();
sum += e;
sumsq += e * e;
}
if sumsq <= 0.0 {
return 0.0;
}
sum * sum / sumsq
}
pub fn stratum_local_birth_residual(
residual: ArrayView2<'_, f64>,
floor: &RoutabilityFloor,
) -> Option<StratumLocalPick> {
let (n, p) = residual.dim();
if n == 0 || p == 0 {
return None;
}
let min_routable = minimum_routable_energy(floor);
let energies: Vec<f64> = (0..n)
.map(|i| residual.row(i).iter().map(|&v| v * v).sum())
.collect();
let all_rows: Vec<usize> = (0..n).collect();
let pooled_fraction = dominant_energy_fraction(residual, &all_rows);
let mut strata = crate::corpus::stratify_row_energies(&energies);
strata.sort_by(|a, b| {
b.mean_energy
.partial_cmp(&a.mean_energy)
.unwrap_or(std::cmp::Ordering::Equal)
});
let min_effective_rows = min_effective_rows_for_birth(p, min_routable);
for (rank, stratum) in strata.iter().enumerate() {
let local = dominant_energy_fraction(residual, &stratum.rows);
if local < min_routable {
continue;
}
let ess = effective_sample_size(residual, &stratum.rows);
if ess < min_effective_rows {
continue;
}
let mut masked = Array2::<f64>::zeros((n, p));
for &i in &stratum.rows {
masked.row_mut(i).assign(&residual.row(i));
}
return Some(StratumLocalPick {
rows: stratum.rows.clone(),
masked_residual: masked,
local_fraction: local,
pooled_fraction,
min_routable_energy: min_routable,
stratum_rank: rank,
effective_sample_size: ess,
min_effective_rows,
});
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use crate::routability::routability_floor;
use ndarray::Array2;
fn gauss(counter: &mut u64) -> f64 {
*counter = counter.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = *counter;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
let a = ((z ^ (z >> 31)) >> 11) as f64 / (1u64 << 53) as f64;
*counter = counter.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut w = *counter;
w = (w ^ (w >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
let b = ((w ^ (w >> 27)) >> 11) as f64 / (1u64 << 53) as f64;
(-2.0 * a.max(1e-12).ln()).sqrt() * (std::f64::consts::TAU * b).cos()
}
#[test]
fn stratum_local_admits_a_planted_signal_pooled_rejects() {
let n = 4000usize;
let p = 256usize;
let k_router = 32_000usize; let n_signal = 40usize; let signal_amp = 6.0; let noise = 1.0;
let mut ctr = 1u64;
let mut r = Array2::<f64>::zeros((n, p));
let mut dir = vec![0.0_f64; p];
for d in dir.iter_mut() {
*d = gauss(&mut ctr);
}
let dn = dir.iter().map(|x| x * x).sum::<f64>().sqrt();
for d in dir.iter_mut() {
*d /= dn;
}
for i in 0..n {
for j in 0..p {
r[[i, j]] = noise * gauss(&mut ctr);
}
}
for i in 0..n_signal {
let a = signal_amp * gauss(&mut ctr);
for j in 0..p {
r[[i, j]] += a * dir[j];
}
}
let floor = routability_floor(p, k_router, 1, 1.0);
let min_routable = minimum_routable_energy(&floor);
let all: Vec<usize> = (0..n).collect();
let pooled = dominant_energy_fraction(r.view(), &all);
assert!(
pooled < min_routable,
"pooled fraction {pooled} must be BELOW the floor {min_routable} \
(the pooled birth would be rejected by the router)"
);
let pick = stratum_local_birth_residual(r.view(), &floor)
.expect("stratum-local screen must admit the planted signal");
assert!(
pick.local_fraction >= min_routable,
"local fraction {} must CLEAR the floor {min_routable}",
pick.local_fraction
);
assert!(
pick.local_fraction > pick.pooled_fraction,
"local {} must exceed pooled {}",
pick.local_fraction,
pick.pooled_fraction
);
assert!(
pick.effective_sample_size >= pick.min_effective_rows,
"admitted stratum ESS {} must clear the floor {}",
pick.effective_sample_size,
pick.min_effective_rows
);
let planted: std::collections::BTreeSet<usize> = (0..n_signal).collect();
let overlap = pick.rows.iter().filter(|i| planted.contains(i)).count();
assert!(
overlap as f64 >= 0.8 * n_signal as f64,
"the picked stratum must hold the planted rows: {overlap}/{n_signal}"
);
for i in 0..n {
let kept = pick.rows.binary_search(&i).is_ok();
let row_energy: f64 = pick.masked_residual.row(i).iter().map(|v| v * v).sum();
if kept {
let orig: f64 = r.row(i).iter().map(|v| v * v).sum();
assert!((row_energy - orig).abs() < 1e-9, "kept row {i} preserved");
} else {
assert_eq!(row_energy, 0.0, "non-stratum row {i} must be zeroed");
}
}
}
#[test]
fn easy_floor_declines_no_stratification_needed() {
let n = 500usize;
let p = 32usize;
let mut ctr = 7u64;
let mut r = Array2::<f64>::zeros((n, p));
let mut dir = vec![0.0_f64; p];
for d in dir.iter_mut() {
*d = gauss(&mut ctr);
}
let dn = dir.iter().map(|x| x * x).sum::<f64>().sqrt();
for d in dir.iter_mut() {
*d /= dn;
}
for i in 0..n {
let a = 5.0 * gauss(&mut ctr);
for j in 0..p {
r[[i, j]] = a * dir[j] + 0.01 * gauss(&mut ctr);
}
}
let floor = routability_floor(p, 8, 1, 1.0);
let all: Vec<usize> = (0..n).collect();
let pooled = dominant_energy_fraction(r.view(), &all);
let min_routable = minimum_routable_energy(&floor);
assert!(
pooled >= min_routable,
"an easy floor with a dominant signal must route pooled: {pooled} ≥ {min_routable}"
);
}
#[test]
fn ess_floor_closed_form_and_kish_sample_size() {
let p = 256usize;
let floor = routability_floor(p, 32_000, 1, 1.0);
let rho = minimum_routable_energy(&floor);
let rho_p = rho * p as f64;
assert!(rho_p > 1.0, "frontier floor must admit a finite m_min");
let expected = p as f64 * (1.0 + rho_p.sqrt()).powi(2) / (rho_p - 1.0).powi(2);
let got = min_effective_rows_for_birth(p, rho);
assert!((got - expected).abs() < 1e-9, "m_min {got} vs {expected}");
assert!(got > 1.0, "a nontrivial router width must demand > 1 row");
assert!(min_effective_rows_for_birth(p, 0.5 / p as f64).is_infinite());
assert!(min_effective_rows_for_birth(0, 0.5).is_infinite());
let mut r = Array2::<f64>::zeros((5, 4));
for j in 0..4 {
r[[0, j]] = 1.0; r[[1, j]] = 1.0; r[[2, j]] = 1.0;
r[[3, j]] = 1.0;
}
r[[4, 0]] = 100.0; assert!((effective_sample_size(r.view(), &[0]) - 1.0).abs() < 1e-12);
assert!((effective_sample_size(r.view(), &[0, 1, 2, 3]) - 4.0).abs() < 1e-12);
assert!(
effective_sample_size(r.view(), &[0, 1, 2, 3, 4]) < 1.01,
"a one-row-dominated selection must have ESS ≈ 1"
);
assert_eq!(effective_sample_size(r.view(), &[]), 0.0);
}
#[test]
fn one_row_stratum_proposes_nothing_adequate_stratum_proposes() {
let n = 1000usize;
let p = 64usize;
let k_router = 32_000usize; let mut ctr = 42u64;
let mut r = Array2::<f64>::zeros((n + 1, p));
for i in 0..n {
let mut e = 0.0_f64;
for j in 0..p {
let v = gauss(&mut ctr);
r[[i, j]] = v;
e += v * v;
}
let s = e.sqrt();
for j in 0..p {
r[[i, j]] /= s;
}
}
let mut e = 0.0_f64;
for j in 0..p {
let v = gauss(&mut ctr);
r[[n, j]] = v;
e += v * v;
}
let s = e.sqrt();
for j in 0..p {
r[[n, j]] = r[[n, j]] / s * 1000.0_f64.sqrt();
}
let floor = routability_floor(p, k_router, 1, 1.0);
let min_routable = minimum_routable_energy(&floor);
assert!(dominant_energy_fraction(r.view(), &[n]) >= min_routable);
assert!(
effective_sample_size(r.view(), &[n]) < min_effective_rows_for_birth(p, min_routable)
);
assert!(
stratum_local_birth_residual(r.view(), &floor).is_none(),
"a 1-row noise outlier must not seed a birth"
);
let n_signal = 50usize;
let mut dir = vec![0.0_f64; p];
for d in dir.iter_mut() {
*d = gauss(&mut ctr);
}
let dn = dir.iter().map(|x| x * x).sum::<f64>().sqrt();
for d in dir.iter_mut() {
*d /= dn;
}
for i in 0..n_signal {
let a = 5.0;
for j in 0..p {
r[[i, j]] = a * dir[j] + 0.05 * gauss(&mut ctr);
}
}
let pick = stratum_local_birth_residual(r.view(), &floor)
.expect("an adequately-sized planted stratum must propose a birth");
assert!(pick.local_fraction >= pick.min_routable_energy);
assert!(
pick.effective_sample_size >= pick.min_effective_rows,
"the proposed stratum must clear the ESS floor: {} < {}",
pick.effective_sample_size,
pick.min_effective_rows
);
assert!(
pick.effective_sample_size > 1.5,
"the proposed stratum must be a genuine multi-row band, not a lone row"
);
}
}