Skip to content

Commit adb5991

Browse files
committed
Clean up UnionArray implementation
Signed-off-by: Connor Tsui <connor.tsui20@gmail.com>
1 parent 21e3100 commit adb5991

12 files changed

Lines changed: 234 additions & 211 deletions

File tree

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

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ use crate::arrays::chunked::ChunkedArrayExt;
2525
use crate::arrays::fixed_size_list::FixedSizeListArrayExt;
2626
use crate::arrays::listview::ListViewArrayExt;
2727
use crate::arrays::listview::ListViewRebuildMode;
28+
use crate::arrays::union::TYPE_IDS_DTYPE;
2829
use crate::arrays::union::UnionArrayExt;
2930
use crate::arrays::variant::VariantArrayExt;
3031
use crate::builders::builder_with_capacity_in;
@@ -104,14 +105,13 @@ fn pack_union_chunks(chunks: Vec<ArrayRef>, ctx: &mut ExecutionCtx) -> VortexRes
104105
.iter()
105106
.map(|chunk| chunk.type_ids().clone())
106107
.collect(),
107-
DType::Primitive(PType::I8, Nullability::NonNullable),
108+
TYPE_IDS_DTYPE,
108109
)?
109110
.into_array();
110-
let children = (0..variants.len())
111-
.map(|index| {
112-
let dtype = variants
113-
.variant_by_index(index)
114-
.vortex_expect("variant index must have a dtype");
111+
let children = variants
112+
.variants()
113+
.enumerate()
114+
.map(|(index, dtype)| {
115115
ChunkedArray::try_new(
116116
union_chunks
117117
.iter()

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

Lines changed: 47 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ use crate::builders::builder_with_capacity;
3232
use crate::dtype::DType;
3333
use crate::dtype::DecimalType;
3434
use crate::dtype::Nullability;
35+
use crate::dtype::UnionVariants;
3536
use crate::match_each_decimal_value;
3637
use crate::match_each_decimal_value_type;
3738
use crate::match_each_native_ptype;
@@ -167,47 +168,7 @@ pub(crate) fn constant_canonicalize(
167168
})
168169
}
169170
DType::Union(variants) => {
170-
if scalar.is_null() {
171-
vortex_bail!(
172-
"Canonicalizing a null union scalar is not supported until null semantics are defined"
173-
)
174-
}
175-
if variants.variants().any(|variant| variant.is_nullable()) {
176-
vortex_bail!("Canonical UnionArray children must be non-nullable")
177-
}
178-
179-
let union = scalar.as_union();
180-
let type_id = union
181-
.type_id()
182-
.vortex_expect("non-null union scalar must have a type ID");
183-
let child_index = union
184-
.child_index()
185-
.vortex_expect("validated union scalar must select a child");
186-
let selected_value = union
187-
.value()
188-
.vortex_expect("non-null union scalar must have a value");
189-
let children = variants
190-
.variants()
191-
.enumerate()
192-
.map(|(index, dtype)| {
193-
let value = if index == child_index {
194-
selected_value.clone()
195-
} else {
196-
Scalar::zero_value(&dtype)
197-
};
198-
ConstantArray::new(value, array.len()).into_array()
199-
})
200-
.collect::<Vec<_>>();
201-
202-
// SAFETY: The scalar's validated type ID selects `child_index`; all sparse children
203-
// have the declared dtype and the same length.
204-
Canonical::Union(unsafe {
205-
UnionArray::new_unchecked(
206-
ConstantArray::new(type_id, array.len()).into_array(),
207-
variants.clone(),
208-
children,
209-
)
210-
})
171+
Canonical::Union(constant_canonical_union(scalar, variants, array.len())?)
211172
}
212173
DType::Variant(_) => Canonical::Variant(VariantArray::try_new(
213174
array.array().clone().into_array(),
@@ -232,6 +193,51 @@ pub(crate) fn constant_canonicalize(
232193
})
233194
}
234195

196+
fn constant_canonical_union(
197+
scalar: &Scalar,
198+
variants: &UnionVariants,
199+
len: usize,
200+
) -> VortexResult<UnionArray> {
201+
if scalar.is_null() {
202+
vortex_bail!(
203+
"Canonicalizing a null union scalar is not supported until null semantics are defined"
204+
)
205+
}
206+
if variants.variants().any(|variant| variant.is_nullable()) {
207+
vortex_bail!("Canonical UnionArray children must be non-nullable")
208+
}
209+
210+
let union = scalar.as_union();
211+
let type_id = union
212+
.type_id()
213+
.vortex_expect("non-null union scalar must have a type ID");
214+
let selected_child = union
215+
.child_index()
216+
.vortex_expect("validated union scalar must select a child");
217+
let selected_value = union
218+
.value()
219+
.vortex_expect("non-null union scalar must have a value");
220+
221+
let children = variants
222+
.variants()
223+
.enumerate()
224+
.map(|(index, dtype)| {
225+
let value = if index == selected_child {
226+
selected_value.clone()
227+
} else {
228+
Scalar::zero_value(&dtype)
229+
};
230+
ConstantArray::new(value, len).into_array()
231+
})
232+
.collect::<Vec<_>>();
233+
234+
UnionArray::try_new(
235+
ConstantArray::new(type_id, len).into_array(),
236+
variants.clone(),
237+
children,
238+
)
239+
}
240+
235241
fn constant_canonical_byte_view(
236242
scalar_bytes: Option<&[u8]>,
237243
dtype: &DType,

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

Lines changed: 5 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -25,13 +25,12 @@ use crate::arrays::Primitive;
2525
use crate::arrays::PrimitiveArray;
2626
use crate::arrays::Struct;
2727
use crate::arrays::StructArray;
28-
use crate::arrays::Union;
29-
use crate::arrays::UnionArray;
3028
use crate::arrays::VarBinView;
3129
use crate::arrays::VarBinViewArray;
3230
use crate::arrays::VariantArray;
3331
use crate::arrays::dict::TakeExecute;
3432
use crate::arrays::dict::TakeReduce;
33+
use crate::arrays::union::compute::take_union;
3534
use crate::arrays::variant::VariantArrayExt;
3635

3736
/// Take from a canonical array using indices (codes), returning a new canonical array.
@@ -55,7 +54,10 @@ pub(crate) fn take_canonical(
5554
Canonical::FixedSizeList(take_fixed_size_list(&a, codes, ctx))
5655
}
5756
Canonical::Struct(a) => Canonical::Struct(take_struct(&a, codes)),
58-
Canonical::Union(a) => Canonical::Union(take_union(&a, codes)?),
57+
Canonical::Union(a) => {
58+
let indices = codes.clone().into_array();
59+
Canonical::Union(take_union(a.as_view(), &indices)?)
60+
}
5961
Canonical::Extension(a) => Canonical::Extension(take_extension(&a, codes, ctx)),
6062
Canonical::Variant(a) => {
6163
let indices = codes.clone().into_array();
@@ -168,15 +170,6 @@ fn take_struct(array: &StructArray, codes: &PrimitiveArray) -> StructArray {
168170
.into_owned()
169171
}
170172

171-
fn take_union(array: &UnionArray, codes: &PrimitiveArray) -> VortexResult<UnionArray> {
172-
let codes_ref = codes.clone().into_array();
173-
let array = array.as_view();
174-
Ok(<Union as TakeReduce>::take(array, &codes_ref)?
175-
.vortex_expect("take UnionArray should be supported")
176-
.as_::<Union>()
177-
.into_owned())
178-
}
179-
180173
fn take_extension(
181174
array: &ExtensionArray,
182175
codes: &PrimitiveArray,

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

Lines changed: 8 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -21,11 +21,10 @@ use crate::arrays::ConstantArray;
2121
use crate::arrays::ExtensionArray;
2222
use crate::arrays::Filter;
2323
use crate::arrays::NullArray;
24-
use crate::arrays::Union;
2524
use crate::arrays::VariantArray;
2625
use crate::arrays::extension::ExtensionArrayExt;
2726
use crate::arrays::filter::FilterArrayExt;
28-
use crate::arrays::filter::FilterReduce;
27+
use crate::arrays::union::compute::filter_union;
2928
use crate::arrays::variant::VariantArrayExt;
3029
use crate::scalar::Scalar;
3130
use crate::validity::Validity;
@@ -97,13 +96,13 @@ pub(super) fn execute_filter(canonical: Canonical, mask: &Arc<MaskValues>) -> Ca
9796
Canonical::FixedSizeList(fixed_size_list::filter_fixed_size_list(&a, mask))
9897
}
9998
Canonical::Struct(a) => Canonical::Struct(struct_::filter_struct(&a, mask)),
100-
Canonical::Union(a) => Canonical::Union(
101-
<Union as FilterReduce>::filter(a.as_view(), &Mask::Values(Arc::clone(mask)))
102-
.vortex_expect("filter UnionArray")
103-
.vortex_expect("UnionArray filter must be supported")
104-
.as_::<Union>()
105-
.into_owned(),
106-
),
99+
Canonical::Union(a) => {
100+
let mask = Mask::Values(Arc::clone(mask));
101+
Canonical::Union(
102+
filter_union(a.as_view(), &mask)
103+
.vortex_expect("UnionArray children must support filter"),
104+
)
105+
}
107106
Canonical::Extension(a) => {
108107
let filtered_storage = a
109108
.storage_array()

0 commit comments

Comments
 (0)