Skip to main content

common_function/aggrs/vector/
sum.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
15use std::sync::Arc;
16
17use arrow::array::{Array, ArrayRef, AsArray, BinaryArray, LargeStringArray, StringArray};
18use arrow_schema::{DataType, Field};
19use datafusion_common::{Result, ScalarValue};
20use datafusion_expr::{
21    Accumulator, AggregateUDF, Signature, SimpleAggregateUDF, TypeSignature, Volatility,
22};
23use datafusion_functions_aggregate_common::accumulator::AccumulatorArgs;
24use nalgebra::{Const, DVectorView, Dyn, OVector};
25
26use crate::scalars::vector::impl_conv::{
27    binlit_as_veclit, parse_veclit_from_strlit, veclit_to_binlit,
28};
29
30/// The accumulator for the `vec_sum` aggregate function.
31///
32/// The result is NULL if any input vector is NULL, so the partial state carries a
33/// `has_null` flag: a NULL `sum` alone can't tell a NULL input from an empty partition.
34#[derive(Debug, Default)]
35pub struct VectorSum {
36    sum: Option<OVector<f32, Dyn>>,
37    has_null: bool,
38}
39
40impl VectorSum {
41    /// Create a new `AggregateUDF` for the `vec_sum` aggregate function.
42    pub fn uadf_impl() -> AggregateUDF {
43        let signature = Signature::one_of(
44            vec![
45                TypeSignature::Exact(vec![DataType::Utf8]),
46                TypeSignature::Exact(vec![DataType::Binary]),
47            ],
48            Volatility::Immutable,
49        );
50        let udaf = SimpleAggregateUDF::new_with_signature(
51            "vec_sum",
52            signature,
53            DataType::Binary,
54            Arc::new(Self::accumulator),
55            vec![
56                Arc::new(Field::new("sum", DataType::Binary, true)),
57                Arc::new(Field::new("has_null", DataType::Boolean, true)),
58            ],
59        );
60        AggregateUDF::from(udaf)
61    }
62
63    fn accumulator(args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
64        if args.exprs.len() != 1 {
65            return Err(datafusion_common::DataFusionError::Internal(format!(
66                "expect creating `VEC_SUM` with only one input field, actual {}",
67                args.exprs.len()
68            )));
69        }
70
71        let t = args.expr_fields[0].data_type();
72        if !matches!(t, DataType::Utf8 | DataType::LargeUtf8 | DataType::Binary) {
73            return Err(datafusion_common::DataFusionError::Internal(format!(
74                "unexpected input datatype {t} when creating `VEC_SUM`"
75            )));
76        }
77
78        Ok(Box::new(VectorSum::default()))
79    }
80
81    fn add(&mut self, vector: &[f32]) {
82        let vector = DVectorView::from_slice(vector, vector.len());
83        *self
84            .sum
85            .get_or_insert_with(|| OVector::zeros_generic(Dyn(vector.len()), Const::<1>)) += vector;
86    }
87
88    fn set_null(&mut self) {
89        self.has_null = true;
90        self.sum = None;
91    }
92}
93
94impl Accumulator for VectorSum {
95    fn state(&mut self) -> Result<Vec<ScalarValue>> {
96        Ok(vec![
97            self.evaluate()?,
98            ScalarValue::Boolean(Some(self.has_null)),
99        ])
100    }
101
102    fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
103        if values.is_empty() || self.has_null {
104            return Ok(());
105        };
106
107        match values[0].data_type() {
108            DataType::Utf8 => {
109                let arr: &StringArray = values[0].as_string();
110                for s in arr.iter() {
111                    let Some(s) = s else {
112                        self.set_null();
113                        return Ok(());
114                    };
115                    self.add(&parse_veclit_from_strlit(s)?);
116                }
117            }
118            DataType::LargeUtf8 => {
119                let arr: &LargeStringArray = values[0].as_string();
120                for s in arr.iter() {
121                    let Some(s) = s else {
122                        self.set_null();
123                        return Ok(());
124                    };
125                    self.add(&parse_veclit_from_strlit(s)?);
126                }
127            }
128            DataType::Binary => {
129                let arr: &BinaryArray = values[0].as_binary();
130                for b in arr.iter() {
131                    let Some(b) = b else {
132                        self.set_null();
133                        return Ok(());
134                    };
135                    self.add(&binlit_as_veclit(b)?);
136                }
137            }
138            _ => {
139                return Err(datafusion_common::DataFusionError::NotImplemented(format!(
140                    "unsupported data type {} for `VEC_SUM`",
141                    values[0].data_type()
142                )));
143            }
144        }
145        Ok(())
146    }
147
148    fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
149        let [sums, has_nulls] = states else {
150            return Err(datafusion_common::DataFusionError::Internal(format!(
151                "expect 2 states for `VEC_SUM`, actual {}",
152                states.len()
153            )));
154        };
155        if self.has_null {
156            return Ok(());
157        }
158        if has_nulls.as_boolean().true_count() > 0 {
159            self.set_null();
160            return Ok(());
161        }
162
163        // A NULL sum without `has_null` comes from a partition without input rows.
164        for b in sums.as_binary::<i32>().iter().flatten() {
165            self.add(&binlit_as_veclit(b)?);
166        }
167        Ok(())
168    }
169
170    fn evaluate(&mut self) -> Result<ScalarValue> {
171        match &self.sum {
172            None => Ok(ScalarValue::Binary(None)),
173            Some(vector) => Ok(ScalarValue::Binary(Some(veclit_to_binlit(
174                vector.as_slice(),
175            )))),
176        }
177    }
178
179    fn size(&self) -> usize {
180        size_of_val(self)
181    }
182}
183
184#[cfg(test)]
185mod tests {
186    use std::sync::Arc;
187
188    use arrow::array::StringArray;
189
190    use super::*;
191
192    #[test]
193    fn test_update_batch() {
194        // test update empty batch, expect not updating anything
195        let mut vec_sum = VectorSum::default();
196        vec_sum.update_batch(&[]).unwrap();
197        assert!(vec_sum.sum.is_none());
198        assert!(!vec_sum.has_null);
199        assert_eq!(ScalarValue::Binary(None), vec_sum.evaluate().unwrap());
200
201        // test update one not-null value
202        let mut vec_sum = VectorSum::default();
203        let v: Vec<ArrayRef> = vec![Arc::new(StringArray::from(vec![Some(
204            "[1.0,2.0,3.0]".to_string(),
205        )]))];
206        vec_sum.update_batch(&v).unwrap();
207        assert_eq!(
208            ScalarValue::Binary(Some(veclit_to_binlit(&[1.0, 2.0, 3.0]))),
209            vec_sum.evaluate().unwrap()
210        );
211
212        // test update one null value
213        let mut vec_sum = VectorSum::default();
214        let v: Vec<ArrayRef> = vec![Arc::new(StringArray::from(vec![Option::<String>::None]))];
215        vec_sum.update_batch(&v).unwrap();
216        assert_eq!(ScalarValue::Binary(None), vec_sum.evaluate().unwrap());
217
218        // test update no null-value batch
219        let mut vec_sum = VectorSum::default();
220        let v: Vec<ArrayRef> = vec![Arc::new(StringArray::from(vec![
221            Some("[1.0,2.0,3.0]".to_string()),
222            Some("[4.0,5.0,6.0]".to_string()),
223            Some("[7.0,8.0,9.0]".to_string()),
224        ]))];
225        vec_sum.update_batch(&v).unwrap();
226        assert_eq!(
227            ScalarValue::Binary(Some(veclit_to_binlit(&[12.0, 15.0, 18.0]))),
228            vec_sum.evaluate().unwrap()
229        );
230
231        // test update null-value batch
232        let mut vec_sum = VectorSum::default();
233        let v: Vec<ArrayRef> = vec![Arc::new(StringArray::from(vec![
234            Some("[1.0,2.0,3.0]".to_string()),
235            None,
236            Some("[7.0,8.0,9.0]".to_string()),
237        ]))];
238        vec_sum.update_batch(&v).unwrap();
239        assert_eq!(ScalarValue::Binary(None), vec_sum.evaluate().unwrap());
240
241        // test update with repeated values
242        let mut vec_sum = VectorSum::default();
243        let v = vec![
244            ScalarValue::Utf8(Some("[1.0,2.0,3.0]".to_string()))
245                .to_array_of_size(4)
246                .unwrap(),
247        ];
248        vec_sum.update_batch(&v).unwrap();
249        assert_eq!(
250            ScalarValue::Binary(Some(veclit_to_binlit(&[4.0, 8.0, 12.0]))),
251            vec_sum.evaluate().unwrap()
252        );
253    }
254
255    #[test]
256    fn test_merge_batch() {
257        let partial = |v: Option<&str>| {
258            let mut acc = VectorSum::default();
259            let v: ArrayRef = Arc::new(StringArray::from(vec![v]));
260            acc.update_batch(&[v]).unwrap();
261            acc.state().unwrap()
262        };
263        let states = |states: Vec<Vec<ScalarValue>>| -> Vec<ArrayRef> {
264            (0..2)
265                .map(|i| ScalarValue::iter_to_array(states.iter().map(|s| s[i].clone())).unwrap())
266                .collect()
267        };
268
269        // An empty partition in the middle of the batch must not stop the merge.
270        let mut merged = VectorSum::default();
271        merged
272            .merge_batch(&states(vec![
273                partial(Some("[1.0,2.0]")),
274                VectorSum::default().state().unwrap(),
275                partial(Some("[3.0,4.0]")),
276            ]))
277            .unwrap();
278        assert_eq!(
279            ScalarValue::Binary(Some(veclit_to_binlit(&[4.0, 6.0]))),
280            merged.evaluate().unwrap()
281        );
282
283        // A NULL input in any partition makes the result NULL.
284        merged
285            .merge_batch(&states(vec![partial(Some("[1.0,2.0]")), partial(None)]))
286            .unwrap();
287        merged
288            .merge_batch(&states(vec![partial(Some("[3.0,4.0]"))]))
289            .unwrap();
290        assert_eq!(ScalarValue::Binary(None), merged.evaluate().unwrap());
291    }
292}