common_function/aggrs/vector/
sum.rs1use 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#[derive(Debug, Default)]
35pub struct VectorSum {
36 sum: Option<OVector<f32, Dyn>>,
37 has_null: bool,
38}
39
40impl VectorSum {
41 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 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 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 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 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 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 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 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 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 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}