Skip to content
Merged
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
10 changes: 10 additions & 0 deletions docs/source/user-guide/latest/compatibility/floating-point.md
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,16 @@ to Spark in some cases, especially when the data contains both positive and nega
case that is not of concern for many users. If it is a concern, setting `spark.comet.exec.strictFloatingPoint=true`
will make relevant operations fall back to Spark.

## Nested equality and membership

For arrays and structs containing `FLOAT` or `DOUBLE`, native `=`, `<>`, `IN`, and `NOT IN`
compare signed zeros as equal and all NaN representations as equal, matching Spark. This also
covers single-candidate membership that Spark rewrites into equality.

Equality and dynamic membership compare nested elements directly and stop at the first mismatch.
Constant membership sets use normalized comparison values for static lookup. These operations
preserve SQL null semantics and do not change the values returned by projections.

## Ordering: NaN and signed zero (`-0.0` vs `+0.0`)

Spark's `ORDER BY`, `RANK`, `DENSE_RANK`, and window frame comparisons route through
Expand Down
16 changes: 8 additions & 8 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -74,8 +74,7 @@ use datafusion::{
logical_expr::Operator as DataFusionOperator,
physical_expr::{
expressions::{
in_list, BinaryExpr, CaseExpr, CastExpr, Column, IsNullExpr,
Literal as DataFusionLiteral,
BinaryExpr, CaseExpr, CastExpr, Column, IsNullExpr, Literal as DataFusionLiteral,
},
PhysicalExpr, PhysicalSortExpr, ScalarFunctionExpr,
},
Expand Down Expand Up @@ -148,11 +147,11 @@ use datafusion_comet_proto::{
spark_partitioning::{partitioning::PartitioningStruct, Partitioning as SparkPartitioning},
};
use datafusion_comet_spark_expr::{
jvm_udf::JvmScalarUdfExpr, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow,
Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow, GetArrayStructFields,
GetStructField, HllPlusPlus, IfExpr, ListExtract, MaxMinBy, Mode, NormalizeNaNAndZero, Regr,
RegrType, SparkCastOptions, Stddev, SumDecimal, ToJson, UnboundColumn, Variance,
WideDecimalBinaryExpr, WideDecimalOp,
jvm_udf::JvmScalarUdfExpr, spark_in_list, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast,
CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow,
GetArrayStructFields, GetStructField, HllPlusPlus, IfExpr, ListExtract, MaxMinBy, Mode,
NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, Stddev, SumDecimal, ToJson,
UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
};
use itertools::Itertools;
use jni::objects::{Global, JObject};
Expand Down Expand Up @@ -762,7 +761,8 @@ impl PhysicalPlanner {
.map(|x| self.create_expr(x, Arc::clone(&input_schema)))
.collect::<Result<Vec<_>, _>>()?;

in_list(value, list, &expr.negated, input_schema.as_ref()).map_err(|e| e.into())
spark_in_list(value, list, expr.negated, input_schema.as_ref())
.map_err(|e| e.into())
}
ExprStruct::If(expr) => {
let if_expr =
Expand Down
10 changes: 7 additions & 3 deletions native/core/src/execution/planner/macros.rs
Original file line number Diff line number Diff line change
Expand Up @@ -93,9 +93,13 @@ macro_rules! binary_expr_builder {
&$operator,
&input_schema,
);
Ok(std::sync::Arc::new(
datafusion::physical_expr::expressions::BinaryExpr::new(left, $operator, right),
))
datafusion_comet_spark_expr::spark_comparison(
left,
$operator,
right,
input_schema.as_ref(),
)
.map_err(Into::into)
}
}
};
Expand Down
4 changes: 4 additions & 0 deletions native/spark-expr/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -430,3 +430,7 @@ harness = false
[[bench]]
name = "utf8_decode"
harness = false

[[bench]]
name = "nested_comparison"
harness = false
170 changes: 170 additions & 0 deletions native/spark-expr/benches/nested_comparison.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,170 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.

//! Compare the base DataFusion path, PR #6073's eager normalization, and Spark equality.
//! Ordinary finite inputs keep the answers identical across all three implementations.

use arrow::array::{ArrayRef, Float64Array, ListArray};
use arrow::buffer::{NullBuffer, OffsetBuffer};
use arrow::datatypes::{DataType, Field, Schema};
use arrow::record_batch::RecordBatch;
use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion};
use datafusion::common::ScalarValue;
use datafusion::logical_expr::Operator;
use datafusion::physical_expr::expressions::{in_list, BinaryExpr, Column, Literal};
use datafusion::physical_expr::PhysicalExpr;
use datafusion_comet_spark_expr::{spark_comparison, spark_in_list, NormalizeNestedFloats};
use std::hint::black_box;
use std::sync::Arc;
use std::time::Duration;

const ROWS: usize = 8192;

fn input(width: usize, mismatch: Option<usize>, null_every: usize) -> ArrayRef {
let values = (0..ROWS * width)
.map(|i| {
if mismatch == Some(i % width) {
2.0
} else {
1.0
}
})
.collect::<Vec<_>>();
let nulls = (null_every != 0).then(|| {
NullBuffer::from(
(0..ROWS)
.map(|i| (i + 1) % null_every != 0)
.collect::<Vec<_>>(),
)
});
Arc::new(ListArray::new(
Arc::new(Field::new("item", DataType::Float64, true)),
OffsetBuffer::from_lengths(std::iter::repeat_n(width, ROWS)),
Arc::new(Float64Array::from(values)),
nulls,
))
}

fn expression(version: &str, mode: &str, batch: &RecordBatch) -> Arc<dyn PhysicalExpr> {
let schema = batch.schema();
let a: Arc<dyn PhysicalExpr> = Arc::new(Column::new("a", 0));
let b: Arc<dyn PhysicalExpr> = Arc::new(Column::new("b", 1));
if mode == "eq" {
return if version == "new" {
spark_comparison(a, Operator::Eq, b, &schema).unwrap()
} else {
Arc::new(BinaryExpr::new(a, Operator::Eq, b))
};
}
let literal: Arc<dyn PhysicalExpr> = Arc::new(Literal::new(
ScalarValue::try_from_array(batch.column(1), 0).unwrap(),
));
let candidates = match mode {
"constant" => vec![literal],
"dynamic" => vec![b],
"mixed" => vec![literal, b],
_ => unreachable!(),
};
match version {
"base" => in_list(a, candidates, &false, &schema).unwrap(),
"head" => in_list(
NormalizeNestedFloats::wrap_if_needed(a, &schema).unwrap(),
candidates
.into_iter()
.map(|e| NormalizeNestedFloats::wrap_if_needed(e, &schema).unwrap())
.collect(),
&false,
&schema,
)
.unwrap(),
"new" => spark_in_list(a, candidates, false, &schema).unwrap(),
_ => unreachable!(),
}
}

fn benchmark(c: &mut Criterion) {
let mut group = c.benchmark_group("nested_comparison");
group
.sample_size(10)
.warm_up_time(Duration::from_millis(100))
.measurement_time(Duration::from_millis(500));
for width in [1, 16, 1024] {
for (shape, mismatch) in [
("first", Some(0)),
("last", Some(width - 1)),
("equal", None),
] {
for (nulls, null_every) in [("no_nulls", 0), ("sparse", 16), ("dense", 2)] {
let a = input(width, None, null_every);
let b = input(width, mismatch, null_every);
let schema = Arc::new(Schema::new(vec![
Field::new("a", a.data_type().clone(), true),
Field::new("b", b.data_type().clone(), true),
]));
let batch = RecordBatch::try_new(schema, vec![a, b]).unwrap();
for mode in ["dynamic", "constant", "mixed", "eq"] {
let expected = expression("base", mode, &batch)
.evaluate(&batch)
.unwrap()
.into_array(ROWS)
.unwrap();
for version in ["base", "head", "new"] {
let expr = expression(version, mode, &batch);
assert_eq!(
expected.as_ref(),
expr.evaluate(&batch)
.unwrap()
.into_array(ROWS)
.unwrap()
.as_ref()
);
group.bench_with_input(
BenchmarkId::new(format!("{mode}/{width}/{shape}/{nulls}"), version),
&expr,
|b, e| {
b.iter(|| black_box(e.evaluate(black_box(&batch)).unwrap()));
},
);
}
}
}
}
}
group.finish();
let mut group = c.benchmark_group("nested_static_build");
group
.sample_size(10)
.warm_up_time(Duration::from_millis(100))
.measurement_time(Duration::from_millis(500));
for width in [1, 16, 1024] {
let a = input(width, None, 0);
let schema = Arc::new(Schema::new(vec![
Field::new("a", a.data_type().clone(), true),
Field::new("b", a.data_type().clone(), true),
]));
let batch = RecordBatch::try_new(schema, vec![Arc::clone(&a), a]).unwrap();
for version in ["base", "head", "new"] {
group.bench_function(BenchmarkId::new(width.to_string(), version), |b| {
b.iter(|| black_box(expression(version, "constant", black_box(&batch))))
});
}
}
group.finish();
}

criterion_group!(benches, benchmark);
criterion_main!(benches);
3 changes: 3 additions & 0 deletions native/spark-expr/src/array_funcs/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ mod arrays_zip;
mod flatten;
mod get_array_struct_fields;
mod list_extract;
mod nested_comparison;
mod nested_float_normalize;
mod sequence;
mod size;
Expand All @@ -35,5 +36,7 @@ pub use arrays_zip::SparkArraysZipFunc;
pub use flatten::SparkFlatten;
pub use get_array_struct_fields::GetArrayStructFields;
pub use list_extract::ListExtract;
pub use nested_comparison::{spark_comparison, spark_in_list};
pub use nested_float_normalize::NormalizeNestedFloats;
pub use sequence::spark_sequence;
pub use size::{spark_size, SparkSizeFunc};
Loading
Loading