use ndarray::{Array2, ArrayView1};
use gam_linalg::utils::splitmix64_hash;
use super::shard_reader::CorpusRowSource;
pub trait RowResidualEnergy {
fn energy(&self, row: ArrayView1<f64>) -> f64;
}
#[derive(Debug, Clone)]
pub struct SpanResidualEnergy {
basis: Array2<f64>,
}
impl SpanResidualEnergy {
pub fn new(basis: Array2<f64>) -> Self {
Self { basis }
}
pub fn width(&self) -> usize {
self.basis.nrows()
}
}
impl RowResidualEnergy for SpanResidualEnergy {
#[inline]
fn energy(&self, row: ArrayView1<f64>) -> f64 {
let full: f64 = row.iter().map(|&v| v * v).sum();
if self.basis.ncols() == 0 {
return full;
}
let mut explained = 0.0_f64;
for col in self.basis.columns() {
let proj: f64 = col.iter().zip(row.iter()).map(|(&q, &x)| q * x).sum();
explained += proj * proj;
}
(full - explained).max(0.0)
}
}
const N_EXPONENT_BINS: usize = 1 << 11;
#[derive(Debug, Clone)]
struct EnergyExponentHistogram {
count: Vec<u64>,
sum: Vec<f64>,
sumsq: Vec<f64>,
total_rows: u64,
}
impl EnergyExponentHistogram {
fn new() -> Self {
Self {
count: vec![0; N_EXPONENT_BINS],
sum: vec![0.0; N_EXPONENT_BINS],
sumsq: vec![0.0; N_EXPONENT_BINS],
total_rows: 0,
}
}
#[inline]
fn bin_of(energy: f64) -> usize {
if !energy.is_finite() || energy <= 0.0 {
return 0;
}
((energy.to_bits() >> 52) & 0x7ff) as usize
}
#[inline]
fn observe(&mut self, energy: f64) {
let e = if energy.is_finite() && energy > 0.0 {
energy
} else {
0.0
};
let bin = Self::bin_of(e);
self.count[bin] += 1;
self.sum[bin] += e;
self.sumsq[bin] += e * e;
self.total_rows += 1;
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Stratum {
pub exp_lo: usize,
pub exp_hi: usize,
pub n_rows: u64,
pub mean_energy: f64,
pub std_energy: f64,
pub pi: f64,
pub censused: bool,
}
#[derive(Debug, Clone)]
pub struct StratumDesign {
strata: Vec<Stratum>,
bin_to_stratum: Vec<usize>,
total_rows: u64,
budget: usize,
}
impl StratumDesign {
pub fn strata(&self) -> &[Stratum] {
&self.strata
}
pub fn total_rows(&self) -> u64 {
self.total_rows
}
pub fn budget(&self) -> usize {
self.budget
}
pub fn uniform_rate(&self) -> f64 {
if self.total_rows == 0 {
0.0
} else {
(self.budget as f64 / self.total_rows as f64).min(1.0)
}
}
fn rate_for_energy(&self, energy: f64) -> f64 {
let bin = EnergyExponentHistogram::bin_of(energy);
let s = self.bin_to_stratum[bin];
if s == usize::MAX {
self.uniform_rate()
} else {
self.strata[s].pi
}
}
}
fn sturges_stratum_cap(total_rows: u64) -> usize {
if total_rows <= 1 {
return 1;
}
(u64::BITS - total_rows.leading_zeros()) as usize
}
#[derive(Clone, Debug, PartialEq)]
pub struct RowStratum {
pub exp_lo: usize,
pub exp_hi: usize,
pub rows: Vec<usize>,
pub mean_energy: f64,
pub std_energy: f64,
}
pub fn stratify_row_energies(energies: &[f64]) -> Vec<RowStratum> {
let n = energies.len();
if n == 0 {
return Vec::new();
}
let mut bin_rows: std::collections::BTreeMap<usize, Vec<usize>> =
std::collections::BTreeMap::new();
for (i, &e) in energies.iter().enumerate() {
bin_rows
.entry(EnergyExponentHistogram::bin_of(e))
.or_default()
.push(i);
}
let clamp = |e: f64| if e.is_finite() && e > 0.0 { e } else { 0.0 };
let mut strata: Vec<RowStratum> = bin_rows
.into_iter()
.map(|(exp, rows)| {
let m = rows.len() as f64;
let sum: f64 = rows.iter().map(|&i| clamp(energies[i])).sum();
let mean = sum / m;
let var = (rows
.iter()
.map(|&i| clamp(energies[i]).powi(2))
.sum::<f64>()
/ m
- mean * mean)
.max(0.0);
RowStratum {
exp_lo: exp,
exp_hi: exp,
rows,
mean_energy: mean,
std_energy: var.sqrt(),
}
})
.collect();
let k_max = sturges_stratum_cap(n as u64);
while strata.len() > k_max && strata.len() >= 2 {
let mut merged = std::mem::take(&mut strata[0]);
let hi = std::mem::take(&mut strata[1]);
let na = merged.rows.len() as f64;
let nb = hi.rows.len() as f64;
let nt = na + nb;
let mean = (merged.mean_energy * na + hi.mean_energy * nb) / nt;
let sumsq_a = (merged.std_energy.powi(2) + merged.mean_energy.powi(2)) * na;
let sumsq_b = (hi.std_energy.powi(2) + hi.mean_energy.powi(2)) * nb;
let var = ((sumsq_a + sumsq_b) / nt - mean * mean).max(0.0);
merged.rows.extend(hi.rows);
merged.rows.sort_unstable();
merged.exp_hi = hi.exp_hi.max(merged.exp_hi);
merged.mean_energy = mean;
merged.std_energy = var.sqrt();
strata[1] = merged;
strata.remove(0);
}
strata
}
impl Default for RowStratum {
fn default() -> Self {
RowStratum {
exp_lo: 0,
exp_hi: 0,
rows: Vec::new(),
mean_energy: 0.0,
std_energy: 0.0,
}
}
}
fn build_strata(hist: &EnergyExponentHistogram) -> (Vec<Stratum>, Vec<usize>) {
let mut strata: Vec<Stratum> = Vec::new();
for (exp, &c) in hist.count.iter().enumerate() {
if c == 0 {
continue;
}
let n = c as f64;
let mean = hist.sum[exp] / n;
let var = (hist.sumsq[exp] / n - mean * mean).max(0.0);
strata.push(Stratum {
exp_lo: exp,
exp_hi: exp,
n_rows: c,
mean_energy: mean,
std_energy: var.sqrt(),
pi: 0.0,
censused: false,
});
}
let k_max = sturges_stratum_cap(hist.total_rows);
while strata.len() > k_max && strata.len() >= 2 {
let merged = merge_stratum(&strata[0], &strata[1]);
strata[1] = merged;
strata.remove(0);
}
let mut bin_to_stratum = vec![usize::MAX; N_EXPONENT_BINS];
for (idx, s) in strata.iter().enumerate() {
for exp in s.exp_lo..=s.exp_hi {
bin_to_stratum[exp] = idx;
}
}
(strata, bin_to_stratum)
}
fn merge_stratum(a: &Stratum, b: &Stratum) -> Stratum {
let na = a.n_rows as f64;
let nb = b.n_rows as f64;
let n = na + nb;
let sum = a.mean_energy * na + b.mean_energy * nb;
let mean = if n > 0.0 { sum / n } else { 0.0 };
let sumsq_a = (a.std_energy * a.std_energy + a.mean_energy * a.mean_energy) * na;
let sumsq_b = (b.std_energy * b.std_energy + b.mean_energy * b.mean_energy) * nb;
let var = if n > 0.0 {
((sumsq_a + sumsq_b) / n - mean * mean).max(0.0)
} else {
0.0
};
Stratum {
exp_lo: a.exp_lo.min(b.exp_lo),
exp_hi: a.exp_hi.max(b.exp_hi),
n_rows: a.n_rows + b.n_rows,
mean_energy: mean,
std_energy: var.sqrt(),
pi: 0.0,
censused: false,
}
}
fn allocate_rates(strata: &mut [Stratum], total_rows: u64, budget: usize) {
let k = strata.len();
if k == 0 || total_rows == 0 {
return;
}
let n_total = total_rows as f64;
let uniform_rate = (budget as f64 / n_total).min(1.0);
if budget as u64 >= total_rows {
for s in strata.iter_mut() {
s.pi = 1.0;
s.censused = true;
}
return;
}
for s in strata.iter_mut() {
s.censused = false;
}
loop {
let remaining: Vec<usize> = (0..k).filter(|&i| !strata[i].censused).collect();
if remaining.is_empty() {
break;
}
let censused_pop: u64 = strata.iter().filter(|s| s.censused).map(|s| s.n_rows).sum();
let budget_left = budget as f64 - censused_pop as f64;
if budget_left <= 0.0 {
break;
}
let share = budget_left / remaining.len() as f64;
let mut newly = false;
for &i in &remaining {
if (strata[i].n_rows as f64) <= share {
strata[i].censused = true;
newly = true;
}
}
if !newly {
break;
}
}
let censused_pop: u64 = strata.iter().filter(|s| s.censused).map(|s| s.n_rows).sum();
let budget_left = (budget as f64 - censused_pop as f64).max(0.0);
let neyman_mass: f64 = strata
.iter()
.filter(|s| !s.censused)
.map(|s| s.n_rows as f64 * s.std_energy)
.sum();
let big_pop: f64 = strata
.iter()
.filter(|s| !s.censused)
.map(|s| s.n_rows as f64)
.sum();
for s in strata.iter_mut() {
if s.censused {
s.pi = 1.0;
continue;
}
let n_h = s.n_rows as f64;
let target = if neyman_mass > 0.0 {
budget_left * (n_h * s.std_energy / neyman_mass) / n_h
} else if big_pop > 0.0 {
budget_left / big_pop
} else {
uniform_rate
};
s.pi = target.max(uniform_rate).min(1.0);
}
}
pub fn design_stratified_subsample(
source: &mut dyn CorpusRowSource,
energy: &dyn RowResidualEnergy,
budget: usize,
) -> Result<StratumDesign, String> {
let total_rows = source.total_rows();
let mut hist = EnergyExponentHistogram::new();
source.reset();
while let Some(batch) = source
.next_batch()
.map_err(|e| format!("design_stratified_subsample: shard read failed: {e}"))?
{
for k in 0..batch.rows.nrows() {
hist.observe(energy.energy(batch.rows.row(k)));
}
}
if hist.total_rows != total_rows {
return Err(format!(
"design_stratified_subsample: streamed {} rows but source declared {total_rows}",
hist.total_rows
));
}
let (mut strata, bin_to_stratum) = build_strata(&hist);
allocate_rates(&mut strata, total_rows, budget);
Ok(StratumDesign {
strata,
bin_to_stratum,
total_rows,
budget,
})
}
const STRATIFY_SALT: u64 = 0x5372_4154_1F9E_0B25;
#[inline]
fn row_included(row_id: u64, seed: u64, pi: f64) -> bool {
if pi >= 1.0 {
return true;
}
if pi <= 0.0 {
return false;
}
let h = splitmix64_hash(row_id ^ seed.wrapping_mul(STRATIFY_SALT) ^ STRATIFY_SALT);
let threshold = (pi * (u64::MAX as f64 + 1.0)) as u64;
h < threshold
}
#[derive(Debug, Clone)]
pub struct StratifiedCorpusTarget {
pub target: Array2<f64>,
pub row_ids: Vec<u64>,
pub likelihood_weights: Vec<f64>,
pub design: StratumDesign,
pub corpus_rows: u64,
}
impl StratifiedCorpusTarget {
pub fn len(&self) -> usize {
self.row_ids.len()
}
pub fn is_empty(&self) -> bool {
self.row_ids.is_empty()
}
pub fn estimated_corpus_rows(&self) -> f64 {
self.likelihood_weights.iter().sum()
}
}
pub fn collect_stratified_target(
source: &mut dyn CorpusRowSource,
energy: &dyn RowResidualEnergy,
budget: usize,
seed: u64,
) -> Result<StratifiedCorpusTarget, String> {
let design = design_stratified_subsample(source, energy, budget)?;
let corpus_rows = source.total_rows();
let p = source.width();
let mut rows_out: Vec<f64> = Vec::new();
let mut row_ids: Vec<u64> = Vec::new();
let mut likelihood_weights: Vec<f64> = Vec::new();
source.reset();
while let Some(batch) = source
.next_batch()
.map_err(|e| format!("collect_stratified_target: shard read failed: {e}"))?
{
for k in 0..batch.rows.nrows() {
let rid = batch.row_ids[k];
let e = energy.energy(batch.rows.row(k));
let pi = design.rate_for_energy(e);
if row_included(rid, seed, pi) {
rows_out.extend(batch.rows.row(k).iter().copied());
row_ids.push(rid);
likelihood_weights.push(1.0 / pi);
}
}
}
let n_sel = row_ids.len();
let target = Array2::from_shape_vec((n_sel, p), rows_out)
.map_err(|e| format!("collect_stratified_target: target assembly failed: {e}"))?;
Ok(StratifiedCorpusTarget {
target,
row_ids,
likelihood_weights,
design,
corpus_rows,
})
}
#[cfg(test)]
mod tests {
use super::super::shard_reader::{MmapShardSource, encode_shard_bytes};
use super::*;
use gam_solve::row_sampling_measure::RowSamplingMeasure;
use ndarray::{Array2, s};
use std::io::Write;
use std::path::PathBuf;
fn gauss(counter: &mut u64) -> f64 {
*counter = counter.wrapping_add(1);
let a = splitmix64_hash(*counter ^ 0x1234_5678);
*counter = counter.wrapping_add(1);
let b = splitmix64_hash(*counter ^ 0x9ABC_DEF0);
let u1 = ((a >> 11) as f64 / (1u64 << 53) as f64).max(1e-12);
let u2 = (b >> 11) as f64 / (1u64 << 53) as f64;
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn planted_corpus(
n: usize,
p: usize,
k_dom: usize,
rare_rows: usize,
) -> (Array2<f64>, Vec<usize>) {
let mut rows = Array2::<f64>::zeros((n, p));
let mut ctr = 1u64;
let mut rare_idx: Vec<usize> = Vec::with_capacity(rare_rows.min(n));
let mut seen: std::collections::BTreeSet<usize> = std::collections::BTreeSet::new();
let mut probe = 0u64;
while rare_idx.len() < rare_rows.min(n) {
let idx = (splitmix64_hash(probe ^ 0x5236_9E11_A73C_D0F5) as usize) % n;
probe = probe.wrapping_add(1);
if seen.insert(idx) {
rare_idx.push(idx);
}
}
rare_idx.sort_unstable();
let rare_set: std::collections::BTreeSet<usize> = rare_idx.iter().copied().collect();
for i in 0..n {
for d in 0..k_dom.min(p) {
rows[[i, d]] = 4.0 * gauss(&mut ctr);
}
for d in 0..p {
rows[[i, d]] += 0.02 * gauss(&mut ctr);
}
if rare_set.contains(&i) && k_dom < p {
rows[[i, k_dom]] += 3.0;
}
}
(rows, rare_idx)
}
fn temp_shard_dir(name: &str, rows: &Array2<f64>, split_at: usize) -> PathBuf {
let mut dir = std::env::temp_dir();
dir.push(format!(
"gam-residual-stratify-test-{}-{}",
std::process::id(),
name
));
std::fs::create_dir_all(&dir).expect("create dir");
let n = rows.nrows();
let split = split_at.min(n);
let parts = [
("a.shard", rows.slice(s![..split, ..])),
("b.shard", rows.slice(s![split.., ..])),
];
for (key, part) in parts {
let bytes = encode_shard_bytes(part);
let mut f = std::fs::File::create(dir.join(key)).expect("create shard");
f.write_all(&bytes).expect("write shard");
f.sync_all().expect("sync");
}
dir
}
fn dominant_basis(p: usize, k_dom: usize) -> Array2<f64> {
let mut q = Array2::<f64>::zeros((p, k_dom));
for d in 0..k_dom {
q[[d, d]] = 1.0;
}
q
}
#[test]
fn stratification_surfaces_rare_structure_uniform_does_not() {
let n = 300_000usize;
let p = 8usize;
let k_dom = 4usize;
let rare_rows = 30usize; let budget = 6_000usize;
let (rows, rare_idx) = planted_corpus(n, p, k_dom, rare_rows);
let dir = temp_shard_dir("surface", &rows, n / 2);
let mut src = MmapShardSource::open_dir(&dir).expect("open");
let screen = SpanResidualEnergy::new(dominant_basis(p, k_dom));
let collected = collect_stratified_target(&mut src, &screen, budget, 7).expect("collect");
let top = collected.design.strata().last().expect("nonempty strata");
assert!(
top.censused && (top.pi - 1.0).abs() < 1e-12,
"top residual-energy stratum must be censused: {top:?}"
);
let selected: std::collections::BTreeSet<u64> = collected.row_ids.iter().copied().collect();
let rare_seen_strat = rare_idx
.iter()
.filter(|&&i| selected.contains(&(i as u64)))
.count();
assert!(
rare_seen_strat as f64 >= 0.8 * rare_idx.len() as f64,
"stratification must surface ≥80% of the rare structure: {rare_seen_strat}/{}",
rare_idx.len()
);
for (k, &rid) in collected.row_ids.iter().enumerate() {
if rare_idx.binary_search(&(rid as usize)).is_ok() {
assert!(
(collected.likelihood_weights[k] - 1.0).abs() < 1e-12,
"censused rare row {rid} must have HT weight 1.0"
);
}
}
let uniform_measure = RowSamplingMeasure::uniform(n);
let uniform_sample = uniform_measure.designed_subsample(budget, 7);
let uniform_selected: std::collections::BTreeSet<usize> =
uniform_sample.rows.iter().copied().collect();
let rare_seen_uniform = rare_idx
.iter()
.filter(|&&i| uniform_selected.contains(&i))
.count();
assert!(
rare_seen_strat > rare_seen_uniform * 4,
"stratification must vastly out-recall uniform: strat={rare_seen_strat} \
uniform={rare_seen_uniform}"
);
assert!(
(rare_seen_uniform as f64) <= 0.25 * rare_idx.len() as f64,
"uniform baseline should surface few rare rows, got {rare_seen_uniform}"
);
if let Err(err) = std::fs::remove_dir_all(&dir) {
log::debug!("fixture cleanup left {} behind: {err}", dir.display());
}
}
#[test]
fn stratified_design_is_horvitz_thompson_unbiased() {
let n = 40_000usize;
let p = 6usize;
let k_dom = 3usize;
let (rows, _rare) = planted_corpus(n, p, k_dom, 40);
let dir = temp_shard_dir("ht", &rows, n / 3);
let mut src = MmapShardSource::open_dir(&dir).expect("open");
let screen = SpanResidualEnergy::new(dominant_basis(p, k_dom));
let collected = collect_stratified_target(&mut src, &screen, 4_000, 3).expect("collect");
let est = collected.estimated_corpus_rows();
assert!(
(est - n as f64).abs() < 0.15 * n as f64,
"HT corpus estimate {est} too far from N = {n}"
);
assert!(
collected
.likelihood_weights
.iter()
.all(|&w| w.is_finite() && w >= 1.0 - 1e-12)
);
if let Err(err) = std::fs::remove_dir_all(&dir) {
log::debug!("fixture cleanup left {} behind: {err}", dir.display());
}
}
#[test]
fn collection_is_deterministic() {
let n = 20_000usize;
let p = 5usize;
let (rows, _rare) = planted_corpus(n, p, 3, 30);
let dir = temp_shard_dir("determinism", &rows, n / 2);
let mut src = MmapShardSource::open_dir(&dir).expect("open");
let screen = SpanResidualEnergy::new(dominant_basis(p, 3));
let a = collect_stratified_target(&mut src, &screen, 2_000, 11).expect("a");
let b = collect_stratified_target(&mut src, &screen, 2_000, 11).expect("b");
assert_eq!(a.row_ids, b.row_ids, "same seed ⇒ identical selection");
assert_eq!(a.likelihood_weights, b.likelihood_weights);
if let Err(err) = std::fs::remove_dir_all(&dir) {
log::debug!("fixture cleanup left {} behind: {err}", dir.display());
}
}
#[test]
fn full_budget_is_the_exact_pass_with_unit_weights() {
let n = 5_000usize;
let p = 4usize;
let (rows, _rare) = planted_corpus(n, p, 2, 20);
let dir = temp_shard_dir("full", &rows, n / 2);
let mut src = MmapShardSource::open_dir(&dir).expect("open");
let screen = SpanResidualEnergy::new(dominant_basis(p, 2));
let collected = collect_stratified_target(&mut src, &screen, n, 1).expect("collect");
assert_eq!(collected.len(), n);
assert_eq!(collected.row_ids, (0..n as u64).collect::<Vec<_>>());
assert!(collected.likelihood_weights.iter().all(|&w| w == 1.0));
assert!(collected.design.strata().iter().all(|s| s.censused));
if let Err(err) = std::fs::remove_dir_all(&dir) {
log::debug!("fixture cleanup left {} behind: {err}", dir.display());
}
}
#[test]
fn no_residual_variation_degrades_to_uniform_rate() {
let n = 10_000usize;
let p = 3usize;
let rows = Array2::<f64>::from_shape_fn((n, p), |(_, d)| if d == 0 { 1.0 } else { 0.0 });
let dir = temp_shard_dir("flat", &rows, n / 2);
let mut src = MmapShardSource::open_dir(&dir).expect("open");
let screen = SpanResidualEnergy::new(Array2::<f64>::zeros((p, 0)));
let budget = 1_000usize;
let design = design_stratified_subsample(&mut src, &screen, budget).expect("design");
assert_eq!(design.strata().len(), 1);
let f = budget as f64 / n as f64;
assert!(
(design.strata()[0].pi - f).abs() < 1e-9,
"flat energy must give the uniform rate f = {f}, got {}",
design.strata()[0].pi
);
if let Err(err) = std::fs::remove_dir_all(&dir) {
log::debug!("fixture cleanup left {} behind: {err}", dir.display());
}
}
#[test]
fn sturges_cap_matches_floor_log2_plus_one() {
assert_eq!(sturges_stratum_cap(1), 1);
assert_eq!(sturges_stratum_cap(2), 2);
assert_eq!(sturges_stratum_cap(255), 8);
assert_eq!(sturges_stratum_cap(256), 9);
assert_eq!(sturges_stratum_cap(100_000_000), 27);
}
}