Skip to main content

common_function/aggrs/approximate/
welford.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Mergeable Welford state for population standard deviation.
16//!
17//! Input samples and intermediate states must contain only finite values.
18
19use std::sync::Arc;
20
21use datafusion::arrow::array::ArrayRef;
22use datafusion::common::cast::{as_binary_array, as_primitive_array};
23use datafusion::common::not_impl_err;
24use datafusion::error::{DataFusionError, Result as DfResult};
25use datafusion::logical_expr::function::AccumulatorArgs;
26use datafusion::logical_expr::{Accumulator as DfAccumulator, AggregateUDF, Volatility};
27use datafusion::prelude::create_udaf;
28use datafusion_common::ScalarValue;
29use datatypes::arrow::datatypes::{DataType, Float64Type};
30
31pub const STDDEV_POP_STATE_NAME: &str = "stddev_pop_state";
32pub const STDDEV_POP_MERGE_NAME: &str = "stddev_pop_merge";
33
34const ENCODED_LEN: usize = 28;
35const MAGIC: &[u8; 4] = b"WLF1";
36
37#[derive(Debug, Clone, Copy, PartialEq)]
38pub(crate) struct WelfordState {
39    pub(crate) count: u64,
40    pub(crate) mean: f64,
41    pub(crate) m2: f64,
42}
43
44impl Default for WelfordState {
45    fn default() -> Self {
46        Self {
47            count: 0,
48            mean: 0.0,
49            m2: 0.0,
50        }
51    }
52}
53
54impl WelfordState {
55    pub(crate) fn encode(&self) -> [u8; ENCODED_LEN] {
56        let mut encoded = [0; ENCODED_LEN];
57        encoded[..4].copy_from_slice(MAGIC);
58        encoded[4..12].copy_from_slice(&self.count.to_le_bytes());
59        encoded[12..20].copy_from_slice(&self.mean.to_bits().to_le_bytes());
60        encoded[20..28].copy_from_slice(&self.m2.to_bits().to_le_bytes());
61        encoded
62    }
63
64    pub(crate) fn decode(encoded: &[u8]) -> DfResult<Self> {
65        if encoded.len() != ENCODED_LEN || &encoded[..4] != MAGIC {
66            return Err(invalid_state());
67        }
68
69        let state = Self {
70            count: decode_u64(encoded, 4),
71            mean: decode_f64(encoded, 12),
72            m2: decode_f64(encoded, 20),
73        };
74        if !state.is_valid() {
75            return Err(invalid_state());
76        }
77
78        Ok(state)
79    }
80
81    fn is_valid(&self) -> bool {
82        match self.count {
83            0 => self.mean.to_bits() == 0 && self.m2.to_bits() == 0,
84            1 => self.mean.is_finite() && self.m2.to_bits() == 0,
85            _ => self.mean.is_finite() && self.m2.is_finite() && self.m2 >= 0.0,
86        }
87    }
88
89    fn update(&mut self, sample: f64) -> DfResult<()> {
90        if !sample.is_finite() {
91            return Err(non_finite_input());
92        }
93
94        let count = self
95            .count
96            .checked_add(1)
97            .ok_or_else(|| DataFusionError::Execution("Welford count overflow".to_string()))?;
98        let candidate_state = if self.count == 0 {
99            Self {
100                count,
101                mean: sample,
102                m2: 0.0,
103            }
104        } else {
105            let delta = sample - self.mean;
106            if !delta.is_finite() {
107                return Err(non_finite_arithmetic());
108            }
109            let mean = self.mean + delta / count as f64;
110            let delta2 = sample - mean;
111            Self {
112                count,
113                mean,
114                m2: self.m2 + delta * delta2,
115            }
116        };
117        self.replace_with_candidate(candidate_state)
118    }
119
120    fn merge(&mut self, other: &Self) -> DfResult<()> {
121        if other.count == 0 {
122            return Ok(());
123        }
124        if self.count == 0 {
125            return self.replace_with_candidate(*other);
126        }
127
128        let count = self
129            .count
130            .checked_add(other.count)
131            .ok_or_else(|| DataFusionError::Execution("Welford count overflow".to_string()))?;
132        let delta = other.mean - self.mean;
133        if !delta.is_finite() {
134            return Err(non_finite_arithmetic());
135        }
136        let self_count = self.count as f64;
137        let other_count = other.count as f64;
138        let count_f64 = count as f64;
139        let mean_delta = if delta.abs() <= f64::MAX / other_count {
140            delta * other_count / count_f64
141        } else {
142            delta * (other_count / count_f64)
143        };
144        let weighted_count = self_count * other_count / count_f64;
145        let candidate_state = Self {
146            count,
147            mean: self.mean + mean_delta,
148            m2: self.m2 + other.m2 + checked_weighted_square(delta, weighted_count)?,
149        };
150        self.replace_with_candidate(candidate_state)
151    }
152
153    fn replace_with_candidate(&mut self, candidate_state: Self) -> DfResult<()> {
154        if !candidate_state.is_valid() {
155            return Err(non_finite_arithmetic());
156        }
157
158        *self = candidate_state;
159        Ok(())
160    }
161
162    pub(crate) fn population_stddev(&self) -> Option<f64> {
163        if self.count == 0 {
164            return None;
165        }
166
167        Some((self.m2 / self.count as f64).sqrt())
168    }
169}
170
171fn checked_weighted_square(delta: f64, weight: f64) -> DfResult<f64> {
172    if delta.abs() <= f64::MAX.sqrt() {
173        return Ok(delta * delta * weight);
174    }
175    if weight <= 1.0 {
176        // Applying the weight first avoids overflow when the weighted square is representable.
177        return Ok(delta * weight * delta);
178    }
179
180    Err(non_finite_arithmetic())
181}
182
183/// Accumulates and merges versioned Welford states.
184#[derive(Debug, Default)]
185pub struct WelfordAccumulator {
186    state: WelfordState,
187}
188
189impl WelfordAccumulator {
190    /// Creates the `stddev_pop_state` aggregate function.
191    pub fn state_udf_impl() -> AggregateUDF {
192        create_udaf(
193            STDDEV_POP_STATE_NAME,
194            vec![DataType::Float64],
195            Arc::new(DataType::Binary),
196            Volatility::Immutable,
197            Arc::new(Self::create_accumulator),
198            Arc::new(vec![DataType::Binary]),
199        )
200    }
201
202    /// Creates the `stddev_pop_merge` aggregate function.
203    pub fn merge_udf_impl() -> AggregateUDF {
204        create_udaf(
205            STDDEV_POP_MERGE_NAME,
206            vec![DataType::Binary],
207            Arc::new(DataType::Binary),
208            Volatility::Immutable,
209            Arc::new(Self::create_accumulator),
210            Arc::new(vec![DataType::Binary]),
211        )
212    }
213
214    fn create_accumulator(args: AccumulatorArgs) -> DfResult<Box<dyn DfAccumulator>> {
215        if args.is_distinct {
216            return not_impl_err!("Welford DISTINCT aggregations are not available");
217        }
218        Ok(Box::new(Self::default()))
219    }
220}
221
222impl DfAccumulator for WelfordAccumulator {
223    fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> {
224        let array = &values[0];
225        match array.data_type() {
226            DataType::Float64 => {
227                for sample in as_primitive_array::<Float64Type>(array)?.iter().flatten() {
228                    self.state.update(sample)?;
229                }
230            }
231            DataType::Binary => self.merge_batch(std::slice::from_ref(array))?,
232            other => {
233                return not_impl_err!("Welford functions do not support data type: {other}");
234            }
235        }
236        Ok(())
237    }
238
239    fn evaluate(&mut self) -> DfResult<ScalarValue> {
240        Ok(ScalarValue::Binary(Some(self.state.encode().to_vec())))
241    }
242
243    fn size(&self) -> usize {
244        std::mem::size_of::<Self>()
245    }
246
247    fn state(&mut self) -> DfResult<Vec<ScalarValue>> {
248        Ok(vec![ScalarValue::Binary(Some(
249            self.state.encode().to_vec(),
250        ))])
251    }
252
253    fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> {
254        let array = as_binary_array(&states[0])?;
255        for encoded in array.iter().flatten() {
256            self.state.merge(&WelfordState::decode(encoded)?)?;
257        }
258        Ok(())
259    }
260}
261
262fn decode_u64(encoded: &[u8], offset: usize) -> u64 {
263    let mut bytes = [0; 8];
264    bytes.copy_from_slice(&encoded[offset..offset + 8]);
265    u64::from_le_bytes(bytes)
266}
267
268fn decode_f64(encoded: &[u8], offset: usize) -> f64 {
269    f64::from_bits(decode_u64(encoded, offset))
270}
271
272fn invalid_state() -> DataFusionError {
273    DataFusionError::Execution("Invalid Welford state".to_string())
274}
275
276fn non_finite_input() -> DataFusionError {
277    DataFusionError::Execution("Welford state requires finite input values".to_string())
278}
279
280fn non_finite_arithmetic() -> DataFusionError {
281    DataFusionError::Execution("Welford arithmetic produced a non-finite state".to_string())
282}
283
284#[cfg(test)]
285mod tests {
286    use std::sync::Arc;
287
288    use datafusion::arrow::array::{ArrayRef, BinaryArray, Float64Array};
289    use datafusion_common::ScalarValue;
290
291    use super::*;
292
293    fn state_from_values(values: &[f64]) -> WelfordState {
294        let mut state = WelfordState::default();
295        for value in values {
296            state.update(*value).unwrap();
297        }
298        state
299    }
300
301    #[test]
302    fn test_welford_state_encoding_contract() {
303        let state = WelfordState {
304            count: 3,
305            mean: 2.0,
306            m2: 6.0,
307        };
308
309        let encoded = state.encode();
310
311        assert_eq!(
312            encoded,
313            [
314                b'W', b'L', b'F', b'1', // magic
315                3, 0, 0, 0, 0, 0, 0, 0, // count
316                0, 0, 0, 0, 0, 0, 0, 64, // mean
317                0, 0, 0, 0, 0, 0, 24, 64, // m2
318            ]
319        );
320        assert_eq!(WelfordState::decode(&encoded).unwrap(), state);
321    }
322
323    #[test]
324    fn test_welford_state_online_update() {
325        let mut state = WelfordState::default();
326        for value in [1.0, 2.0, 3.0, 4.0] {
327            state.update(value).unwrap();
328        }
329
330        assert_eq!(state.count, 4);
331        assert_eq!(state.mean, 2.5);
332        assert_eq!(state.m2, 5.0);
333        assert_eq!(state.population_stddev(), Some(1.25_f64.sqrt()));
334    }
335
336    #[test]
337    fn test_welford_state_empty_and_single_value() {
338        let mut state = WelfordState::default();
339        assert_eq!(state.population_stddev(), None);
340
341        state.update(42.0).unwrap();
342        assert_eq!(state.population_stddev(), Some(0.0));
343    }
344
345    #[test]
346    fn test_welford_non_finite_values_fail_independent_of_partitioning() {
347        for sample in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
348            let mut one_pass = WelfordState::default();
349            let update_failed = one_pass.update(sample).is_err();
350
351            let mut merged = WelfordState::default();
352            let non_finite_state = WelfordState {
353                count: 1,
354                mean: sample,
355                m2: 0.0,
356            };
357            let merge_failed = merged.merge(&non_finite_state).is_err();
358
359            assert_eq!((update_failed, merge_failed), (true, true));
360            assert_eq!(one_pass, WelfordState::default());
361            assert_eq!(merged, WelfordState::default());
362        }
363    }
364
365    #[test]
366    fn test_welford_state_rejects_malformed_encoding() {
367        assert!(WelfordState::decode(b"").is_err());
368        assert!(WelfordState::decode(&[0; 28]).is_err());
369
370        let mut encoded = WelfordState::default().encode().to_vec();
371        encoded.push(0);
372        assert!(WelfordState::decode(&encoded).is_err());
373
374        let noncanonical_empty = WelfordState {
375            count: 0,
376            mean: 1.0,
377            m2: 0.0,
378        };
379        assert!(WelfordState::decode(&noncanonical_empty.encode()).is_err());
380
381        let negative_m2 = WelfordState {
382            count: 2,
383            mean: 1.0,
384            m2: -1.0,
385        };
386        assert!(WelfordState::decode(&negative_m2.encode()).is_err());
387
388        for m2 in [1.0, -0.0] {
389            let noncanonical_singleton = WelfordState {
390                count: 1,
391                mean: 0.0,
392                m2,
393            };
394            assert!(WelfordState::decode(&noncanonical_singleton.encode()).is_err());
395        }
396
397        for (mean, m2) in [
398            (f64::NAN, 0.0),
399            (f64::INFINITY, 0.0),
400            (f64::NEG_INFINITY, 0.0),
401            (0.0, f64::NAN),
402            (0.0, f64::INFINITY),
403            (0.0, f64::NEG_INFINITY),
404        ] {
405            let non_finite = WelfordState { count: 1, mean, m2 };
406            assert!(WelfordState::decode(&non_finite.encode()).is_err());
407        }
408    }
409
410    #[test]
411    fn test_welford_state_merge_matches_one_pass_update() {
412        let mut merged = state_from_values(&[1.0, 2.0]);
413        merged.merge(&state_from_values(&[3.0, 4.0])).unwrap();
414
415        assert_eq!(merged, state_from_values(&[1.0, 2.0, 3.0, 4.0]));
416    }
417
418    #[test]
419    fn test_welford_large_finite_variance_matches_partitioned_merge() {
420        let large_sample = f64::MAX.sqrt() * 1.1;
421        let one_pass = state_from_values(&[0.0, large_sample]);
422        let mut merged = state_from_values(&[0.0]);
423
424        merged.merge(&state_from_values(&[large_sample])).unwrap();
425
426        assert_eq!(merged, one_pass);
427    }
428
429    #[test]
430    fn test_welford_extreme_values_fail_independent_of_partitioning() {
431        for values in [[f64::MAX, -f64::MAX], [-f64::MAX, f64::MAX]] {
432            let mut one_pass = state_from_values(&values[..1]);
433            let original_one_pass = one_pass;
434            let update_failed = one_pass.update(values[1]).is_err();
435
436            let mut merged = state_from_values(&values[..1]);
437            let original_merged = merged;
438            let merge_failed = merged.merge(&state_from_values(&values[1..])).is_err();
439
440            assert_eq!((update_failed, merge_failed), (true, true));
441            assert_eq!(one_pass, original_one_pass);
442            assert_eq!(merged, original_merged);
443        }
444    }
445
446    #[test]
447    fn test_welford_state_empty_merge_identity() {
448        let populated = state_from_values(&[1.0, 2.0]);
449        let mut left = WelfordState::default();
450        left.merge(&populated).unwrap();
451        assert_eq!(left, populated);
452
453        let mut right = populated;
454        right.merge(&WelfordState::default()).unwrap();
455        assert_eq!(right, populated);
456    }
457
458    #[test]
459    fn test_welford_state_merge_rejects_count_overflow() {
460        let mut state = WelfordState {
461            count: u64::MAX,
462            mean: 1.0,
463            m2: 0.0,
464        };
465        let other = WelfordState {
466            count: 1,
467            mean: 1.0,
468            m2: 0.0,
469        };
470
471        assert!(state.merge(&other).is_err());
472    }
473
474    #[test]
475    fn test_welford_accumulator_ignores_nulls() {
476        let mut accumulator = WelfordAccumulator::default();
477        let array = Arc::new(Float64Array::from(vec![Some(1.0), None, Some(3.0)])) as ArrayRef;
478
479        accumulator.update_batch(&[array]).unwrap();
480
481        let ScalarValue::Binary(Some(encoded)) = accumulator.evaluate().unwrap() else {
482            panic!("Expected binary scalar value");
483        };
484        assert_eq!(
485            WelfordState::decode(&encoded).unwrap(),
486            state_from_values(&[1.0, 3.0])
487        );
488    }
489
490    #[test]
491    fn test_welford_accumulator_merges_binary_states() {
492        let first = state_from_values(&[1.0, 2.0]).encode();
493        let second = state_from_values(&[3.0, 4.0]).encode();
494        let states = Arc::new(BinaryArray::from(vec![
495            Some(first.as_slice()),
496            None,
497            Some(second.as_slice()),
498        ])) as ArrayRef;
499        let mut accumulator = WelfordAccumulator::default();
500
501        accumulator.merge_batch(&[states]).unwrap();
502
503        let ScalarValue::Binary(Some(encoded)) = accumulator.state().unwrap().remove(0) else {
504            panic!("Expected binary scalar value");
505        };
506        assert_eq!(
507            WelfordState::decode(&encoded).unwrap(),
508            state_from_values(&[1.0, 2.0, 3.0, 4.0])
509        );
510    }
511
512    #[test]
513    fn test_welford_accumulator_rejects_malformed_state() {
514        let noncanonical_singleton = WelfordState {
515            count: 1,
516            mean: 0.0,
517            m2: 1.0,
518        }
519        .encode();
520
521        for encoded in [b"invalid".to_vec(), noncanonical_singleton.to_vec()] {
522            let states = Arc::new(BinaryArray::from(vec![Some(encoded.as_slice())])) as ArrayRef;
523            let mut accumulator = WelfordAccumulator::default();
524
525            assert!(accumulator.merge_batch(&[states]).is_err());
526        }
527    }
528}