Skip to content

Commit ed16fc6

Browse files
committed
Pushdown some expressions to Dict layout reader
Signed-off-by: Mikhail Kot <mikhail@spiraldb.com>
1 parent 9444d20 commit ed16fc6

3 files changed

Lines changed: 148 additions & 8 deletions

File tree

vortex-array/src/expr/transform/partition.rs

Lines changed: 117 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,12 +20,62 @@ use crate::expr::analysis::AnnotationFn;
2020
use crate::expr::analysis::Annotations;
2121
use crate::expr::analysis::descendent_annotations;
2222
use crate::expr::get_item;
23+
use crate::expr::is_root;
24+
use crate::expr::label_is_fallible;
25+
use crate::expr::label_null_sensitive;
2326
use crate::expr::pack;
2427
use crate::expr::root;
2528
use crate::expr::traversal::NodeExt;
2629
use crate::expr::traversal::NodeRewriter;
2730
use crate::expr::traversal::Transformed;
2831
use crate::expr::traversal::TraversalOrder;
32+
use crate::scalar_fn::is_negative_cost;
33+
34+
fn references_root(expr: &Expression) -> bool {
35+
is_root(expr) || expr.children().iter().any(references_root)
36+
}
37+
38+
/// Split expression into two parts:
39+
///
40+
/// left is the optional outer part that we want to apply to array after
41+
/// canonicalizing.
42+
/// right is the optional inner part that we want to apply to array before
43+
/// canonicalizing.
44+
///
45+
/// We want to push to array only if expression has a negative cost, is
46+
/// infallible and null-insensitive.
47+
///
48+
/// TODO(myrrc): This is a specialized version of partition(), and we want to
49+
/// unify this with expression partitioning logic.
50+
pub fn split_expression_for_pushdown(expr: Expression) -> (Option<Expression>, Option<Expression>) {
51+
let labelled_expr = expr.clone();
52+
let fallible = label_is_fallible(&labelled_expr);
53+
let null_sensitive = label_null_sensitive(&labelled_expr);
54+
let mut inner: Option<Expression> = None;
55+
56+
let outer = expr
57+
.transform_down(|node| {
58+
if is_negative_cost(node.id())
59+
&& references_root(&node)
60+
&& !fallible.get(&node).copied().unwrap_or(true)
61+
&& !null_sensitive.get(&node).copied().unwrap_or(true)
62+
{
63+
inner = Some(node);
64+
Ok(Transformed {
65+
value: root(),
66+
changed: true,
67+
order: TraversalOrder::Skip,
68+
})
69+
} else {
70+
Ok(Transformed::no(node))
71+
}
72+
})
73+
.vortex_expect("infallible")
74+
.into_inner();
75+
76+
let outer = (!is_root(&outer)).then_some(outer);
77+
(outer, inner)
78+
}
2979

3080
/// Partition an expression into sub-expressions that are uniquely associated with an annotation.
3181
/// A root expression is also returned that can be used to recombine the results of the partitions
@@ -205,10 +255,19 @@ where
205255
mod tests {
206256
use rstest::fixture;
207257
use rstest::rstest;
208-
258+
use vortex_array::expr::Expression;
259+
use vortex_array::expr::byte_length;
260+
use vortex_array::expr::cast;
261+
use vortex_array::expr::like;
262+
use vortex_array::expr::traversal::NodeExt;
263+
use vortex_array::expr::traversal::Transformed;
264+
use vortex_array::expr::traversal::TraversalOrder;
265+
266+
use super::split_expression_for_pushdown;
209267
use super::*;
210268
use crate::dtype::DType;
211269
use crate::dtype::Nullability::NonNullable;
270+
use crate::dtype::PType;
212271
use crate::dtype::PType::I32;
213272
use crate::dtype::StructFields;
214273
use crate::expr::analysis::make_free_field_annotator;
@@ -348,4 +407,61 @@ mod tests {
348407
let expected_b = pack([("b_0", pack([("b", col("b"))], NonNullable))], NonNullable);
349408
assert_eq!(part_b, &expected_b, "{part_b} {expected_b}");
350409
}
410+
411+
fn join_split_expr(initial: &Expression, outer: Option<Expression>, inner: Option<Expression>) {
412+
let outer_expr = outer.unwrap_or_else(root);
413+
let inner_expr = inner.unwrap_or_else(root);
414+
let expected = outer_expr
415+
.transform_down(|node| {
416+
if !is_root(&node) {
417+
return Ok(Transformed::no(node));
418+
}
419+
Ok(Transformed {
420+
value: inner_expr.clone(),
421+
changed: true,
422+
order: TraversalOrder::Skip,
423+
})
424+
})
425+
.vortex_expect("infallible");
426+
assert_eq!(&expected.into_inner(), initial);
427+
}
428+
429+
#[test]
430+
fn split_expr_cast_root() {
431+
let (outer, inner) = split_expression_for_pushdown(root());
432+
assert_eq!(outer, None);
433+
assert_eq!(inner, None); // Applying root to array is useless work
434+
}
435+
436+
#[test]
437+
fn split_expr_partial_pushdown() {
438+
let dtype = DType::Primitive(PType::U64, NonNullable);
439+
let expr = cast(byte_length(root()), dtype.clone());
440+
let (outer, inner) = split_expression_for_pushdown(expr.clone());
441+
// [0] = cast([1], dtype)
442+
// [1] = byte_length(root)
443+
assert_eq!(outer, Some(cast(root(), dtype)));
444+
assert_eq!(inner, Some(byte_length(root())));
445+
join_split_expr(&expr, outer, inner);
446+
}
447+
448+
#[test]
449+
fn split_expr_full_pushdown() {
450+
let expr = byte_length(root());
451+
let (outer, inner) = split_expression_for_pushdown(expr.clone());
452+
assert_eq!(outer, None);
453+
assert_eq!(inner, Some(byte_length(root())));
454+
join_split_expr(&expr, outer, inner);
455+
}
456+
457+
#[test]
458+
fn split_expr_no_pushdown() {
459+
// We can push down lit(), but it we replace
460+
// lit() with root(), the semantics change.
461+
let expr = like(root(), lit(1u64));
462+
let (outer, inner) = split_expression_for_pushdown(expr.clone());
463+
assert_eq!(outer, Some(expr.clone()));
464+
assert_eq!(inner, None);
465+
join_split_expr(&expr, outer, inner);
466+
}
351467
}

vortex-array/src/scalar_fn/mod.rs

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,10 @@
99
1010
use vortex_session::registry::Id;
1111

12+
use crate::scalar_fn::fns::byte_length::ByteLength;
13+
use crate::scalar_fn::fns::get_item::GetItem;
14+
use crate::scalar_fn::fns::literal::Literal;
15+
1216
mod vtable;
1317
pub use vtable::*;
1418

@@ -48,3 +52,16 @@ mod sealed {
4852
/// This can be the **only** implementor for [`super::typed::DynScalarFn`].
4953
impl<V: ScalarFnVTable> Sealed for TypedScalarFnInstance<V> {}
5054
}
55+
56+
/// A scalar function has a negative cost if applying it to an array and
57+
/// canonicalizing is cheaper than canonicalizing an array and applying it.
58+
///
59+
/// Example of negative cost expressions are byte_length() and get_item() since
60+
/// they don't depend on input size.
61+
///
62+
/// Example of non-negative cost expression is like()
63+
pub fn is_negative_cost(id: ScalarFnId) -> bool {
64+
id == ScalarFnVTable::id(&ByteLength)
65+
|| id == ScalarFnVTable::id(&GetItem)
66+
|| id == ScalarFnVTable::id(&Literal)
67+
}

vortex-layout/src/layouts/dict/reader.rs

Lines changed: 14 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ use vortex_array::dtype::DType;
2020
use vortex_array::dtype::FieldMask;
2121
use vortex_array::expr::Expression;
2222
use vortex_array::expr::root;
23+
use vortex_array::expr::transform::split_expression_for_pushdown;
2324
use vortex_array::optimizer::ArrayOptimizer;
2425
use vortex_error::VortexError;
2526
use vortex_error::VortexExpect;
@@ -100,10 +101,7 @@ impl DictReader {
100101
)
101102
.vortex_expect("must construct dict values array evaluation")
102103
.map_err(Arc::new)
103-
.map(move |array| {
104-
let array = array?;
105-
Ok(SharedArray::new(array).into_array())
106-
})
104+
.map(move |array| Ok(SharedArray::new(array?).into_array()))
107105
.boxed()
108106
.shared()
109107
})
@@ -229,13 +227,18 @@ impl LayoutReader for DictReader {
229227
mask: MaskFuture,
230228
) -> VortexResult<BoxFuture<'static, VortexResult<ArrayRef>>> {
231229
// TODO: fix up expr partitioning with fallible & null sensitive annotations
232-
let values_eval = self.values_array();
233230
let codes_eval = self
234231
.codes
235232
.projection_evaluation(row_range, &root(), mask)
236233
.map_err(|err| err.with_context("While evaluating projection on codes"))?;
237-
let expr = expr.clone();
238234

235+
let (expr_outer, expr_inner) = split_expression_for_pushdown(expr.clone());
236+
237+
let values_eval = if let Some(inner) = expr_inner {
238+
self.values_eval(inner)
239+
} else {
240+
self.values_array()
241+
};
239242
let all_values_referenced = self.layout.has_all_values_referenced();
240243
Ok(async move {
241244
let (values, codes) = try_join!(values_eval.map_err(VortexError::from), codes_eval)?;
@@ -252,7 +255,11 @@ impl LayoutReader for DictReader {
252255
.into_array()
253256
.optimize()?;
254257

255-
array.apply(&expr)
258+
if let Some(expr) = expr_outer {
259+
array.apply(&expr)
260+
} else {
261+
Ok(array)
262+
}
256263
}
257264
.boxed())
258265
}

0 commit comments

Comments
 (0)