Skip to main content

sim_lib_numbers_tensor/implementation/
reduction.rs

1//! Shared, auditable floating-point reduction policies.
2/// Explicit addition policy for Tensor `sum` and `cumsum`.
3#[derive(Clone, Copy, Debug, Eq, PartialEq)]
4pub enum SumMode {
5    /// Sequential.
6    Naive,
7    /// Balanced tree.
8    Pairwise,
9    /// Neumaier compensated.
10    Neumaier,
11}
12/// Reduces a slice with the named policy.
13pub fn sum_f64(v: &[f64], m: SumMode) -> f64 {
14    match m {
15        SumMode::Naive => v.iter().sum(),
16        SumMode::Pairwise => {
17            if v.len() < 2 {
18                v.first().copied().unwrap_or(0.)
19            } else {
20                let n = v.len() / 2;
21                sum_f64(&v[..n], m) + sum_f64(&v[n..], m)
22            }
23        }
24        SumMode::Neumaier => {
25            let (mut s, mut c) = (0., 0.);
26            for &x in v {
27                let t = s + x;
28                c += if s.abs() >= x.abs() {
29                    (s - t) + x
30                } else {
31                    (x - t) + s
32                };
33                s = t
34            }
35            s + c
36        }
37    }
38}
39/// Produces prefix sums with the named policy.
40pub fn cumsum_f64(v: &[f64], m: SumMode) -> Vec<f64> {
41    match m {
42        SumMode::Naive => {
43            let mut s = 0.;
44            v.iter()
45                .map(|&x| {
46                    s += x;
47                    s
48                })
49                .collect()
50        }
51        SumMode::Pairwise => (1..=v.len()).map(|n| sum_f64(&v[..n], m)).collect(),
52        SumMode::Neumaier => {
53            let (mut s, mut c) = (0., 0.);
54            v.iter()
55                .map(|&x| {
56                    let t = s + x;
57                    c += if s.abs() >= x.abs() {
58                        (s - t) + x
59                    } else {
60                        (x - t) + s
61                    };
62                    s = t;
63                    s + c
64                })
65                .collect()
66        }
67    }
68}