Skip to main content

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