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