Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions docs/source/contributor-guide/adding_a_new_expression.md
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,10 @@ object CometLevenshtein extends CometExpressionSerde[Levenshtein] {

When the return type is set on the proto, the native planner skips the registry lookup entirely and routes straight to the Comet UDF registered in `create_comet_physical_fun_with_eval_mode`.

#### Arguments that Spark skips after a NULL

A null-intolerant `BinaryExpression` or `TernaryExpression` returns NULL as soon as one of its arguments is NULL, without evaluating the arguments after it. Native execution evaluates every argument of a `ScalarFunc` over the whole batch, so an argument that can fail, such as an ANSI cast of a malformed string, would fail on rows that Spark skips. Wrap the serialized call of such an expression in `withNullShortCircuit`, as `CometArrayContains` and `CometSlice` do, after checking the conditions in its Scaladoc.

#### Registering the Expression Handler

Once you've created your `CometExpressionSerde` implementation, register it in `QueryPlanSerde.scala` by adding it to the appropriate expression map (e.g., `mathExpressions`, `stringExpressions`, `predicateExpressions`, etc.):
Expand Down
1 change: 1 addition & 0 deletions native/core/src/execution/jni_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2777,6 +2777,7 @@ mod tests {
}],
return_type: None,
fail_on_error: false,
null_short_circuit: false,
})),
query_context: None,
expr_id: None,
Expand Down
20 changes: 16 additions & 4 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -146,8 +146,8 @@ use datafusion_comet_spark_expr::{
ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct,
DecimalRescaleCheckOverflow, FloatOperands, GetArrayStructFields, GetStructField, HllPlusPlus,
HllSketchAgg, HllUnionAgg, IfExpr, ListExtract, MaxMinBy, Mode, NormalizeNaNAndZero,
NormalizeNestedFloats, Regr, RegrType, SparkCastOptions, SparkMinMax, Stddev, SumDecimal,
ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
NormalizeNestedFloats, NullShortCircuit, Regr, RegrType, SparkCastOptions, SparkMinMax, Stddev,
SumDecimal, ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
};
use itertools::Itertools;
use jni::objects::{Global, JObject};
Expand Down Expand Up @@ -955,15 +955,20 @@ impl PhysicalPlanner {
.as_ref()
.map(|e| self.create_expr(e, Arc::clone(&input_schema)))
.transpose()?;
Ok(Arc::new(ListExtract::new(
let list_extract: Arc<dyn PhysicalExpr> = Arc::new(ListExtract::new(
child,
ordinal,
default_value,
expr.one_based,
expr.fail_on_error,
spark_expr.expr_id,
Arc::clone(&self.query_context_registry),
)))
));
if expr.null_short_circuit {
Ok(NullShortCircuit::wrap(list_extract, &input_schema)?)
} else {
Ok(list_extract)
}
}
ExprStruct::GetArrayStructFields(expr) => {
let child =
Expand Down Expand Up @@ -4033,6 +4038,11 @@ impl PhysicalPlanner {
Arc::new(Field::new(fun_name, data_type.clone(), true)),
Arc::new(ConfigOptions::default()),
));
let scalar_expr = if expr.null_short_circuit {
NullShortCircuit::wrap(scalar_expr, &input_schema)?
} else {
scalar_expr
};

// DF53 changed some UDFs (e.g. md5) to return StringViewArray at execution
// time (apache/datafusion#20045). Comet does not yet support view types, so
Expand Down Expand Up @@ -6688,6 +6698,7 @@ mod tests {
args: vec![array_col, array_col_1],
return_type: None,
fail_on_error: false,
null_short_circuit: false,
})),
query_context: None,
expr_id: None,
Expand Down Expand Up @@ -6814,6 +6825,7 @@ mod tests {
args: vec![array_col, array_col_1],
return_type: None,
fail_on_error: false,
null_short_circuit: false,
})),
query_context: None,
expr_id: None,
Expand Down
6 changes: 6 additions & 0 deletions native/proto/src/proto/expr.proto
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,10 @@ message ScalarFunc {
repeated Expr args = 2;
DataType return_type = 3;
bool fail_on_error = 4;
// Evaluate each argument after the first only for the rows where no earlier argument is NULL,
// as Spark does for a null-intolerant BinaryExpression or TernaryExpression. See
// QueryPlanSerde.withNullShortCircuit.
bool null_short_circuit = 5;
}

message CaseWhen {
Expand Down Expand Up @@ -631,6 +635,8 @@ message ListExtract {
Expr default_value = 3;
bool one_based = 4;
bool fail_on_error = 5;
// Evaluate the ordinal only for the rows where the child is not NULL, as in ScalarFunc
bool null_short_circuit = 6;
}

message GetArrayStructFields {
Expand Down
60 changes: 55 additions & 5 deletions native/spark-expr/benches/conditional.rs
Original file line number Diff line number Diff line change
Expand Up @@ -45,9 +45,10 @@ type Expr = Arc<dyn PhysicalExpr>;
type Shape<'a> = (&'a str, fn() -> Expr, &'a RecordBatch);

/// Columns shaped like the ones in `CometConditionalExpressionBenchmark`:
/// `c1` a random long, `c2` an int in `0..100`, `c3` another random long, and `c4` / `c5` short
/// and `c6` / `c7` long strings. `null_density` applies to every column. With `sorted`, `c1` and
/// `c2` ascend through the batch, so every predicate over them selects one contiguous run of rows.
/// `c1` a random long, `c2` an int in `0..100`, `c3` another random long, `c4` / `c5` short
/// and `c6` / `c7` long strings, and `c8` / `c9` integers written as strings. `null_density`
/// applies to every column. With `sorted`, `c1` and `c2` ascend through the batch, so every
/// predicate over them selects one contiguous run of rows.
fn make_batch(null_density: f32, sorted: bool) -> RecordBatch {
let mut rng = StdRng::seed_from_u64(42);
let mut c1: Vec<i64> = (0..NUM_ROWS).map(|_| rng.random::<i64>()).collect();
Expand All @@ -57,7 +58,7 @@ fn make_batch(null_density: f32, sorted: bool) -> RecordBatch {
c2.sort_unstable();
}
let c3: Vec<i64> = (0..NUM_ROWS).map(|_| rng.random::<i64>()).collect();
let valid: Vec<Vec<bool>> = (0..7)
let valid: Vec<Vec<bool>> = (0..9)
.map(|_| {
(0..NUM_ROWS)
.map(|_| rng.random::<f32>() >= null_density)
Expand All @@ -84,6 +85,11 @@ fn make_batch(null_density: f32, sorted: bool) -> RecordBatch {
.map(|i| valid[col][i].then(|| format!("{tag}{i}-").repeat(12)))
.collect()
};
let numeric = |col: usize, seed: usize| -> StringArray {
(0..NUM_ROWS)
.map(|i| valid[col][i].then(|| format!("{}", (i * 7919 + seed) % 2_000_000)))
.collect()
};
let columns: Vec<ArrayRef> = vec![
Arc::new(c1),
Arc::new(c2),
Expand All @@ -92,6 +98,8 @@ fn make_batch(null_density: f32, sorted: bool) -> RecordBatch {
Arc::new(short(4, "t")),
Arc::new(long(5, "long value ")),
Arc::new(long(6, "other value ")),
Arc::new(numeric(7, 17)),
Arc::new(numeric(8, 4242)),
];
RecordBatch::try_new(Arc::new(schema()), columns).unwrap()
}
Expand All @@ -105,6 +113,8 @@ fn schema() -> Schema {
Field::new("c5", DataType::Utf8, true),
Field::new("c6", DataType::Utf8, true),
Field::new("c7", DataType::Utf8, true),
Field::new("c8", DataType::Utf8, true),
Field::new("c9", DataType::Utf8, true),
])
}

Expand Down Expand Up @@ -309,13 +319,45 @@ fn case_fallible_branch() -> Expr {
)
}

/// A string-to-integer cast with ANSI off, which returns NULL for a string it cannot parse.
fn parse_int(column: &str) -> Expr {
spark_cast(col(column), DataType::Int32)
}

/// A selective predicate whose branch parses a string: 5% of the rows choose it.
fn if_selective_parse() -> Expr {
if_expr(c2_below(5), parse_int("c8"), lit(ScalarValue::Int32(None)))
}

/// A predicate that most rows match, with a branch that parses a string.
fn if_mostly_parse() -> Expr {
if_expr(c2_below(95), parse_int("c8"), lit(ScalarValue::Int32(None)))
}

/// Both branches parse a string, and half of the rows choose each.
fn if_parse_either() -> Expr {
if_expr(c2_below(50), parse_int("c8"), parse_int("c9"))
}

/// Three selective branches, each parsing a string for 2% of the rows.
fn case_selective_parse_3_branches() -> Expr {
case_when(
vec![
(c2_below(2), parse_int("c8")),
(c2_below(4), parse_int("c9")),
(c2_below(6), parse_int("c8")),
],
None,
)
}

fn criterion_benchmark(c: &mut Criterion) {
let no_nulls = make_batch(0.0, false);
let sparse_nulls = make_batch(0.1, false);
let dense_nulls = make_batch(0.9, false);
let sorted = make_batch(0.0, true);

let shapes: [Shape; 26] = [
let shapes: [Shape; 30] = [
(
"case literal 3 branches",
case_literal_3_branches,
Expand Down Expand Up @@ -374,6 +416,14 @@ fn criterion_benchmark(c: &mut Criterion) {
case_column_10_branches,
&sorted,
),
("if selective parse", if_selective_parse, &no_nulls),
("if mostly parse", if_mostly_parse, &no_nulls),
("if parse either", if_parse_either, &no_nulls),
(
"case selective parse 3 branches",
case_selective_parse_3_branches,
&no_nulls,
),
];
let mut group = c.benchmark_group("conditional");
group.throughput(Throughput::Elements(NUM_ROWS as u64));
Expand Down
68 changes: 56 additions & 12 deletions native/spark-expr/src/conditional_funcs/case_when.rs
Original file line number Diff line number Diff line change
Expand Up @@ -385,12 +385,27 @@ impl PhysicalExpr for CaseWhenExpr {
}
}

/// Whether `expr` can be evaluated for rows that Spark would not evaluate it for.
/// Whether `expr` can be evaluated for rows that Spark would not evaluate it for, at little cost.
///
/// That holds when it can neither fail nor return something different for seeing more rows: a
/// column, a literal, and comparisons, boolean logic, null checks, widening casts and wrapping
/// arithmetic over them. Anything else, including every function, is assumed to be able to fail.
fn is_infallible(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema) -> bool {
/// That holds when it can neither fail nor return something different for seeing more rows, and is
/// cheap: a column, a literal, and comparisons, boolean logic, null checks, widening casts and
/// wrapping arithmetic over them. Anything else, including every function and a cast that parses a
/// string, is treated as able to fail. `CASE` uses this to decide whether to evaluate every branch
/// over the whole batch, which pays off only when that costs little more than selecting rows.
pub(super) fn is_infallible(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema) -> bool {
infallible(expr, input_schema, false)
}

/// Whether `expr` can be evaluated for rows that Spark would not evaluate it for, at any cost: what
/// [`is_infallible`] accepts, and casts that [`Cast::cannot_fail`] accepts beyond it, such as a
/// LEGACY or TRY cast from a string to an integer, which returns NULL for a string it cannot parse.
/// `NullShortCircuit` uses this, since it masks an argument only to keep it from failing on rows
/// that Spark skips.
pub(super) fn cannot_fail(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema) -> bool {
infallible(expr, input_schema, true)
}

fn infallible(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema, parsing: bool) -> bool {
if expr.is::<Column>() || expr.is::<Literal>() {
return true;
}
Expand All @@ -399,7 +414,11 @@ fn is_infallible(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema) -> bool {
} else if let Some(comparison) = expr.downcast_ref::<SparkComparison>() {
comparison.is_infallible(input_schema)
} else if let Some(cast) = expr.downcast_ref::<Cast>() {
cast.is_infallible(input_schema)
if parsing {
cast.cannot_fail(input_schema)
} else {
cast.is_infallible(input_schema)
}
} else if let Some(normalize) = expr.downcast_ref::<NormalizeNaNAndZero>() {
normalize.is_infallible()
} else if let Some(not) = expr.downcast_ref::<NotExpr>() {
Expand All @@ -417,7 +436,7 @@ fn is_infallible(expr: &Arc<dyn PhysicalExpr>, input_schema: &Schema) -> bool {
&& expr
.children()
.into_iter()
.all(|child| is_infallible(child, input_schema))
.all(|child| infallible(child, input_schema, parsing))
}

fn binary_is_infallible(binary: &BinaryExpr, input_schema: &Schema) -> bool {
Expand Down Expand Up @@ -1147,6 +1166,7 @@ mod tests {
Field::new("d", DataType::Float64, true),
Field::new("e", DataType::Float64, true),
Field::new("s", DataType::Utf8, true),
Field::new("ls", DataType::LargeUtf8, true),
Field::new("dec", DataType::Decimal128(10, 2), true),
Field::new("l", list_of(DataType::Float64), true),
Field::new("m", list_of(DataType::Float64), true),
Expand All @@ -1160,15 +1180,17 @@ mod tests {
),
]);
let c = |name: &str| col(name, &schema).unwrap();
let cast = |e: Arc<dyn PhysicalExpr>, to: DataType| -> Arc<dyn PhysicalExpr> {
Arc::new(Cast::new(
let cast_in = |mode: EvalMode, e: Arc<dyn PhysicalExpr>, to: DataType| {
let cast: Arc<dyn PhysicalExpr> = Arc::new(Cast::new(
e,
to,
SparkCastOptions::new_without_timezone(EvalMode::Ansi),
SparkCastOptions::new_without_timezone(mode),
None,
None,
))
));
cast
};
let cast = |e: Arc<dyn PhysicalExpr>, to: DataType| cast_in(EvalMode::Ansi, e, to);
let infallible = |e: Arc<dyn PhysicalExpr>| is_infallible(&e, &schema);
// A comparison as the planner builds it, which follows Spark's ordering for floats
let compare = |left: &str, op: Operator, right: Arc<dyn PhysicalExpr>| {
Expand Down Expand Up @@ -1225,9 +1247,31 @@ mod tests {
assert!(!infallible(checked));
// Decimal arithmetic can overflow
assert!(!infallible(binary(c("dec"), Operator::Plus, c("dec"))));
// Narrowing and parsing casts can fail
// Narrowing and parsing casts can fail under ANSI
assert!(!infallible(cast(c("i"), DataType::Int32)));
assert!(!infallible(cast(c("s"), DataType::Int32)));
// A string that LEGACY or TRY mode cannot parse as an integer becomes NULL, so the cast
// cannot fail. Parsing costs too much for CASE to do it for rows it would skip, though.
for mode in [EvalMode::Legacy, EvalMode::Try] {
for to in [
DataType::Int8,
DataType::Int16,
DataType::Int32,
DataType::Int64,
] {
for column in ["s", "ls"] {
let parse = cast_in(mode, c(column), to.clone());
assert!(!infallible(Arc::clone(&parse)), "{mode:?} {column} {to}");
assert!(cannot_fail(&parse, &schema), "{mode:?} {column} {to}");
}
}
// Other string casts are not claimed
assert!(
!cannot_fail(&cast_in(mode, c("s"), DataType::Float64), &schema),
"{mode:?}"
);
}
assert!(!cannot_fail(&cast(c("s"), DataType::Int32), &schema));
// So can anything built from a part that can fail
assert!(!infallible(binary(
binary(c("i"), Operator::Divide, c("i")),
Expand Down
2 changes: 2 additions & 0 deletions native/spark-expr/src/conditional_funcs/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@

mod case_when;
mod if_expr;
mod null_short_circuit;

pub use case_when::{create_case_when, create_if_expr, CaseWhenExpr};
pub use if_expr::IfExpr;
pub use null_short_circuit::NullShortCircuit;
Loading
Loading