Skip to main content

polars_utils/
kahan_sum.rs

1use std::ops::{Add, AddAssign};
2
3use num_traits::Num;
4
5use crate::float::IsFloat;
6
7#[derive(Debug, Clone)]
8pub struct KahanSum<T> {
9    sum: T,
10    err: T,
11}
12
13impl<T: IsFloat + Num + Copy> KahanSum<T> {
14    #[inline]
15    pub fn new(v: T) -> Self {
16        KahanSum {
17            sum: v,
18            err: T::zero(),
19        }
20    }
21
22    #[inline(always)]
23    pub fn sum(&self) -> T {
24        self.sum
25    }
26}
27
28impl<T: Num> Default for KahanSum<T> {
29    #[inline]
30    fn default() -> Self {
31        KahanSum {
32            sum: T::zero(),
33            err: T::zero(),
34        }
35    }
36}
37
38impl<T: IsFloat + Num + AddAssign + Copy> AddAssign<T> for KahanSum<T> {
39    #[inline]
40    fn add_assign(&mut self, rhs: T) {
41        let y = rhs - self.err;
42        let new_sum = self.sum + y;
43        let new_err = (new_sum - self.sum) - y;
44        self.sum = new_sum;
45        if new_err.is_finite() {
46            // Ensure err stays finite so we don't introduce NaNs through Inf - Inf.
47            self.err = new_err;
48        }
49    }
50}
51
52impl<T: IsFloat + Num + AddAssign + Copy> Add<T> for KahanSum<T> {
53    type Output = Self;
54
55    #[inline]
56    fn add(self, rhs: T) -> Self::Output {
57        let mut rv = self;
58        rv += rhs;
59        rv
60    }
61}