fix: make vector aggregates work with GROUP BY and partial aggregation (#9338)

* fix: make vector aggregates work with GROUP BY and partial aggregation

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

* test: make partitioned vec_avg case distinguish weighted averages

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

---------

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>
This commit is contained in:
dennis zhuang
2026-09-24 08:45:40 +00:00
committed by GitHub
parent 822042c7bb
commit 357c935804
5 changed files with 321 additions and 72 deletions
+3 -3
View File
@@ -62,14 +62,14 @@ impl VectorAvg {
}
fn accumulator(args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
if args.schema.fields().len() != 1 {
if args.exprs.len() != 1 {
return Err(datafusion_common::DataFusionError::Internal(format!(
"expect creating `VEC_AVG` with only one input field, actual {}",
args.schema.fields().len()
args.exprs.len()
)));
}
let t = args.schema.field(0).data_type();
let t = args.expr_fields[0].data_type();
if !matches!(t, DataType::Utf8 | DataType::LargeUtf8 | DataType::Binary) {
return Err(datafusion_common::DataFusionError::Internal(format!(
"unexpected input datatype {t} when creating `VEC_AVG`"
+94 -29
View File
@@ -28,6 +28,9 @@ use crate::scalars::vector::impl_conv::{
};
/// Aggregates by multiplying elements across the same dimension, returns a vector.
///
/// The result is NULL if any input vector is NULL, so the partial state carries a
/// `has_null` flag: a NULL `product` alone can't tell a NULL input from an empty partition.
#[derive(Debug, Default)]
pub struct VectorProduct {
product: Option<OVector<f32, Dyn>>,
@@ -49,20 +52,23 @@ impl VectorProduct {
signature,
DataType::Binary,
Arc::new(Self::accumulator),
vec![Arc::new(Field::new("x", DataType::Binary, true))],
vec![
Arc::new(Field::new("product", DataType::Binary, true)),
Arc::new(Field::new("has_null", DataType::Boolean, true)),
],
);
AggregateUDF::from(udaf)
}
fn accumulator(args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
if args.schema.fields().len() != 1 {
if args.exprs.len() != 1 {
return Err(datafusion_common::DataFusionError::Internal(format!(
"expect creating `VEC_PRODUCT` with only one input field, actual {}",
args.schema.fields().len()
args.exprs.len()
)));
}
let t = args.schema.field(0).data_type();
let t = args.expr_fields[0].data_type();
if !matches!(t, DataType::Utf8 | DataType::LargeUtf8 | DataType::Binary) {
return Err(datafusion_common::DataFusionError::Internal(format!(
"unexpected input datatype {t} when creating `VEC_PRODUCT`"
@@ -72,13 +78,33 @@ impl VectorProduct {
Ok(Box::new(VectorProduct::default()))
}
fn inner(&mut self, len: usize) -> &mut OVector<f32, Dyn> {
self.product.get_or_insert_with(|| {
OVector::from_iterator_generic(Dyn(len), Const::<1>, (0..len).map(|_| 1.0))
})
fn mul(&mut self, vector: &[f32]) {
let vector = DVectorView::from_slice(vector, vector.len());
let product = self.product.get_or_insert_with(|| {
OVector::from_iterator_generic(
Dyn(vector.len()),
Const::<1>,
(0..vector.len()).map(|_| 1.0),
)
});
*product = product.component_mul(&vector);
}
fn update(&mut self, values: &[ArrayRef], is_update: bool) -> Result<()> {
fn set_null(&mut self) {
self.has_null = true;
self.product = None;
}
}
impl Accumulator for VectorProduct {
fn state(&mut self) -> Result<Vec<ScalarValue>> {
Ok(vec![
self.evaluate()?,
ScalarValue::Boolean(Some(self.has_null)),
])
}
fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
if values.is_empty() || self.has_null {
return Ok(());
};
@@ -112,33 +138,34 @@ impl VectorProduct {
}
};
if vectors.len() != values[0].len() {
if is_update {
self.has_null = true;
self.product = None;
}
self.set_null();
return Ok(());
}
vectors.iter().for_each(|v| {
let v = DVectorView::from_slice(v, v.len());
let inner = self.inner(v.len());
*inner = inner.component_mul(&v);
});
vectors.iter().for_each(|v| self.mul(v));
Ok(())
}
}
impl Accumulator for VectorProduct {
fn state(&mut self) -> Result<Vec<ScalarValue>> {
self.evaluate().map(|v| vec![v])
}
fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
self.update(values, true)
}
fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
self.update(states, false)
let [products, has_nulls] = states else {
return Err(datafusion_common::DataFusionError::Internal(format!(
"expect 2 states for `VEC_PRODUCT`, actual {}",
states.len()
)));
};
if self.has_null {
return Ok(());
}
if has_nulls.as_boolean().true_count() > 0 {
self.set_null();
return Ok(());
}
// A NULL product without `has_null` comes from a partition without input rows.
for b in products.as_binary::<i32>().iter().flatten() {
self.mul(&binlit_as_veclit(b)?);
}
Ok(())
}
fn evaluate(&mut self) -> Result<ScalarValue> {
@@ -225,4 +252,42 @@ mod tests {
vec_product.evaluate().unwrap()
);
}
#[test]
fn test_merge_batch() {
let partial = |v: Option<&str>| {
let mut acc = VectorProduct::default();
let v: ArrayRef = Arc::new(StringArray::from(vec![v]));
acc.update_batch(&[v]).unwrap();
acc.state().unwrap()
};
let states = |states: Vec<Vec<ScalarValue>>| -> Vec<ArrayRef> {
(0..2)
.map(|i| ScalarValue::iter_to_array(states.iter().map(|s| s[i].clone())).unwrap())
.collect()
};
// An empty partition in the middle of the batch must not stop the merge.
let mut merged = VectorProduct::default();
merged
.merge_batch(&states(vec![
partial(Some("[1.0,2.0]")),
VectorProduct::default().state().unwrap(),
partial(Some("[3.0,4.0]")),
]))
.unwrap();
assert_eq!(
ScalarValue::Binary(Some(veclit_to_binlit(&[3.0, 8.0]))),
merged.evaluate().unwrap()
);
// A NULL input in any partition makes the result NULL.
merged
.merge_batch(&states(vec![partial(Some("[1.0,2.0]")), partial(None)]))
.unwrap();
merged
.merge_batch(&states(vec![partial(Some("[3.0,4.0]"))]))
.unwrap();
assert_eq!(ScalarValue::Binary(None), merged.evaluate().unwrap());
}
}
+93 -40
View File
@@ -28,6 +28,9 @@ use crate::scalars::vector::impl_conv::{
};
/// The accumulator for the `vec_sum` aggregate function.
///
/// The result is NULL if any input vector is NULL, so the partial state carries a
/// `has_null` flag: a NULL `sum` alone can't tell a NULL input from an empty partition.
#[derive(Debug, Default)]
pub struct VectorSum {
sum: Option<OVector<f32, Dyn>>,
@@ -49,20 +52,23 @@ impl VectorSum {
signature,
DataType::Binary,
Arc::new(Self::accumulator),
vec![Arc::new(Field::new("x", DataType::Binary, true))],
vec![
Arc::new(Field::new("sum", DataType::Binary, true)),
Arc::new(Field::new("has_null", DataType::Boolean, true)),
],
);
AggregateUDF::from(udaf)
}
fn accumulator(args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
if args.schema.fields().len() != 1 {
if args.exprs.len() != 1 {
return Err(datafusion_common::DataFusionError::Internal(format!(
"expect creating `VEC_SUM` with only one input field, actual {}",
args.schema.fields().len()
args.exprs.len()
)));
}
let t = args.schema.field(0).data_type();
let t = args.expr_fields[0].data_type();
if !matches!(t, DataType::Utf8 | DataType::LargeUtf8 | DataType::Binary) {
return Err(datafusion_common::DataFusionError::Internal(format!(
"unexpected input datatype {t} when creating `VEC_SUM`"
@@ -72,12 +78,28 @@ impl VectorSum {
Ok(Box::new(VectorSum::default()))
}
fn inner(&mut self, len: usize) -> &mut OVector<f32, Dyn> {
self.sum
.get_or_insert_with(|| OVector::zeros_generic(Dyn(len), Const::<1>))
fn add(&mut self, vector: &[f32]) {
let vector = DVectorView::from_slice(vector, vector.len());
*self
.sum
.get_or_insert_with(|| OVector::zeros_generic(Dyn(vector.len()), Const::<1>)) += vector;
}
fn update(&mut self, values: &[ArrayRef], is_update: bool) -> Result<()> {
fn set_null(&mut self) {
self.has_null = true;
self.sum = None;
}
}
impl Accumulator for VectorSum {
fn state(&mut self) -> Result<Vec<ScalarValue>> {
Ok(vec![
self.evaluate()?,
ScalarValue::Boolean(Some(self.has_null)),
])
}
fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
if values.is_empty() || self.has_null {
return Ok(());
};
@@ -87,45 +109,30 @@ impl VectorSum {
let arr: &StringArray = values[0].as_string();
for s in arr.iter() {
let Some(s) = s else {
if is_update {
self.has_null = true;
self.sum = None;
}
self.set_null();
return Ok(());
};
let values = parse_veclit_from_strlit(s)?;
let vec_column = DVectorView::from_slice(&values, values.len());
*self.inner(vec_column.len()) += vec_column;
self.add(&parse_veclit_from_strlit(s)?);
}
}
DataType::LargeUtf8 => {
let arr: &LargeStringArray = values[0].as_string();
for s in arr.iter() {
let Some(s) = s else {
if is_update {
self.has_null = true;
self.sum = None;
}
self.set_null();
return Ok(());
};
let values = parse_veclit_from_strlit(s)?;
let vec_column = DVectorView::from_slice(&values, values.len());
*self.inner(vec_column.len()) += vec_column;
self.add(&parse_veclit_from_strlit(s)?);
}
}
DataType::Binary => {
let arr: &BinaryArray = values[0].as_binary();
for b in arr.iter() {
let Some(b) = b else {
if is_update {
self.has_null = true;
self.sum = None;
}
self.set_null();
return Ok(());
};
let values = binlit_as_veclit(b)?;
let vec_column = DVectorView::from_slice(&values, values.len());
*self.inner(vec_column.len()) += vec_column;
self.add(&binlit_as_veclit(b)?);
}
}
_ => {
@@ -137,19 +144,27 @@ impl VectorSum {
}
Ok(())
}
}
impl Accumulator for VectorSum {
fn state(&mut self) -> Result<Vec<ScalarValue>> {
self.evaluate().map(|v| vec![v])
}
fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
self.update(values, true)
}
fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
self.update(states, false)
let [sums, has_nulls] = states else {
return Err(datafusion_common::DataFusionError::Internal(format!(
"expect 2 states for `VEC_SUM`, actual {}",
states.len()
)));
};
if self.has_null {
return Ok(());
}
if has_nulls.as_boolean().true_count() > 0 {
self.set_null();
return Ok(());
}
// A NULL sum without `has_null` comes from a partition without input rows.
for b in sums.as_binary::<i32>().iter().flatten() {
self.add(&binlit_as_veclit(b)?);
}
Ok(())
}
fn evaluate(&mut self) -> Result<ScalarValue> {
@@ -236,4 +251,42 @@ mod tests {
vec_sum.evaluate().unwrap()
);
}
#[test]
fn test_merge_batch() {
let partial = |v: Option<&str>| {
let mut acc = VectorSum::default();
let v: ArrayRef = Arc::new(StringArray::from(vec![v]));
acc.update_batch(&[v]).unwrap();
acc.state().unwrap()
};
let states = |states: Vec<Vec<ScalarValue>>| -> Vec<ArrayRef> {
(0..2)
.map(|i| ScalarValue::iter_to_array(states.iter().map(|s| s[i].clone())).unwrap())
.collect()
};
// An empty partition in the middle of the batch must not stop the merge.
let mut merged = VectorSum::default();
merged
.merge_batch(&states(vec![
partial(Some("[1.0,2.0]")),
VectorSum::default().state().unwrap(),
partial(Some("[3.0,4.0]")),
]))
.unwrap();
assert_eq!(
ScalarValue::Binary(Some(veclit_to_binlit(&[4.0, 6.0]))),
merged.evaluate().unwrap()
);
// A NULL input in any partition makes the result NULL.
merged
.merge_batch(&states(vec![partial(Some("[1.0,2.0]")), partial(None)]))
.unwrap();
merged
.merge_batch(&states(vec![partial(Some("[3.0,4.0]"))]))
.unwrap();
assert_eq!(ScalarValue::Binary(None), merged.evaluate().unwrap());
}
}
@@ -448,3 +448,89 @@ FROM (
| [4,5,6,10,-8] |
+---------------------------------------------------+
SELECT h, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM (
SELECT 'a' AS h, '[1.0, 1.0]' AS v
UNION ALL
SELECT 'a' AS h, '[2.0, 2.0]' AS v
UNION ALL
SELECT 'b' AS h, '[3.0, 3.0]' AS v
) GROUP BY h ORDER BY h;
+---+---------------------------+---------------------------+-------------------------------+
| h | vec_to_string(vec_sum(v)) | vec_to_string(vec_avg(v)) | vec_to_string(vec_product(v)) |
+---+---------------------------+---------------------------+-------------------------------+
| a | [3,3] | [1.5,1.5] | [2,2] |
| b | [3,3] | [3,3] | [3,3] |
+---+---------------------------+---------------------------+-------------------------------+
-- On partitioned tables the aggregates are split into partial state and merge.
-- Rows only land in two of the three partitions, the third one has an empty state.
-- The two non-empty partitions differ in row count and mean, so vec_avg must weight by count.
CREATE TABLE vector_aggr_partitioned (
ts TIMESTAMP TIME INDEX,
k INT,
g STRING,
v VECTOR(2),
PRIMARY KEY(k)
)
PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20);
Affected Rows: 0
INSERT INTO vector_aggr_partitioned VALUES
(1000, 1, 'a', '[1.0, 1.0]'),
(2000, 11, 'a', '[8.0, 8.0]'),
(3000, 2, 'a', '[3.0, 3.0]'),
(3500, 3, 'a', '[5.0, 5.0]'),
(4000, 12, 'b', '[4.0, 4.0]');
Affected Rows: 5
SELECT vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned;
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| vec_to_string(vec_sum(vector_aggr_partitioned.v)) | vec_to_string(vec_avg(vector_aggr_partitioned.v)) | vec_to_string(vec_product(vector_aggr_partitioned.v)) |
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| [21,21] | [4.2,4.2] | [480,480] |
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
SELECT g, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned GROUP BY g ORDER BY g;
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| g | vec_to_string(vec_sum(vector_aggr_partitioned.v)) | vec_to_string(vec_avg(vector_aggr_partitioned.v)) | vec_to_string(vec_product(vector_aggr_partitioned.v)) |
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| a | [17,17] | [4.25,4.25] | [120,120] |
| b | [4,4] | [4,4] | [4,4] |
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
-- A NULL vector makes vec_sum and vec_product NULL, vec_avg skips it.
INSERT INTO vector_aggr_partitioned VALUES (5000, 21, 'b', NULL);
Affected Rows: 1
SELECT vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned;
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| vec_to_string(vec_sum(vector_aggr_partitioned.v)) | vec_to_string(vec_avg(vector_aggr_partitioned.v)) | vec_to_string(vec_product(vector_aggr_partitioned.v)) |
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| | [4.2,4.2] | |
+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
SELECT g, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned GROUP BY g ORDER BY g;
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| g | vec_to_string(vec_sum(vector_aggr_partitioned.v)) | vec_to_string(vec_avg(vector_aggr_partitioned.v)) | vec_to_string(vec_product(vector_aggr_partitioned.v)) |
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
| a | [17,17] | [4.25,4.25] | [120,120] |
| b | | [4,4] | |
+---+---------------------------------------------------+---------------------------------------------------+-------------------------------------------------------+
DROP TABLE vector_aggr_partitioned;
Affected Rows: 0
@@ -150,3 +150,48 @@ FROM (
UNION ALL
SELECT '[4.0, 5.0, 6.0, 10, -8, 100]' AS v
) ORDER BY v;
SELECT h, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM (
SELECT 'a' AS h, '[1.0, 1.0]' AS v
UNION ALL
SELECT 'a' AS h, '[2.0, 2.0]' AS v
UNION ALL
SELECT 'b' AS h, '[3.0, 3.0]' AS v
) GROUP BY h ORDER BY h;
-- On partitioned tables the aggregates are split into partial state and merge.
-- Rows only land in two of the three partitions, the third one has an empty state.
-- The two non-empty partitions differ in row count and mean, so vec_avg must weight by count.
CREATE TABLE vector_aggr_partitioned (
ts TIMESTAMP TIME INDEX,
k INT,
g STRING,
v VECTOR(2),
PRIMARY KEY(k)
)
PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20);
INSERT INTO vector_aggr_partitioned VALUES
(1000, 1, 'a', '[1.0, 1.0]'),
(2000, 11, 'a', '[8.0, 8.0]'),
(3000, 2, 'a', '[3.0, 3.0]'),
(3500, 3, 'a', '[5.0, 5.0]'),
(4000, 12, 'b', '[4.0, 4.0]');
SELECT vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned;
SELECT g, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned GROUP BY g ORDER BY g;
-- A NULL vector makes vec_sum and vec_product NULL, vec_avg skips it.
INSERT INTO vector_aggr_partitioned VALUES (5000, 21, 'b', NULL);
SELECT vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned;
SELECT g, vec_to_string(vec_sum(v)), vec_to_string(vec_avg(v)), vec_to_string(vec_product(v))
FROM vector_aggr_partitioned GROUP BY g ORDER BY g;
DROP TABLE vector_aggr_partitioned;