1use 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 return Ok(delta * weight * delta);
178 }
179
180 Err(non_finite_arithmetic())
181}
182
183#[derive(Debug, Default)]
185pub struct WelfordAccumulator {
186 state: WelfordState,
187}
188
189impl WelfordAccumulator {
190 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 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', 3, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 24, 64, ]
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}