const MAX_UNITS: u64 = 1 << 53;
#[derive(Clone, Copy, Debug, PartialEq)]
pub(crate) struct SumDomain {
max: f32,
grain: i32,
finite: bool,
}
impl SumDomain {
pub(crate) const EMPTY: Self = Self {
max: 0.0,
grain: i32::MAX,
finite: true,
};
pub(crate) fn of_slice<T: Sync>(values: &[T], project: impl Fn(&T) -> f32 + Sync) -> Self {
use rayon::prelude::*;
const CHUNK: usize = 4096;
values
.par_chunks(CHUNK)
.map(|chunk| Self::of(chunk.iter().map(&project)))
.reduce(|| Self::EMPTY, Self::combine)
}
fn combine(self, other: Self) -> Self {
Self {
max: self.max.max(other.max),
grain: self.grain.min(other.grain),
finite: self.finite && other.finite,
}
}
pub(crate) fn of(values: impl IntoIterator<Item = f32>) -> Self {
values.into_iter().fold(Self::EMPTY, |domain, v| {
if !v.is_finite() {
Self {
finite: false,
..domain
}
} else if v == 0.0 {
domain
} else {
Self {
max: domain.max.max(v.abs()),
grain: domain.grain.min(grain(v)),
finite: domain.finite,
}
}
})
}
pub(crate) fn sums_exact(&self, n: usize) -> bool {
if !self.finite {
return false;
}
if self.grain == i32::MAX {
return true;
}
let max_units = f64::from(self.max) * pow2(-self.grain);
max_units <= MAX_UNITS as f64
&& (n as u128) * u128::from(max_units as u64) <= u128::from(MAX_UNITS)
}
pub(crate) fn max_units(&self) -> u64 {
if self.grain == i32::MAX {
return 0;
}
(f64::from(self.max) * pow2(-self.grain)) as u64
}
pub(crate) fn units(&self, x: f32) -> i64 {
if self.grain == i32::MAX {
return 0;
}
(f64::from(x) * pow2(-self.grain)) as i64
}
pub(crate) fn value(&self, units: i64) -> f64 {
if self.grain == i32::MAX {
return 0.0;
}
units as f64 * pow2(self.grain)
}
}
fn grain(v: f32) -> i32 {
let bits = v.to_bits();
let biased = ((bits >> 23) & 0xFF) as i32;
let mantissa = bits & 0x7F_FFFF;
let (significand, exponent) = if biased == 0 {
(mantissa, -149)
} else {
(mantissa | 0x80_0000, biased - 150)
};
exponent + significand.trailing_zeros() as i32
}
fn pow2(k: i32) -> f64 {
debug_assert!((-1022..=1023).contains(&k));
f64::from_bits(((k + 1023) as u64) << 52)
}
#[cfg(test)]
mod tests {
use super::*;
fn gpu_sum(values: &[f32], slice_len: usize) -> f64 {
let domain = SumDomain::of(values.iter().copied());
let total: i64 = values
.chunks(slice_len)
.map(|slice| slice.iter().map(|&x| domain.units(x)).sum::<i64>())
.sum();
domain.value(total)
}
fn cpu_sum(values: &[f32]) -> f64 {
values.iter().fold(0.0f64, |s, &x| s + f64::from(x))
}
#[test]
fn max_units_matches_the_exactness_bound() {
let values = [0.75f32, -3.0, 1.0 + f32::EPSILON];
let domain = SumDomain::of(values.iter().copied());
let max_units = domain.max_units();
assert!(max_units >= 1);
assert!(domain.sums_exact(1));
assert!(domain.sums_exact((MAX_UNITS / max_units) as usize));
assert!(!domain.sums_exact((MAX_UNITS / max_units) as usize + 1));
assert_eq!(SumDomain::EMPTY.max_units(), 0);
assert_eq!(SumDomain::of([0.0f32, -0.0].iter().copied()).max_units(), 0);
}
#[test]
fn the_parallel_fold_matches_the_serial_one() {
let mut values = vec![0.0, -0.0, -3.0, 0.75, 1.0 + f32::EPSILON, 2f32.powi(50)];
values.extend((0..10_000).map(|i| (i as f32 * 0.37).sin() * 2f32.powi(i % 40 - 20)));
for extra in [
vec![],
vec![f32::INFINITY, -1.0],
vec![f32::NAN],
vec![f32::from_bits(1)],
] {
let all: Vec<f32> = values
.iter()
.copied()
.chain(extra.iter().copied())
.collect();
assert_eq!(
SumDomain::of_slice(&all, |v| *v),
SumDomain::of(all.iter().copied())
);
}
}
#[test]
fn grain_is_the_lowest_set_bit() {
assert_eq!(grain(1.0), 0);
assert_eq!(grain(-3.0), 0);
assert_eq!(grain(0.75), -2);
assert_eq!(grain(2f32.powi(50)), 50);
assert_eq!(grain(8_388_609.0), 0); assert_eq!(grain(f32::from_bits(1)), -149); assert_eq!(grain(f32::from_bits(0x40_0000)), -127); assert_eq!(grain(1.0 + f32::EPSILON), -23);
}
#[test]
fn grains_round_trip_exactly() {
for values in [
vec![
f32::from_bits(1),
f32::from_bits(0x7F_FFFF),
-f32::MIN_POSITIVE,
],
vec![f32::MAX, -(2f32.powi(104))],
vec![0.75, -3.0, 1.0 + f32::EPSILON],
] {
let domain = SumDomain::of(values.iter().copied());
assert!(domain.sums_exact(1));
for &x in &values {
assert_eq!(domain.value(domain.units(x)), f64::from(x), "{x:e}");
}
}
let zeros = SumDomain::of([0.0f32, -0.0]);
assert_eq!(zeros.units(0.0), 0);
assert_eq!(zeros.value(0), 0.0);
}
#[test]
fn wide_dynamic_range_bounds_the_row_count() {
let six = [
2f32.powi(50),
2f32.powi(26),
2f32.powi(23) + 1.0,
-(2f32.powi(50)),
-(2f32.powi(26)),
-(2f32.powi(23)),
];
let domain = SumDomain::of(six);
assert!(domain.sums_exact(8));
assert!(!domain.sums_exact(9));
assert!(!domain.sums_exact(8192));
assert_eq!(gpu_sum(&six, 4), cpu_sum(&six));
}
#[test]
fn domain_boundary() {
let ones = SumDomain::of([1.0f32; 3]);
assert!(ones.sums_exact(1 << 53));
assert!(!ones.sums_exact((1 << 53) + 1));
let mixed = SumDomain::of([4.0, -0.5, 3.0 * 2f32.powi(-10), 0.0]);
assert!(mixed.sums_exact(1 << 41));
assert!(!mixed.sums_exact((1 << 41) + 1));
assert!(SumDomain::of([f32::from_bits(1)]).sums_exact(1 << 53));
assert!(!SumDomain::of([2f32.powi(60), 1.0]).sums_exact(1));
assert!(SumDomain::of([0.0f32, -0.0]).sums_exact(usize::MAX));
assert!(SumDomain::EMPTY.sums_exact(usize::MAX));
assert!(!SumDomain::of([1.0, f32::NAN]).sums_exact(1));
assert!(!SumDomain::of([f32::INFINITY]).sums_exact(1));
}
#[test]
fn in_domain_sums_match_the_cpu_chain() {
let mut state = 0x9E37_79B9_7F4A_7C15u64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for case in 0..200 {
let n = 4096 + (next() % 4096) as usize;
let top = 53 - (n as f64).log2().ceil() as i32;
let values: Vec<f32> = (0..n)
.map(|_| {
let r = next();
let sign = if (r >> 40) % 8 == 0 { -1.0 } else { 1.0 };
let v = match (r >> 1) % 4 {
0 | 1 => ((r >> 8) % (1 << 24)) as f32 * 2f32.powi(top - 24),
2 => ((r >> 8) % 1000 * 2 + 1) as f32,
_ => 2f32.powi(top),
};
sign * v
})
.collect();
let domain = SumDomain::of(values.iter().copied());
assert!(domain.sums_exact(n), "case {case} must be in the domain");
for slice_len in [1, 7, 128, 1024, n] {
assert_eq!(
gpu_sum(&values, slice_len),
cpu_sum(&values),
"case {case}, slices of {slice_len}"
);
}
}
}
#[test]
fn the_cpu_chain_rounds_just_past_the_bound() {
let k = (1 << 23) + 1;
let values: Vec<f32> = std::iter::once(1.0)
.chain(std::iter::repeat_n(2f32.powi(30), k))
.collect();
let exact = 1 + (k as i64) * (1 << 30);
assert_eq!(cpu_sum(&values) as i64, exact - 1);
assert!(!SumDomain::of(values.iter().copied()).sums_exact(values.len()));
}
}