use crate::scalar::Numeric;
const BLOCK: usize = 128;
const LEVELS: usize = 32;
pub(crate) struct PairwiseSum<T: Numeric> {
tree: [T; LEVELS],
occupied: u64,
block_sum: T,
block_count: usize,
}
impl<T: Numeric> PairwiseSum<T> {
#[inline]
pub(crate) fn new() -> Self {
Self {
tree: [T::ZERO; LEVELS],
occupied: 0,
block_sum: T::ZERO,
block_count: 0,
}
}
#[inline]
pub(crate) fn add(&mut self, value: T) {
self.block_sum += value;
self.block_count += 1;
if self.block_count == BLOCK {
let block = self.block_sum;
self.push(block);
self.block_sum = T::ZERO;
self.block_count = 0;
}
}
fn push(&mut self, mut value: T) {
let mut level = 0;
while level < LEVELS && self.occupied & (1u64 << level) != 0 {
value = self.tree[level & (LEVELS - 1)] + value;
self.occupied &= !(1u64 << level);
level += 1;
}
if level == LEVELS {
level = LEVELS - 1;
}
self.tree[level & (LEVELS - 1)] = value;
self.occupied |= 1u64 << level;
}
#[inline]
pub(crate) fn total(&self) -> T {
let mut sum = self.block_sum;
let top = (u64::BITS - self.occupied.leading_zeros()) as usize;
#[allow(clippy::needless_range_loop)]
for level in 0..top {
if self.occupied & (1u64 << level) != 0 {
sum += self.tree[level & (LEVELS - 1)];
}
}
sum
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_input_is_zero() {
let acc = PairwiseSum::<f64>::new();
assert_eq!(acc.total(), 0.0);
}
#[test]
fn matches_naive_on_short_input() {
let values = [0.1, 0.2, 0.3, 0.4, 1.5, -0.7, 2.25, 8.0];
let mut acc = PairwiseSum::new();
let mut naive = 0.0;
for &v in &values {
acc.add(v);
naive += v;
}
assert!((acc.total() - naive).abs() < 1e-12);
}
#[test]
fn f32_matches_naive_on_short_input() {
let values = [0.1f32, 0.2, 0.3, 0.4, 1.5, -0.7, 2.25, 8.0];
let mut acc = PairwiseSum::new();
let mut naive = 0.0f32;
for &v in &values {
acc.add(v);
naive += v;
}
assert!((acc.total() - naive).abs() < 1e-5);
}
#[test]
fn beats_naive_on_long_sum() {
let tiny = 2f64.powi(-53);
let n: u64 = 1 << 24;
let analytic = 1.0 + (n as f64) * tiny;
let mut acc = PairwiseSum::new();
acc.add(1.0);
let mut naive = 1.0f64;
for _ in 0..n {
acc.add(tiny);
naive += tiny;
}
let pairwise = acc.total();
assert_eq!(naive, 1.0, "naive should lose every tiny term");
let pairwise_err = (pairwise - analytic).abs();
let naive_err = (naive - analytic).abs();
assert!(
pairwise_err < 1e-12,
"pairwise error {pairwise_err:e} too large"
);
assert!(
pairwise_err < naive_err,
"pairwise ({pairwise_err:e}) should be strictly closer than naive ({naive_err:e})"
);
}
}