Skip to content

Commit 640d802

Browse files
authored
Add UnionArray compute functions (#8884)
## Rationale for this change Tracking issue: #7882 This PR adds the straightforward compute support for canonical sparse Union arrays while keeping operations with unresolved nullability or placeholder semantics in focused follow-ups. ## What changes are included in this PR? - Adds canonical filter execution and slice and mask reductions for UnionArray, following the same structural execution patterns as neighboring canonical dtypes. - Adds validity-mask execution while preserving row alignment across type IDs and sparse children. - Recursively compresses the type IDs and every sparse child. - Computes uncompressed size as the checked sum of the type IDs and sparse children. - Adds focused coverage for structural operations, outer null masking, and size accounting. ## Deferred follow-ups - Take, including nullable-index outer-null propagation, and dictionary execution that depends on it. - Union casts and outer-nullability conformance coverage. - Constant Union canonicalization and inactive-child placeholder construction. - Chunked Union canonicalization, including a representation for empty chunked unions. The source TODOs record the required semantics at each deferred dispatch point. Signed-off-by: Connor Tsui <connor.tsui20@gmail.com>
1 parent 2794f89 commit 640d802

18 files changed

Lines changed: 293 additions & 32 deletions

File tree

vortex-array/src/aggregate_fn/fns/uncompressed_size_in_bytes/mod.rs

Lines changed: 31 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ mod list_view;
99
mod null;
1010
mod primitive;
1111
mod struct_;
12+
mod union;
1213
mod varbinview;
1314

1415
use std::mem::size_of;
@@ -21,6 +22,7 @@ use list_view::list_view_uncompressed_size_in_bytes;
2122
use null::null_uncompressed_size_in_bytes;
2223
use primitive::primitive_uncompressed_size_in_bytes;
2324
use struct_::struct_uncompressed_size_in_bytes;
25+
use union::union_uncompressed_size_in_bytes;
2426
use varbinview::varbinview_uncompressed_size_in_bytes;
2527
use vortex_error::VortexExpect;
2628
use vortex_error::VortexResult;
@@ -199,9 +201,7 @@ pub(crate) fn canonical_uncompressed_size_in_bytes(
199201
Canonical::List(array) => list_view_uncompressed_size_in_bytes(array, ctx),
200202
Canonical::FixedSizeList(array) => fixed_size_list_uncompressed_size_in_bytes(array, ctx),
201203
Canonical::Struct(array) => struct_uncompressed_size_in_bytes(array, ctx),
202-
Canonical::Union(_) => {
203-
todo!("TODO(connor)[Union]: implement UncompressedSizeInBytes for Union arrays")
204-
}
204+
Canonical::Union(array) => union_uncompressed_size_in_bytes(array, ctx),
205205
Canonical::Extension(array) => extension_uncompressed_size_in_bytes(array, ctx),
206206
Canonical::Variant(_) => {
207207
vortex_bail!("UncompressedSizeInBytes is not supported for Variant arrays")
@@ -236,7 +236,12 @@ pub(crate) fn constant_uncompressed_size_in_bytes(
236236
let canonical = array.array().clone().execute::<Canonical>(ctx)?;
237237
return canonical_uncompressed_size_in_bytes(&canonical, ctx);
238238
}
239-
DType::Union(..) => todo!("TODO(connor)[Union]: unimplemented"),
239+
DType::Union(..) => {
240+
todo!(
241+
"TODO(connor)[Union]: support constant Union size accounting after constant Union \
242+
canonicalization defines inactive sparse-child placeholders"
243+
)
244+
}
240245
DType::Variant(_) => {
241246
vortex_bail!("UncompressedSizeInBytes is not supported for Variant arrays")
242247
}
@@ -342,6 +347,7 @@ mod tests {
342347
use crate::arrays::NullArray;
343348
use crate::arrays::PrimitiveArray;
344349
use crate::arrays::StructArray;
350+
use crate::arrays::UnionArray;
345351
use crate::arrays::VarBinViewArray;
346352
use crate::arrays::VariantArray;
347353
use crate::builders::builder_with_capacity;
@@ -350,6 +356,7 @@ mod tests {
350356
use crate::dtype::FieldNames;
351357
use crate::dtype::Nullability;
352358
use crate::dtype::PType;
359+
use crate::dtype::UnionVariants;
353360
use crate::expr::stats::Precision;
354361
use crate::expr::stats::Stat;
355362
use crate::expr::stats::StatsProvider;
@@ -539,6 +546,26 @@ mod tests {
539546
Ok(())
540547
}
541548

549+
#[test]
550+
fn union_sums_type_ids_and_sparse_children() -> VortexResult<()> {
551+
let type_ids = PrimitiveArray::from_iter([5_u8, 9, 5]).into_array();
552+
let numbers = PrimitiveArray::from_iter([10_i32, 0, 30]).into_array();
553+
let flags = BoolArray::from_iter([false, true, false]).into_array();
554+
let expected = aggregate(&type_ids)? + aggregate(&numbers)? + aggregate(&flags)?;
555+
let variants = UnionVariants::try_new(
556+
["number", "flag"].into(),
557+
vec![
558+
DType::Primitive(PType::I32, Nullability::NonNullable),
559+
DType::Bool(Nullability::NonNullable),
560+
],
561+
vec![5, 9],
562+
)?;
563+
let array = UnionArray::try_new(type_ids, variants, vec![numbers, flags])?.into_array();
564+
565+
assert_eq!(aggregate(&array)?, expected);
566+
Ok(())
567+
}
568+
542569
#[test]
543570
fn extension_matches_materialized_size() -> VortexResult<()> {
544571
let storage = PrimitiveArray::from_option_iter([Some(1i32), None, Some(3)]).into_array();
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
// SPDX-FileCopyrightText: Copyright the Vortex contributors
3+
4+
use vortex_error::VortexResult;
5+
use vortex_error::vortex_err;
6+
7+
use super::uncompressed_size_in_bytes_u64;
8+
use crate::ExecutionCtx;
9+
use crate::arrays::UnionArray;
10+
use crate::arrays::union::UnionArrayExt;
11+
12+
pub(super) fn union_uncompressed_size_in_bytes(
13+
array: &UnionArray,
14+
ctx: &mut ExecutionCtx,
15+
) -> VortexResult<u64> {
16+
let mut size = uncompressed_size_in_bytes_u64(array.type_ids(), ctx)?;
17+
18+
for child in array.iter_children() {
19+
size = size
20+
.checked_add(uncompressed_size_in_bytes_u64(child, ctx)?)
21+
.ok_or_else(|| vortex_err!("uncompressed size in bytes overflowed u64"))?;
22+
}
23+
24+
Ok(size)
25+
}

vortex-array/src/arrays/chunked/vtable/mod.rs

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -251,6 +251,12 @@ impl VTable for Chunked {
251251

252252
fn execute(array: Array<Self>, ctx: &mut ExecutionCtx) -> VortexResult<ExecutionResult> {
253253
match array.dtype() {
254+
DType::Union(..) => {
255+
todo!(
256+
"TODO(connor)[Union]: canonicalize chunked Union arrays by packing type IDs and \
257+
every sparse child along identical chunk boundaries"
258+
)
259+
}
254260
// Struct, List, FixedSizeList, and Variant need child swizzling that the builder path
255261
// cannot express.
256262
DType::Struct(..) | DType::List(..) | DType::FixedSizeList(..) | DType::Variant(..) => {

vortex-array/src/arrays/constant/vtable/canonical.rs

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -164,7 +164,13 @@ pub(crate) fn constant_canonicalize(
164164
StructArray::new_unchecked(fields, struct_dtype.clone(), array.len(), validity)
165165
})
166166
}
167-
DType::Union(..) => todo!("TODO(connor)[Union]: unimplemented"),
167+
DType::Union(..) => {
168+
todo!(
169+
"TODO(connor)[Union]: canonicalize constant Union arrays in a focused follow-up \
170+
after defining placeholder values for every inactive sparse child, including \
171+
nested Struct and Union variants"
172+
)
173+
}
168174
DType::Variant(_) => Canonical::Variant(VariantArray::try_new(
169175
array.array().clone().into_array(),
170176
None,

vortex-array/src/arrays/dict/execute.rs

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,10 @@ pub(crate) fn take_canonical(
5454
}
5555
Canonical::Struct(a) => Canonical::Struct(take_struct(&a, codes)),
5656
Canonical::Union(_) => {
57-
todo!("TODO(connor)[Union]: implement dictionary execution for Union arrays")
57+
todo!(
58+
"TODO(connor)[Union]: implement dictionary execution after Union take supports \
59+
nullable indices and outer null propagation"
60+
)
5861
}
5962
Canonical::Extension(a) => Canonical::Extension(take_extension(&a, codes, ctx)),
6063
Canonical::Variant(a) => {

vortex-array/src/arrays/filter/execute/mod.rs

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -28,17 +28,21 @@ use crate::arrays::variant::VariantArrayExt;
2828
use crate::scalar::Scalar;
2929
use crate::validity::Validity;
3030

31+
pub(crate) mod byte_compress;
32+
33+
mod slice;
34+
mod take;
35+
3136
mod bitbuffer;
32-
mod bool;
3337
mod buffer;
34-
pub(crate) mod byte_compress;
38+
39+
mod bool;
3540
mod decimal;
3641
mod fixed_size_list;
3742
mod listview;
3843
mod primitive;
39-
mod slice;
4044
mod struct_;
41-
pub mod take;
45+
mod union;
4246
mod varbinview;
4347

4448
/// A helper function that lazily filters a [`Validity`] with selection mask values.
@@ -95,9 +99,7 @@ pub(super) fn execute_filter(canonical: Canonical, mask: &Arc<MaskValues>) -> Ca
9599
Canonical::FixedSizeList(fixed_size_list::filter_fixed_size_list(&a, mask))
96100
}
97101
Canonical::Struct(a) => Canonical::Struct(struct_::filter_struct(&a, mask)),
98-
Canonical::Union(_) => {
99-
todo!("TODO(connor)[Union]: implement filter for Union arrays")
100-
}
102+
Canonical::Union(a) => Canonical::Union(union::filter_union(&a, mask)),
101103
Canonical::Extension(a) => {
102104
let filtered_storage = a
103105
.storage_array()
Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
// SPDX-FileCopyrightText: Copyright the Vortex contributors
3+
4+
use std::sync::Arc;
5+
6+
use vortex_error::VortexExpect;
7+
use vortex_mask::Mask;
8+
use vortex_mask::MaskValues;
9+
10+
use crate::ArrayRef;
11+
use crate::arrays::UnionArray;
12+
use crate::arrays::union::UnionArrayExt;
13+
14+
pub fn filter_union(array: &UnionArray, mask: &Arc<MaskValues>) -> UnionArray {
15+
let filter_mask = Mask::Values(Arc::clone(mask));
16+
17+
let type_ids = array
18+
.type_ids()
19+
.filter(filter_mask.clone())
20+
.vortex_expect("UnionArray type IDs are guaranteed to support filter");
21+
22+
let children: Vec<ArrayRef> = array
23+
.iter_children()
24+
.map(|child| {
25+
child
26+
.filter(filter_mask.clone())
27+
.vortex_expect("UnionArray children are guaranteed to support filter")
28+
})
29+
.collect();
30+
31+
UnionArray::try_new(type_ids, array.variants().clone(), children)
32+
.vortex_expect("filtered UnionArray children have consistent dtypes and lengths")
33+
}

vortex-array/src/arrays/masked/execute.rs

Lines changed: 22 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -17,13 +17,15 @@ use crate::arrays::ListViewArray;
1717
use crate::arrays::MaskedArray;
1818
use crate::arrays::PrimitiveArray;
1919
use crate::arrays::StructArray;
20+
use crate::arrays::UnionArray;
2021
use crate::arrays::VarBinViewArray;
2122
use crate::arrays::VariantArray;
2223
use crate::arrays::bool::BoolArrayExt;
2324
use crate::arrays::extension::ExtensionArrayExt;
2425
use crate::arrays::fixed_size_list::FixedSizeListArrayExt;
2526
use crate::arrays::listview::ListViewArrayExt;
2627
use crate::arrays::struct_::StructArrayExt;
28+
use crate::arrays::union::UnionArrayExt;
2729
use crate::arrays::variant::VariantArrayExt;
2830
use crate::builtins::ArrayBuiltins;
2931
use crate::executor::ExecutionCtx;
@@ -50,9 +52,7 @@ pub fn mask_validity_canonical(
5052
Canonical::FixedSizeList(mask_validity_fixed_size_list(a, validity)?)
5153
}
5254
Canonical::Struct(a) => Canonical::Struct(mask_validity_struct(a, validity)?),
53-
Canonical::Union(_) => {
54-
todo!("TODO(connor)[Union]: implement masking for Union arrays")
55-
}
55+
Canonical::Union(a) => Canonical::Union(mask_validity_union(a, validity)?),
5656
Canonical::Extension(a) => Canonical::Extension(mask_validity_extension(a, validity, ctx)?),
5757
Canonical::Variant(a) => Canonical::Variant(mask_validity_variant(a, validity, ctx)?),
5858
})
@@ -69,7 +69,7 @@ fn mask_validity_primitive(
6969
) -> VortexResult<PrimitiveArray> {
7070
let ptype = array.ptype();
7171
let new_validity = Validity::and(array.validity()?, validity)?;
72-
// SAFETY: validity has same length as values
72+
// SAFETY: We're only changing validity, not the data structure.
7373
Ok(unsafe {
7474
PrimitiveArray::new_unchecked_from_handle(
7575
array.buffer_handle().clone(),
@@ -81,7 +81,7 @@ fn mask_validity_primitive(
8181

8282
fn mask_validity_decimal(array: DecimalArray, validity: Validity) -> VortexResult<DecimalArray> {
8383
let new_validity = Validity::and(array.validity()?, validity)?;
84-
// SAFETY: We're only changing validity, not the data structure
84+
// SAFETY: We're only changing validity, not the data structure.
8585
Ok(unsafe {
8686
DecimalArray::new_unchecked_handle(
8787
array.buffer_handle().clone(),
@@ -99,7 +99,7 @@ fn mask_validity_varbinview(
9999
) -> VortexResult<VarBinViewArray> {
100100
let dtype = array.dtype().as_nullable();
101101
let new_validity = Validity::and(array.validity()?, validity)?;
102-
// SAFETY: We're only changing validity, not the data structure
102+
// SAFETY: We're only changing validity, not the data structure.
103103
Ok(unsafe {
104104
VarBinViewArray::new_handle_unchecked(
105105
array.views_handle().clone(),
@@ -112,7 +112,7 @@ fn mask_validity_varbinview(
112112

113113
fn mask_validity_listview(array: ListViewArray, validity: Validity) -> VortexResult<ListViewArray> {
114114
let new_validity = Validity::and(array.validity()?, validity)?;
115-
// SAFETY: We're only changing validity, not the data structure
115+
// SAFETY: We're only changing validity, not the data structure.
116116
let is_zctl = array.is_zero_copy_to_list();
117117
Ok(unsafe {
118118
ListViewArray::new_unchecked(
@@ -132,7 +132,7 @@ fn mask_validity_fixed_size_list(
132132
let len = array.len();
133133
let list_size = array.list_size();
134134
let new_validity = Validity::and(array.validity()?, validity)?;
135-
// SAFETY: We're only changing validity, not the data structure
135+
// SAFETY: We're only changing validity, not the data structure.
136136
Ok(unsafe {
137137
FixedSizeListArray::new_unchecked(array.elements().clone(), list_size, new_validity, len)
138138
})
@@ -143,16 +143,28 @@ fn mask_validity_struct(array: StructArray, validity: Validity) -> VortexResult<
143143
let new_validity = Validity::and(array.validity()?, validity)?;
144144
let fields = array.unmasked_fields();
145145
let struct_fields = array.struct_fields();
146-
// SAFETY: We're only changing validity, not the data structure
146+
// SAFETY: We're only changing validity, not the data structure.
147147
Ok(unsafe { StructArray::new_unchecked(fields, struct_fields.clone(), len, new_validity) })
148148
}
149149

150+
fn mask_validity_union(array: UnionArray, validity: Validity) -> VortexResult<UnionArray> {
151+
let type_ids = array
152+
.type_ids()
153+
.clone()
154+
.mask(validity.to_array(array.len()))?;
155+
let variants = array.variants().clone();
156+
let children = array.children();
157+
158+
// SAFETY: We're only changing validity, not the data structure.
159+
Ok(unsafe { UnionArray::new_unchecked(type_ids, variants, children) })
160+
}
161+
150162
fn mask_validity_extension(
151163
array: ExtensionArray,
152164
validity: Validity,
153165
ctx: &mut ExecutionCtx,
154166
) -> VortexResult<ExtensionArray> {
155-
// For extension arrays, we need to mask the underlying storage
167+
// For extension arrays, we need to mask the underlying storage.
156168
let storage = array.storage_array().clone().execute::<Canonical>(ctx)?;
157169
let masked_storage = mask_validity_canonical(storage, validity, ctx)?;
158170
let masked_storage = masked_storage.into_array();

vortex-array/src/arrays/union/array.rs

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ use crate::arrays::PrimitiveArray;
2020
use crate::arrays::Union;
2121
use crate::dtype::DType;
2222
use crate::dtype::Nullability;
23+
use crate::dtype::PType;
2324
use crate::dtype::UnionVariants;
2425

2526
/// The row-aligned array of type IDs selecting a union child.
@@ -144,10 +145,7 @@ impl Array<Union> {
144145
children: impl Into<Arc<[ArrayRef]>>,
145146
) -> VortexResult<Self> {
146147
vortex_ensure!(
147-
matches!(
148-
type_ids.dtype(),
149-
DType::Primitive(crate::dtype::PType::U8, _)
150-
),
148+
matches!(type_ids.dtype(), DType::Primitive(PType::U8, _)),
151149
"UnionArray type_ids must be u8, got {}",
152150
type_ids.dtype()
153151
);
Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
// SPDX-FileCopyrightText: Copyright the Vortex contributors
3+
4+
use vortex_error::VortexResult;
5+
6+
use crate::ArrayRef;
7+
use crate::IntoArray;
8+
use crate::array::ArrayView;
9+
use crate::arrays::Union;
10+
use crate::arrays::UnionArray;
11+
use crate::arrays::union::UnionArrayExt;
12+
use crate::builtins::ArrayBuiltins;
13+
use crate::scalar_fn::fns::mask::MaskReduce;
14+
15+
impl MaskReduce for Union {
16+
fn mask(array: ArrayView<'_, Union>, mask: &ArrayRef) -> VortexResult<Option<ArrayRef>> {
17+
UnionArray::try_new(
18+
array.type_ids().clone().mask(mask.clone())?,
19+
array.variants().clone(),
20+
array.children(),
21+
)
22+
.map(|a| Some(a.into_array()))
23+
}
24+
}

0 commit comments

Comments
 (0)