Skip to content
Draft
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
26 changes: 11 additions & 15 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -149,11 +149,12 @@ use datafusion_comet_proto::{
spark_partitioning::{partitioning::PartitioningStruct, Partitioning as SparkPartitioning},
};
use datafusion_comet_spark_expr::{
jvm_udf::JvmScalarUdfExpr, spark_in_list, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast,
CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow,
GetArrayStructFields, GetStructField, HllPlusPlus, HllSketchAgg, HllUnionAgg, IfExpr,
ListExtract, MaxMinBy, Mode, NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, Stddev,
SumDecimal, ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
jvm_udf::JvmScalarUdfExpr, normalize_floats, spark_in_list, ApproxPercentile, ArrayInsert, Avg,
AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct,
DecimalRescaleCheckOverflow, GetArrayStructFields, GetStructField, HllPlusPlus, HllSketchAgg,
HllUnionAgg, 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 @@ -1015,15 +1016,10 @@ impl PhysicalPlanner {
input_schema: SchemaRef,
) -> Result<Arc<dyn PhysicalExpr>, ExecutionError> {
let child = self.create_expr(spark_expr, Arc::clone(&input_schema))?;
let data_type = child.data_type(input_schema.as_ref())?;
// Spark may already have normalized a partition or join key.
if matches!(data_type, DataType::Float32 | DataType::Float64)
&& child.downcast_ref::<NormalizeNaNAndZero>().is_none()
{
Ok(Arc::new(NormalizeNaNAndZero::new(data_type, child)))
} else {
Ok(child)
}
Ok(NormalizeNaNAndZero::wrap_if_needed(
child,
input_schema.as_ref(),
)?)
}

/// Only constant literals are supported as scan defaults.
Expand Down Expand Up @@ -3705,7 +3701,7 @@ impl PhysicalPlanner {
.iter()
.map(|scalar_vec| {
ScalarValue::iter_to_array(scalar_vec.iter().cloned())
.map(|array| NormalizeNaNAndZero::normalize_array(&array))
.map(|array| normalize_floats(&array))
})
.collect::<Result<Vec<_>, _>>()?;

Expand Down
5 changes: 3 additions & 2 deletions native/core/src/parquet/cast_column/variant.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ use arrow::{
};
use datafusion::common::{DataFusionError, Result as DataFusionResult};
use datafusion_comet_common::SparkError;
use datafusion_comet_spark_expr::canonicalize_nan;
use parquet::variant::{
unshred_variant, ListBuilder, MetadataBuilder, ObjectBuilder, ObjectFieldBuilder, ParentState,
ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilderExt,
Expand Down Expand Up @@ -852,8 +853,8 @@ fn spark_typed_scalar<'m, 'v>(value: Variant<'m, 'v>) -> Variant<'m, 'v> {
.map(Variant::Decimal4)
.unwrap_or(value),
Variant::String(s) => Variant::from(s),
Variant::Float(v) if v.is_nan() => Variant::Float(f32::NAN),
Variant::Double(v) if v.is_nan() => Variant::Double(f64::NAN),
Variant::Float(v) => Variant::Float(canonicalize_nan(v)),
Variant::Double(v) => Variant::Double(canonicalize_nan(v)),
_ => value,
}
}
Expand Down
24 changes: 24 additions & 0 deletions native/core/src/parquet/cast_column/variant/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -980,3 +980,27 @@ fn normalize_null_parent_ignores_empty_children() {
assert!(output.is_null(0));
assert!(output.as_struct().column(0).is_null(0));
}

/// Spark writes typed floats through `floatToIntBits`/`doubleToLongBits`, so every NaN comes out
/// canonical while `-0.0` keeps its sign.
#[test]
fn typed_scalar_canonicalizes_nan_and_keeps_negative_zero() {
for bits in [0xfff8_0000_0000_0000, 0x7ff0_0000_0000_0001] {
let Variant::Double(v) = spark_typed_scalar(Variant::Double(f64::from_bits(bits))) else {
panic!("expected a double");
};
assert_eq!(v.to_bits(), f64::NAN.to_bits());
}
let Variant::Float(v) = spark_typed_scalar(Variant::Float(f32::from_bits(0xffc0_0000))) else {
panic!("expected a float");
};
assert_eq!(v.to_bits(), f32::NAN.to_bits());
let Variant::Double(v) = spark_typed_scalar(Variant::Double(-0.0)) else {
panic!("expected a double");
};
assert_eq!(v.to_bits(), (-0.0f64).to_bits());
let Variant::Float(v) = spark_typed_scalar(Variant::Float(-0.0)) else {
panic!("expected a float");
};
assert_eq!(v.to_bits(), (-0.0f32).to_bits());
}
76 changes: 40 additions & 36 deletions native/spark-expr/src/agg_funcs/hll_plus_plus.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,19 +17,15 @@

//! Spark-compatible `approx_count_distinct`, a faithful port of Spark's
//! `HyperLogLogPlusPlus` / `HyperLogLogPlusPlusHelper`. Values are hashed with Spark's
//! `XxHash64` (seed 42, floats normalized first) and the registers are stored using the exact
//! same packed-`Long` buffer layout Spark uses (10 six-bit registers per 64-bit word). Keeping
//! the wire format identical means the partial-aggregation state matches Spark's
//! `aggBufferSchema`, and the cardinality estimate uses the same bias-correction tables, so
//! results are bit-identical to Spark.
//! `XxHash64` (seed 42) and the registers are stored using the exact same packed-`Long` buffer
//! layout Spark uses (10 six-bit registers per 64-bit word). Keeping the wire format identical
//! means the partial-aggregation state matches Spark's `aggBufferSchema`, and the cardinality
//! estimate uses the same bias-correction tables, so results are bit-identical to Spark.

use crate::agg_funcs::hll_plus_plus_const::{BIAS_DATA, RAW_ESTIMATE_DATA, THRESHOLDS};
use crate::hash_funcs::create_xxhash64_hashes;
use crate::math_funcs::internal::normalize_float;
use arrow::array::{
Array, ArrayRef, AsArray, BooleanArray, Float32Array, Float64Array, Int64Array,
};
use arrow::datatypes::{DataType, Field, FieldRef, Float32Type, Float64Type};
use arrow::array::{Array, ArrayRef, AsArray, BooleanArray, Int64Array};
use arrow::datatypes::{DataType, Field, FieldRef};
use datafusion::common::{not_impl_err, Result, ScalarValue};
use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
use datafusion::logical_expr::{
Expand Down Expand Up @@ -120,33 +116,17 @@ impl AggregateUDFImpl for HllPlusPlus {
}
}

/// Normalize a float/double column the way Spark's `NormalizeNaNAndZero` does before hashing:
/// every NaN becomes the canonical NaN and `-0.0` becomes `0.0`. Returns the input unchanged for
/// non-floating-point types.
fn normalize_floats(array: &ArrayRef) -> ArrayRef {
match array.data_type() {
DataType::Float32 => {
let normalized: Float32Array =
array.as_primitive::<Float32Type>().unary(normalize_float);
Arc::new(normalized)
}
DataType::Float64 => {
let normalized: Float64Array =
array.as_primitive::<Float64Type>().unary(normalize_float);
Arc::new(normalized)
}
_ => Arc::clone(array),
}
}

/// Hash a value column with Spark's `XxHash64` (seed 42), normalizing floats first, reusing
/// `buf` as scratch to avoid a per-batch allocation. The buffer is re-seeded (not just cleared)
/// because `create_xxhash64_hashes` folds each value into the existing seed.
/// Hash a value column with Spark's `XxHash64` (seed 42), reusing `buf` as scratch to avoid a
/// per-batch allocation. The buffer is re-seeded (not just cleared) because
/// `create_xxhash64_hashes` folds each value into the existing seed.
///
/// Spark runs float inputs through `NormalizeFloatingNumbers` before hashing, but
/// `create_xxhash64_hashes` already hashes `-0.0` as `0.0` and every NaN as the canonical NaN, so
/// the result is the same without a separate pass.
fn hash_values_into(array: &ArrayRef, buf: &mut Vec<u64>) -> Result<()> {
let normalized = normalize_floats(array);
buf.clear();
buf.resize(normalized.len(), HASH_SEED);
create_xxhash64_hashes(&[normalized], buf)?;
buf.resize(array.len(), HASH_SEED);
create_xxhash64_hashes(&[Arc::clone(array)], buf)?;
Ok(())
}

Expand Down Expand Up @@ -484,7 +464,7 @@ impl GroupsAccumulator for HllPlusPlusGroupsAccumulator {
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{Int32Array, StringArray};
use arrow::array::{Float32Array, Float64Array, Int32Array, StringArray};

fn acc(p: usize) -> HllPlusPlusAccumulator {
HllPlusPlusAccumulator::new(p)
Expand Down Expand Up @@ -524,6 +504,30 @@ mod tests {
assert_eq!(eval(&mut a), 2);
}

/// Spark counts the two zeros as one value and every NaN as one value.
#[test]
fn floats_fold_negative_zero_and_nan() {
let values: ArrayRef = Arc::new(Float64Array::from(vec![
0.0,
-0.0,
f64::NAN,
f64::from_bits(0xfff8_0000_0000_0000),
f64::from_bits(0x7ff0_0000_0000_0001),
]));
let mut a = acc(9);
a.update_batch(&[values]).unwrap();
assert_eq!(eval(&mut a), 2);
let values: ArrayRef = Arc::new(Float32Array::from(vec![
0.0,
-0.0,
f32::NAN,
f32::from_bits(0xffc0_0000),
]));
let mut a = acc(9);
a.update_batch(&[values]).unwrap();
assert_eq!(eval(&mut a), 2);
}

#[test]
fn strings_small_cardinality() {
let mut a = acc(9);
Expand Down
66 changes: 12 additions & 54 deletions native/spark-expr/src/agg_funcs/max_min_by.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,10 @@
// specific language governing permissions and limitations
// under the License.

use arrow::array::{new_null_array, Array, ArrayRef, AsArray, BooleanArray};
use crate::float_semantics::normalize_floats;
use arrow::array::{new_null_array, Array, ArrayRef, BooleanArray};
use arrow::compute::SortOptions;
use arrow::datatypes::{DataType, Field, FieldRef, Float32Type, Float64Type};
use arrow::datatypes::{DataType, Field, FieldRef};
use arrow::row::{OwnedRow, RowConverter, SortField};
use datafusion::common::{not_impl_err, Result, ScalarValue};
use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
Expand All @@ -39,6 +40,13 @@ use std::sync::Arc;
/// buffer and, on a tie in the ordering, the later row wins. Because ties across
/// partitions are processed in an unspecified order, Spark documents the function as
/// non-deterministic when several rows share the extremum ordering.
///
/// # Float orderings
///
/// Spark compares the ordering with `SQLOrderingUtil.compareDoubles`, while the accumulators rank
/// it in Arrow's row format, which orders floats by IEEE 754 total order. The ordering column goes
/// through [`normalize_floats`] first, which makes the two agree. This is verified identical on
/// Spark 3.4 through master.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MaxMinBy {
name: String,
Expand Down Expand Up @@ -145,56 +153,6 @@ fn extremum_sort_options(is_max: bool) -> SortOptions {
}
}

/// Canonicalize a floating-point ordering column so that Arrow's row-format byte order reproduces
/// Spark's comparison for this aggregate.
///
/// Spark compares the ordering with `SQLOrderingUtil.compareDoubles`/`compareFloats`, wired in via
/// `PhysicalDoubleType.ordering`/`PhysicalFloatType.ordering`. That is
/// `if (x == y) 0 else java.lang.Double.compare(x, y)`, which has two consequences Arrow's row
/// format does not share:
///
/// * `-0.0` and `0.0` tie, because the `x == y` short-circuit is IEEE equality. Arrow encodes
/// floats by flipping the bits off the sign, a total order placing `-0.0` strictly below `0.0`.
/// * every `NaN` is one value and sorts above `+Infinity`, because `Double.compare` goes through
/// `doubleToLongBits`. Arrow uses the raw bits, so a sign-bit-set `NaN` would sort below
/// `-Infinity` instead.
///
/// Folding `-0.0` into `0.0` and every `NaN` into the canonical `NaN` makes the row bytes agree
/// with `compareDoubles` on both counts. This is verified identical on Spark 3.4 through master.
///
/// Note this is the opposite of what `mode` needs: `mode` keys a hash map via
/// `OpenHashSet`'s `equals` (`java.lang.Double.equals`), which distinguishes `-0.0` from `0.0`, so
/// it must *not* fold them. Same two input values, different Spark comparison path, opposite
/// correct behaviour.
fn canonicalize_float_ordering(array: &ArrayRef) -> ArrayRef {
match array.data_type() {
DataType::Float32 => Arc::new(array.as_primitive::<Float32Type>().unary::<_, Float32Type>(
|v| {
if v.is_nan() {
f32::NAN
} else if v == 0.0 {
// `-0.0 == 0.0` in IEEE 754, so this catches negative zero only.
0.0
} else {
v
}
},
)),
DataType::Float64 => Arc::new(array.as_primitive::<Float64Type>().unary::<_, Float64Type>(
|v| {
if v.is_nan() {
f64::NAN
} else if v == 0.0 {
0.0
} else {
v
}
},
)),
_ => Arc::clone(array),
}
}

/// Accumulator that tracks the running `(value, ordering)` pair for the extremum ordering.
#[derive(Debug)]
struct MaxMinByAccumulator {
Expand Down Expand Up @@ -231,7 +189,7 @@ impl MaxMinByAccumulator {
return Ok(());
}

let ordering_arr = canonicalize_float_ordering(ordering_arr);
let ordering_arr = normalize_floats(ordering_arr);
let rows = self
.ordering_converter
.convert_columns(&[Arc::clone(&ordering_arr)])?;
Expand Down Expand Up @@ -389,7 +347,7 @@ impl MaxMinByGroupsAccumulator {
let value_rows = self
.value_converter
.convert_columns(&[Arc::clone(&values[0])])?;
let ordering_arr = canonicalize_float_ordering(&values[1]);
let ordering_arr = normalize_floats(&values[1]);
let ordering_rows = self
.ordering_converter
.convert_columns(&[Arc::clone(&ordering_arr)])?;
Expand Down
23 changes: 9 additions & 14 deletions native/spark-expr/src/agg_funcs/mode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
* under the License.
*/

use crate::float_semantics::{canonicalize_nan, normalize_float};
use arrow::array::{Array, ArrayRef, AsArray, BooleanArray, StructArray};
use arrow::datatypes::{DataType, Field, FieldRef, Fields, Int64Type};
use datafusion::common::{internal_datafusion_err, not_impl_err, Result, ScalarValue};
Expand Down Expand Up @@ -171,21 +172,15 @@ impl AggregateUDFImpl for Mode {
/// equality, which is what Spark's `OpenHashSet` uses. `-0.0` is folded into `0.0` only when
/// `normalize_neg_zero` is set, i.e. only on Spark 4.2.0+ (SPARK-57329); see [`Mode`].
fn normalize_key(value: ScalarValue, normalize_neg_zero: bool) -> ScalarValue {
macro_rules! normalize_float {
($variant:path, $f:expr, $nan:expr) => {
if $f.is_nan() {
$variant(Some($nan))
} else if normalize_neg_zero && $f == 0.0 {
// `-0.0 == 0.0` in IEEE 754, so this catches negative zero only.
$variant(Some(0.0))
} else {
$variant(Some($f))
}
};
}
match value {
ScalarValue::Float32(Some(f)) => normalize_float!(ScalarValue::Float32, f, f32::NAN),
ScalarValue::Float64(Some(f)) => normalize_float!(ScalarValue::Float64, f, f64::NAN),
ScalarValue::Float32(Some(f)) if normalize_neg_zero => {
ScalarValue::Float32(Some(normalize_float(f)))
}
ScalarValue::Float64(Some(f)) if normalize_neg_zero => {
ScalarValue::Float64(Some(normalize_float(f)))
}
ScalarValue::Float32(Some(f)) => ScalarValue::Float32(Some(canonicalize_nan(f))),
ScalarValue::Float64(Some(f)) => ScalarValue::Float64(Some(canonicalize_nan(f))),
other => other,
}
}
Expand Down
Loading
Loading