Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
140 changes: 140 additions & 0 deletions src/hpc/cascade.rs
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,57 @@ impl Cascade {
}
}

/// Fold a whole batch of distances into the rolling floor at once.
///
/// The batch is reduced to exact integer moments with
/// [`moments_u32`](crate::hpc::statistics::moments_u32) and merged into
/// the running `(n, μ, σ)` with the parallel-variance merge
/// (Chan, Golub & LeVeque), so the resulting `mu`/`sigma`/`observations`
/// match calling [`observe`](Self::observe) once per distance up to f64
/// rounding. Shards computed on different threads can each be folded in
/// this way, in any order.
///
/// Drift is judged once per batch, not per element: an alert fires when
/// the batch moves μ by more than 2σ of the pre-batch state (and the state
/// had at least 10 observations and σ > 0), the same test `observe`
/// applies to a single step, so `observe_batch(&[x])` alerts exactly when
/// `observe(x)` does.
pub fn observe_batch(&mut self, distances: &[u32]) -> Option<ShiftAlert> {
let b = crate::hpc::statistics::moments_u32(distances);
if b.n == 0 {
return None;
}
let old_mu = self.mu;
let old_sigma = self.sigma;
let old_n = self.observations;
let n_a = old_n as f64;
let n_b = b.n as f64;
let n = n_a + n_b;
let (mean_b, m2_b) = (b.mean(), b.variance() * n_b);
if old_n == 0 {
self.mu = mean_b;
self.sigma = (m2_b / n_b).sqrt();
} else {
let delta = mean_b - old_mu;
self.mu = old_mu + delta * n_b / n;
let m2 = old_sigma * old_sigma * n_a + m2_b + delta * delta * n_a * n_b / n;
self.sigma = (m2 / n).sqrt();
}
self.observations = old_n + b.n as usize;

if old_n >= 10 && old_sigma > 0.0 && (self.mu - old_mu).abs() > 2.0 * old_sigma {
Some(ShiftAlert {
old_mu,
new_mu: self.mu,
old_sigma,
new_sigma: self.sigma,
observations: self.observations,
})
} else {
None
}
}

pub fn recalibrate(&mut self, alert: &ShiftAlert) {
self.mu = alert.new_mu;
self.sigma = alert.new_sigma;
Expand Down Expand Up @@ -768,6 +819,95 @@ mod tests {
assert_eq!(got, exact_hits(&expected, threshold));
}

fn close(a: f64, b: f64) -> bool {
(a - b).abs() <= 1e-9 * a.abs().max(b.abs()).max(1.0)
}

fn noisy(n: usize, base: u32, spread: u32, mut s: u64) -> Vec<u32> {
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
base + (s as u32) % spread
})
.collect()
}

/// One `observe_batch` lands on the same rolling floor as `observe`
/// called per distance, from an empty state and from a warm one.
#[test]
fn observe_batch_matches_sequential_observe() {
let warm = noisy(300, 8000, 400, 1);
let batch = noisy(2000, 8100, 500, 2);
for start in [&[][..], &warm[..]] {
let mut seq = Cascade::from_threshold(8000, 2048);
let mut bat = Cascade::from_threshold(8000, 2048);
for &d in start {
seq.observe(d);
bat.observe(d);
}
for &d in &batch {
seq.observe(d);
}
bat.observe_batch(&batch);
assert_eq!(seq.observations(), bat.observations());
assert!(close(seq.mu(), bat.mu()), "mu {} vs {}", seq.mu(), bat.mu());
assert!(close(seq.sigma(), bat.sigma()), "sigma {} vs {}", seq.sigma(), bat.sigma());
}
}

/// Shard-parallel use: folding shards in any order gives the same floor
/// as folding the whole batch.
#[test]
fn observe_batch_shards_merge_in_any_order() {
let x = noisy(3001, 8000, 700, 3);
let mut whole = Cascade::from_threshold(8000, 2048);
whole.observe_batch(&x);
let shards: Vec<&[u32]> = x.chunks(640).collect();
for order in [[0usize, 1, 2, 3, 4], [4, 2, 0, 3, 1]] {
let mut c = Cascade::from_threshold(8000, 2048);
for i in order {
c.observe_batch(shards[i]);
}
assert!(close(c.mu(), whole.mu()));
assert!(close(c.sigma(), whole.sigma()));
assert_eq!(c.observations(), whole.observations());
}
}

/// The drift alert can fire (a batch from a distribution shifted far
/// past 2σ) and stays silent on a batch from the same distribution.
#[test]
fn observe_batch_alerts_on_a_shift_only() {
let mut c = Cascade::from_threshold(8000, 2048);
assert!(c.observe_batch(&noisy(500, 8000, 100, 4)).is_none(), "first batch has no prior");
assert!(c.observe_batch(&noisy(500, 8000, 100, 5)).is_none(), "same distribution");
let alert = c
.observe_batch(&noisy(5000, 9000, 100, 6))
.expect("shifted batch must alert");
assert!(alert.new_mu > alert.old_mu + 2.0 * alert.old_sigma);
}

/// A singleton batch is the same event as one `observe` call, so both
/// must agree on the alert gate at the boundary: exactly 10 prior
/// observations, then an outlier. `observe` counts the new element
/// before testing `> 10`; `observe_batch` must admit the same state.
#[test]
fn observe_batch_singleton_matches_observe_alert_gate() {
let mut seq = Cascade::from_threshold(8000, 2048);
let mut batch = Cascade::from_threshold(8000, 2048);
for i in 0..10u32 {
assert!(seq.observe(8000 + 100 * (i % 2)).is_none());
assert!(batch.observe(8000 + 100 * (i % 2)).is_none());
}
assert_eq!(seq.observations(), 10);
let a = seq.observe(20000);
let b = batch.observe_batch(&[20000]);
assert!(a.is_some(), "observe must alert on the outlier");
assert!(b.is_some(), "observe_batch(&[x]) must alert exactly like observe(x)");
}

#[test]
fn packed_database_roundtrip() {
let vec_bytes = 256;
Expand Down
230 changes: 230 additions & 0 deletions src/hpc/statistics.rs
Original file line number Diff line number Diff line change
Expand Up @@ -360,3 +360,233 @@ mod tests {
}
}
}

// ── Batch moments for shard-parallel Welford ──────────────────────────────

/// Exact first and second moments of a `u32` sample: count, `Σx` and `Σx²`.
///
/// These three integers are the sufficient statistics of a Welford rolling
/// floor. Unlike a running `(mean, M2)` pair they merge by plain addition, so
/// [`MomentsU32::merge`] is exact, associative and commutative: shards of a
/// sample can be reduced in parallel, in any order, and combined into the
/// same result as one sequential pass. Floats appear only in [`mean`] and
/// [`variance`], at the end.
///
/// [`mean`]: MomentsU32::mean
/// [`variance`]: MomentsU32::variance
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct MomentsU32 {
/// Number of values.
pub n: u64,
/// `Σx`.
pub sum: u128,
/// `Σx²`.
pub sum_sq: u128,
}

impl MomentsU32 {
/// Moments of the union of two samples — exact integer addition.
#[inline]
#[must_use]
pub fn merge(self, other: Self) -> Self {
Self {
n: self.n + other.n,
sum: self.sum + other.sum,
sum_sq: self.sum_sq + other.sum_sq,
}
}

/// Sample mean; `0.0` for an empty sample.
pub fn mean(&self) -> f64 {
if self.n == 0 {
0.0
} else {
self.sum as f64 / self.n as f64
}
}

/// Population variance `E[(X - μ)²]`; `0.0` for an empty sample.
///
/// Computed as `(n·Σx² − (Σx)²) / n²`. The numerator is formed exactly in
/// `u128` whenever it fits (always, for Hamming-scale data: n ≤ 2³², x ≤
/// 2¹⁷), so there is no cancellation between two large floats; only the
/// final division rounds. Past that range it centres the sums on the
/// integer part of the mean in `u128` first, so the `f64` step still
/// works on small, variance-sized quantities.
pub fn variance(&self) -> f64 {
if self.n == 0 {
return 0.0;
}
let n = u128::from(self.n);
match (n.checked_mul(self.sum_sq), self.sum.checked_mul(self.sum)) {
(Some(a), Some(b)) => (a - b) as f64 / (self.n as f64 * self.n as f64),
_ => {
// Centre on the integer part of the mean, q = ⌊Σx / n⌋, with
// remainder r = Σx − n·q < n. Then, exactly in u128,
// Σ(x − q)² = Σx² − q·Σx − q·r, which is < n·2⁶⁴ and never
// negative at any step. The true M2 is that minus r²/n, and
// both terms are O(n·(σ² + 1)), so the one float subtraction
// cannot cancel the variance away.
let q = self.sum / n;
let r = self.sum % n;
let centred = self.sum_sq - q * self.sum - q * r;
let rf = r as f64;
let nf = self.n as f64;
((centred as f64 - rf * rf / nf) / nf).max(0.0)
}
}
}
}

/// Exact [`MomentsU32`] of `values`, eight lanes at a time through `U64x8`.
///
/// Each value is widened to `u64` and its square split into 32-bit halves, so
/// every lane add is below 2³², and the lanes are drained into `u128` totals
/// every 2²⁸ chunks, before an 8-lane reduction could overflow `u64`. The
/// widening multiply is plain lane-wise Rust that the compiler vectorizes
/// (`vpmuludq` on x86); the accumulation and the final reduction go through
/// the polyfill, so every backend runs the same code.
///
/// # Example
///
/// ```
/// use ndarray::hpc::statistics::moments_u32;
///
/// let m = moments_u32(&[1, 2, 3, 4]);
/// assert_eq!((m.n, m.sum, m.sum_sq), (4, 10, 30));
/// assert_eq!(m.mean(), 2.5);
/// assert_eq!(m.variance(), 1.25);
/// ```
pub fn moments_u32(values: &[u32]) -> MomentsU32 {
use crate::simd::U64x8;

/// Chunks per drain. Every lane add is below 2^32, so after 2^28 chunks
/// a lane is below 2^60 and the 8-lane `reduce_sum` below 2^63 — the
/// reduction, not the lane, is the binding limit.
const DRAIN: usize = 1 << 28;

let (chunks, tail) = values.as_chunks::<8>();
let mut out = MomentsU32 {
n: values.len() as u64,
..MomentsU32::default()
};
for block in chunks.chunks(DRAIN) {
let (mut s, mut lo, mut hi) = (U64x8::splat(0), U64x8::splat(0), U64x8::splat(0));
for c in block {
let x: [u64; 8] = core::array::from_fn(|i| u64::from(c[i]));
let sq: [u64; 8] = core::array::from_fn(|i| x[i] * x[i]);
s += U64x8::from_array(x);
lo += U64x8::from_array(core::array::from_fn(|i| sq[i] & 0xFFFF_FFFF));
hi += U64x8::from_array(core::array::from_fn(|i| sq[i] >> 32));
}
out.sum += u128::from(s.reduce_sum());
out.sum_sq += u128::from(lo.reduce_sum()) + (u128::from(hi.reduce_sum()) << 32);
}
for &v in tail {
out.sum += u128::from(v);
out.sum_sq += u128::from(v) * u128::from(v);
}
out
}

#[cfg(test)]
mod moments_tests {
use super::*;

/// Past the exact-`u128` range the variance must not cancel: 2³³ values
/// split evenly between `u32::MAX` and `u32::MAX - 1` have variance
/// exactly 0.25, and `n·Σx²` overflows `u128`, so this takes the
/// fallback path.
#[test]
fn variance_fallback_does_not_cancel() {
let half = 1u128 << 32;
let hi = u128::from(u32::MAX);
let lo = hi - 1;
let m = MomentsU32 {
n: 1u64 << 33,
sum: half * (hi + lo),
sum_sq: half * (hi * hi + lo * lo),
};
assert!(u128::from(m.n).checked_mul(m.sum_sq).is_none(), "fixture must take the fallback");
assert!((m.variance() - 0.25).abs() < 1e-9, "variance {}", m.variance());
}

fn xorshift(n: usize, mut s: u64, mask: u32) -> Vec<u32> {
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s as u32) & mask
})
.collect()
}

fn reference(x: &[u32]) -> MomentsU32 {
MomentsU32 {
n: x.len() as u64,
sum: x.iter().map(|&v| u128::from(v)).sum(),
sum_sq: x.iter().map(|&v| u128::from(v) * u128::from(v)).sum(),
}
}

/// Exact against a u128 reference at every length across the 8-lane
/// chunk boundary, for small (Hamming-scale) and full-range values.
#[test]
fn moments_u32_is_exact() {
for mask in [0x3FFF, u32::MAX] {
let x = xorshift(1000, 0x9E37_79B9_7F4A_7C15, mask);
for n in (0..=40).chain([63, 64, 65, 999, 1000]) {
assert_eq!(moments_u32(&x[..n]), reference(&x[..n]), "n={n} mask={mask:#x}");
}
}
}

/// Worst case for the square accumulator: every square is (2^32-1)^2,
/// which alone nearly fills a u64. A lane that summed whole squares
/// would overflow on the second element.
#[test]
fn moments_u32_does_not_overflow_at_u32_max() {
let x = vec![u32::MAX; 100_003];
let m = moments_u32(&x);
let v = u128::from(u32::MAX);
assert_eq!(m.n, 100_003);
assert_eq!(m.sum, 100_003 * v);
assert_eq!(m.sum_sq, 100_003 * v * v);
}

/// Merging is exact integer addition: the moments of a concatenation
/// equal the merge of the parts' moments at any split point and in
/// either order — which is what makes shard-parallel statistics exact.
#[test]
fn merge_is_exact_and_order_independent() {
let x = xorshift(777, 42, u32::MAX);
let whole = moments_u32(&x);
for split in [0, 1, 7, 8, 9, 400, 776, 777] {
let (a, b) = x.split_at(split);
assert_eq!(moments_u32(a).merge(moments_u32(b)), whole, "split={split}");
assert_eq!(moments_u32(b).merge(moments_u32(a)), whole, "reversed split={split}");
}
}

/// Mean and population variance against a two-pass f64 reference.
#[test]
fn mean_and_variance_match_two_pass() {
for mask in [0x3FFF, u32::MAX] {
let x = xorshift(5000, 7, mask);
let n = x.len() as f64;
let mean = x.iter().map(|&v| f64::from(v)).sum::<f64>() / n;
let var = x
.iter()
.map(|&v| (f64::from(v) - mean).powi(2))
.sum::<f64>()
/ n;
let m = moments_u32(&x);
assert!((m.mean() - mean).abs() <= 1e-9 * mean.abs(), "mean mask={mask:#x}");
assert!((m.variance() - var).abs() <= 1e-9 * var, "variance mask={mask:#x}");
}
assert_eq!(MomentsU32::default().mean(), 0.0);
assert_eq!(MomentsU32::default().variance(), 0.0);
assert_eq!(moments_u32(&[5, 5, 5]).variance(), 0.0, "constant input has zero variance");
}
}
1 change: 1 addition & 0 deletions src/simd.rs
Original file line number Diff line number Diff line change
Expand Up @@ -617,6 +617,7 @@ pub use crate::hpc::fft::{wht_f32, wht_f32_new};
pub use crate::hpc::fingerprint::{
vector_config, Fingerprint, Fingerprint1K, Fingerprint2K, Fingerprint64K, VectorConfig, VectorWidth,
};
pub use crate::hpc::statistics::{moments_u32, MomentsU32};

// PR-X1 — SoA carrier + const-size slice helpers, dispatched from their
// respective `simd_{type}.rs` modules. The W1a consumer contract forbids
Expand Down
Loading