diff --git a/datafusion/common/src/utils/mod.rs b/datafusion/common/src/utils/mod.rs index 107226b42bad1..5f0c18a4f3195 100644 --- a/datafusion/common/src/utils/mod.rs +++ b/datafusion/common/src/utils/mod.rs @@ -30,8 +30,8 @@ use crate::error::{ }; use crate::{Result, ScalarValue}; use arrow::array::{ - Array, ArrayRef, FixedSizeListArray, LargeListArray, ListArray, OffsetSizeTrait, - cast::AsArray, + Array, ArrayData, ArrayRef, FixedSizeListArray, LargeListArray, ListArray, + OffsetSizeTrait, cast::AsArray, downcast_array, }; use arrow::array::{ ArrowPrimitiveType, BooleanArray, Datum, GenericListArray, Int32Array, Int64Array, @@ -1432,9 +1432,9 @@ fn fsl_values_row_number(list_size: i32, array_len: usize) -> Result Ok(PrimitiveArray::new(rows_number.into(), None)) } -/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array. -/// For non-float arrays returns the input unchanged. NaN payloads are -/// preserved. +/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array, +/// including floats inside dictionaries and nested arrays. For other arrays, +/// returns the input unchanged. NaN payloads are preserved. /// /// Arrow's comparison kernels (`arrow::compute::kernels::cmp::eq` etc.) and /// row-encoding (`arrow::row::RowConverter`) use IEEE 754 totalOrder @@ -1544,21 +1544,38 @@ pub fn has_float_leaf(data_type: &DataType) -> bool { } /// Replace `-0.0` with `+0.0` in `Float16`, `Float32`, or `Float64` scalar -/// values. Other variants are returned unchanged. See [`normalize_float_zero`] -/// for context. -pub fn normalize_float_zero_scalar(scalar: ScalarValue) -> ScalarValue { - match scalar { - ScalarValue::Float32(Some(v)) if v.to_bits() << 1 == 0 => { - ScalarValue::Float32(Some(0.0)) - } - ScalarValue::Float64(Some(v)) if v.to_bits() << 1 == 0 => { - ScalarValue::Float64(Some(0.0)) - } - ScalarValue::Float16(Some(v)) if v.to_bits() << 1 == 0 => { - ScalarValue::Float16(Some(half::f16::from_bits(0))) +/// values, including floats inside nested and encoded values. Other variants +/// are returned unchanged. See [`normalize_float_zero`] for context. +pub fn normalize_float_zero_scalar(mut scalar: ScalarValue) -> ScalarValue { + fn normalize_array + 'static>(array: &mut Arc) { + *array = Arc::new(downcast_array( + normalize_float_zero(&(Arc::clone(array) as ArrayRef)).as_ref(), + )); + } + + fn normalize(scalar: &mut ScalarValue) { + match scalar { + ScalarValue::Float32(Some(v)) if v.to_bits() << 1 == 0 => *v = 0.0, + ScalarValue::Float64(Some(v)) if v.to_bits() << 1 == 0 => *v = 0.0, + ScalarValue::Float16(Some(v)) if v.to_bits() << 1 == 0 => { + *v = half::f16::from_bits(0); + } + ScalarValue::FixedSizeList(array) => normalize_array(array), + ScalarValue::List(array) => normalize_array(array), + ScalarValue::LargeList(array) => normalize_array(array), + ScalarValue::ListView(array) => normalize_array(array), + ScalarValue::LargeListView(array) => normalize_array(array), + ScalarValue::Struct(array) => normalize_array(array), + ScalarValue::Map(array) => normalize_array(array), + ScalarValue::Union(Some((_, value)), _, _) + | ScalarValue::Dictionary(_, value) + | ScalarValue::RunEndEncoded(_, _, value) => normalize(value), + _ => {} } - other => other, } + + normalize(&mut scalar); + scalar } /// Apply a struct's nulls to one of its fields. @@ -1611,8 +1628,8 @@ mod tests { use super::*; use crate::ScalarValue::Null; use arrow::{ - array::{Float64Array, Int32Array, NullArray}, - datatypes::Int32Type, + array::{DictionaryArray, Float64Array, Int8Array, Int32Array, NullArray}, + datatypes::{Float64Type, Int8Type, Int32Type}, }; #[cfg(feature = "sql")] use sqlparser::ast::Ident; @@ -1692,6 +1709,104 @@ mod tests { } } + #[test] + fn normalize_float_zero_in_dictionary_arrays_and_scalars() -> Result<()> { + let nan = f64::from_bits(0x7ff8_0000_0000_0001); + let keys = Int8Array::from(vec![Some(0), Some(1), None, Some(2)]); + let array: ArrayRef = Arc::new(DictionaryArray::try_new( + keys.clone(), + Arc::new(Float64Array::from(vec![-0.0, nan, 1.0])), + )?); + + let normalized = normalize_float_zero(&array); + let dictionary = normalized.as_dictionary::(); + assert_eq!(dictionary.keys(), &keys); + let values = dictionary.values().as_primitive::(); + assert_eq!(values.value(0).to_bits(), 0.0_f64.to_bits()); + assert_eq!(values.value(1).to_bits(), nan.to_bits()); + assert_eq!(values.value(2), 1.0); + + assert!(Arc::ptr_eq(&normalize_float_zero(&normalized), &normalized)); + + assert_eq!( + normalize_float_zero_scalar(ScalarValue::try_from_array(&array, 0)?), + ScalarValue::try_from_array(&normalized, 0)? + ); + + Ok(()) + } + + #[test] + fn normalize_float_zero_scalar_preserves_union_and_run_metadata() -> Result<()> { + use arrow::datatypes::{UnionFields, UnionMode}; + + let run_ends = Arc::new(Field::new("ends", DataType::Int32, false)); + let values = Arc::new(Field::new("samples", DataType::Float64, true)); + let fields = UnionFields::try_new( + [7, 42], + [ + Field::new( + "runs", + DataType::RunEndEncoded(Arc::clone(&run_ends), Arc::clone(&values)), + true, + ), + Field::new("other", DataType::Int32, true), + ], + )?; + let nan = f64::from_bits(0x7ff8_0000_0000_0001); + for mode in [UnionMode::Sparse, UnionMode::Dense] { + let wrap = |value| { + ScalarValue::Union( + Some(( + 7, + Box::new(ScalarValue::RunEndEncoded( + Arc::clone(&run_ends), + Arc::clone(&values), + Box::new(ScalarValue::Float64(value)), + )), + )), + fields.clone(), + mode, + ) + }; + for (value, expected) in [ + (Some(-0.0), Some(0.0)), + (Some(nan), Some(nan)), + (None, None), + ] { + assert_eq!(normalize_float_zero_scalar(wrap(value)), wrap(expected)); + } + let null = ScalarValue::Union(None, fields.clone(), mode); + assert_eq!(normalize_float_zero_scalar(null.clone()), null); + } + Ok(()) + } + + #[test] + fn normalize_float_zero_scalar_preserves_sliced_list() { + let nan = f64::from_bits(0x7ff8_0000_0000_0001); + let list = Arc::new( + ListArray::from_iter_primitive::([ + Some(vec![Some(99.0)]), + Some(vec![Some(-0.0), Some(nan), None]), + Some(vec![Some(42.0)]), + ]) + .slice(1, 1), + ); + let ScalarValue::List(normalized) = + normalize_float_zero_scalar(ScalarValue::List(list)) + else { + panic!("normalization must preserve the scalar variant"); + }; + assert_eq!(normalized.len(), 1); + let child = normalized.value(0); + let child = child.as_primitive::(); + assert_eq!(child.len(), 3); + assert_eq!(child.value(0).to_bits(), 0.0_f64.to_bits()); + assert_eq!(child.value(1).to_bits(), nan.to_bits()); + assert!(child.is_null(2)); + } + #[test] fn test_bisect_linear_left_and_right() -> Result<()> { let arrays: Vec = vec![ diff --git a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs index 506692f0dfa6d..974c95452a103 100644 --- a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs +++ b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs @@ -30,6 +30,7 @@ use std::sync::LazyLock; use datafusion_common::config::ConfigOptions; use datafusion_common::nested_struct::has_one_of_more_common_fields; +use datafusion_common::utils::has_float_leaf; use datafusion_common::{ DFSchema, DataFusionError, Result, ScalarValue, exec_datafusion_err, internal_err, }; @@ -2318,8 +2319,9 @@ fn are_inlist_and_eq_and_match_neg( fn inlists_have_set_comparable_literals(left: &Expr, right: &Expr) -> bool { match (left, right) { (Expr::InList(l), Expr::InList(r)) => l.list.iter().chain(&r.list).all(|item| { - item.as_literal() - .is_some_and(|value| !value.is_null() && !value.data_type().is_floating()) + item.as_literal().is_some_and(|value| { + !value.is_null() && !has_float_leaf(&value.data_type()) + }) }), _ => false, } diff --git a/datafusion/physical-expr/src/expressions/in_list.rs b/datafusion/physical-expr/src/expressions/in_list.rs index 7a7cb25317c1d..56aff29650af1 100644 --- a/datafusion/physical-expr/src/expressions/in_list.rs +++ b/datafusion/physical-expr/src/expressions/in_list.rs @@ -31,6 +31,7 @@ use arrow::compute::kernels::boolean::{not, or_kleene}; use arrow::compute::kernels::cmp::eq as arrow_eq; use arrow::datatypes::*; +use datafusion_common::utils::{normalize_float_zero, normalize_float_zero_scalar}; use datafusion_common::{ DFSchema, Result, ScalarValue, assert_or_internal_err, exec_err, }; @@ -82,6 +83,15 @@ fn supports_arrow_eq(dt: &DataType) -> bool { } } +fn normalize_in_list_float_zero_value(value: ColumnarValue) -> ColumnarValue { + match value { + ColumnarValue::Array(array) => ColumnarValue::Array(normalize_float_zero(&array)), + ColumnarValue::Scalar(scalar) => { + ColumnarValue::Scalar(normalize_float_zero_scalar(scalar)) + } + } +} + /// Evaluates the list of expressions into an array, flattening any dictionaries fn evaluate_list( list: &[Arc], @@ -370,12 +380,15 @@ impl PhysicalExpr for InListExpr { // Use Arrow's vectorized eq kernel for types it supports (primitive, // boolean, string, binary, dictionary), falling back to row-by-row // comparator for unsupported types (nested, RunEndEncoded, etc.). - let value = value.into_array(num_rows)?; + // Normalize the left side once for the whole list. Doing this + // outside `compare_one` avoids rescanning it for every item. + let value = + normalize_in_list_float_zero_value(value).into_array(num_rows)?; let lhs_supports_arrow_eq = supports_arrow_eq(value.data_type()); // Helper: compare value against a single list expression let compare_one = |expr: &Arc| -> Result { - match expr.evaluate(batch)? { + match normalize_in_list_float_zero_value(expr.evaluate(batch)?) { ColumnarValue::Array(array) => { if lhs_supports_arrow_eq && supports_arrow_eq(array.data_type()) @@ -3364,6 +3377,44 @@ mod tests { Ok(()) } + #[test] + fn test_in_list_with_columns_float_signed_zero() -> Result<()> { + use arrow::compute::cast; + + for (data_type, dictionary) in [ + (DataType::Float32, false), + (DataType::Float64, false), + (DataType::Float64, true), + ] { + let mut left = cast(&Float64Array::from(vec![0.0, -0.0, 1.0]), &data_type)?; + let mut right = cast(&Float64Array::from(vec![-0.0, 0.0, 2.0]), &data_type)?; + if dictionary { + left = wrap_in_dict(left); + right = wrap_in_dict(right); + } + let scalar = lit(ScalarValue::try_from_array(left.as_ref(), 1)?); + let batch = RecordBatch::try_from_iter([("a", left), ("b", right)])?; + let schema = batch.schema(); + + // Exercise array and scalar values on both sides, including scalar + // normalization before the left-hand side is broadcast. + for (left, right) in [ + (col("a", &schema)?, col("b", &schema)?), + (col("a", &schema)?, Arc::clone(&scalar)), + (scalar, col("a", &schema)?), + ] { + let expr = make_in_list_with_columns(left, vec![right], false); + let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?; + assert_eq!( + as_boolean_array(&result), + &BooleanArray::from(vec![true, true, false]), + "{data_type:?}, dictionary={dictionary}: {expr}" + ); + } + } + Ok(()) + } + /// Tests that short-circuit evaluation produces correct results. /// When all rows match after the first list item, remaining items /// should be skipped without affecting correctness. @@ -3889,6 +3940,23 @@ mod tests { Ok(()) } + #[test] + fn test_try_new_from_array_dict_haystack_float64_signed_zero() -> Result<()> { + // One value beyond the branchless limit selects the hash-set strategy. + let list_len = + ::MAX_LIST_LEN + 1; + let haystack = make_f64_dict_array(vec![Some(-0.0); list_len]); + let needles: ArrayRef = Arc::new(Float64Array::from(vec![0.0])); + for needles in [Arc::clone(&needles), wrap_in_dict(needles)] { + assert_eq!( + eval_in_list_from_array(needles, Arc::clone(&haystack))?, + BooleanArray::from(vec![true]) + ); + } + + Ok(()) + } + #[test] fn test_try_new_from_array_type_mismatch_rejects() -> Result<()> { let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]); diff --git a/datafusion/physical-expr/src/expressions/in_list/array_static_filter.rs b/datafusion/physical-expr/src/expressions/in_list/array_static_filter.rs index c2d3b4728274e..0aef552d88e8f 100644 --- a/datafusion/physical-expr/src/expressions/in_list/array_static_filter.rs +++ b/datafusion/physical-expr/src/expressions/in_list/array_static_filter.rs @@ -15,13 +15,14 @@ // specific language governing permissions and limitations // under the License. -use arrow::array::{Array, ArrayRef, BooleanArray, make_comparator}; +use arrow::array::{Array, ArrayRef, BooleanArray, make_array, make_comparator}; use arrow::buffer::{BooleanBuffer, NullBuffer}; use arrow::compute::SortOptions; use arrow::datatypes::DataType; use arrow::util::bit_iterator::BitIndexIterator; use datafusion_common::Result; use datafusion_common::hash_utils::{RandomState, with_hashes}; +use datafusion_common::utils::{has_float_leaf, normalize_float_zero}; use hashbrown::HashTable; use super::result::build_in_list_result; @@ -53,6 +54,8 @@ impl ArrayStaticFilter { }); } + // Hashing treats both signed zeros alike; the comparator must do so too. + let in_array = normalize_float_zero(&in_array); let state = RandomState::default(); let table = Self::build_haystack_table(&in_array, &state)?; @@ -138,6 +141,11 @@ impl StaticFilter for ArrayStaticFilter { )); } - self.find_needles_in_haystack(v, negated) + if has_float_leaf(v.data_type()) { + let normalized = normalize_float_zero(&make_array(v.to_data())); + self.find_needles_in_haystack(normalized.as_ref(), negated) + } else { + self.find_needles_in_haystack(v, negated) + } } } diff --git a/datafusion/physical-expr/src/expressions/in_list/branchless_filter.rs b/datafusion/physical-expr/src/expressions/in_list/branchless_filter.rs index 86cea37b3a98d..b92eb869ccae4 100644 --- a/datafusion/physical-expr/src/expressions/in_list/branchless_filter.rs +++ b/datafusion/physical-expr/src/expressions/in_list/branchless_filter.rs @@ -38,8 +38,9 @@ //! `Float32` and a `UInt32` both use four bytes per value. The filter compares //! those stored bits through an unsigned type of the same size, without copying //! the value buffer. A bit pattern is simply the bytes Arrow uses to store a -//! value. Comparing it preserves details such as `0.0` versus `-0.0` and -//! different NaN values. [`BranchlessFilterType`] defines these safe, +//! value. Float comparisons treat `0.0` and `-0.0` as equal under SQL +//! semantics, while preserving distinctions between different NaN values. +//! [`BranchlessFilterType`] defines these safe, //! same-sized mappings and checks their sizes at compile time. //! //! The fast path is limited to short lists: @@ -75,6 +76,7 @@ use arrow::buffer::{BooleanBuffer, ScalarBuffer}; use arrow::datatypes::*; use arrow::util::bit_iterator::BitIndexIterator; use datafusion_common::{Result, exec_datafusion_err, internal_datafusion_err}; +use half::f16; use super::result::build_result_from_contains; use super::static_filter::StaticFilter; @@ -117,20 +119,24 @@ const BRANCHLESS_MAX_16B: usize = 4; /// /// `T` is the logical Arrow type accepted by the filter. `CompareType` is the /// same-width type used for the fixed comparison chain. Signed integers, -/// floats, and temporal values use an unsigned comparison type so they compare -/// by their raw bit pattern. +/// floats, and temporal values use an unsigned comparison type. Most values +/// compare by raw bit pattern; floats additionally treat both signed-zero +/// patterns as equal. pub(super) trait BranchlessFilterType: - ArrowPrimitiveType + Send + Sync + 'static + ArrowPrimitiveType + Send + Sync + Sized + 'static { type CompareType: ArrowPrimitiveType + Send + Sync + 'static; /// Maximum number of non-null IN-list values to handle with /// [`BranchlessFilter`] for this primitive type. const MAX_LIST_LEN: usize; + + /// The two signed-zero encodings for float types. + const SIGNED_ZERO_BITS: Option<[BranchlessNative; 2]> = None; } macro_rules! branchless_filter_type { - ($logical:ty, $compare:ty, $max_len:expr) => { + ($logical:ty, $compare:ty, $max_len:expr $(, $zero_bits:expr)?) => { // The branchless filter reads the same Arrow value buffer as the // comparison type. That is only valid when both native types have the // same width, so catch any bad mapping here at compile time. @@ -143,6 +149,7 @@ macro_rules! branchless_filter_type { impl BranchlessFilterType for $logical { type CompareType = $compare; const MAX_LIST_LEN: usize = $max_len; + $(const SIGNED_ZERO_BITS: Option<[BranchlessNative; 2]> = Some($zero_bits);)? } }; } @@ -151,18 +158,33 @@ branchless_filter_type!(Int8Type, UInt8Type, BRANCHLESS_MAX_1B); branchless_filter_type!(UInt8Type, UInt8Type, BRANCHLESS_MAX_1B); branchless_filter_type!(Int16Type, UInt16Type, BRANCHLESS_MAX_2B); branchless_filter_type!(UInt16Type, UInt16Type, BRANCHLESS_MAX_2B); -branchless_filter_type!(Float16Type, UInt16Type, BRANCHLESS_MAX_2B); +branchless_filter_type!( + Float16Type, + UInt16Type, + BRANCHLESS_MAX_2B, + [f16::ZERO.to_bits(), f16::NEG_ZERO.to_bits()] +); branchless_filter_type!(Int32Type, UInt32Type, BRANCHLESS_MAX_4B); branchless_filter_type!(UInt32Type, UInt32Type, BRANCHLESS_MAX_4B); -branchless_filter_type!(Float32Type, UInt32Type, BRANCHLESS_MAX_4B); +branchless_filter_type!( + Float32Type, + UInt32Type, + BRANCHLESS_MAX_4B, + [0.0_f32.to_bits(), (-0.0_f32).to_bits()] +); branchless_filter_type!(Date32Type, UInt32Type, BRANCHLESS_MAX_4B); branchless_filter_type!(Time32SecondType, UInt32Type, BRANCHLESS_MAX_4B); branchless_filter_type!(Time32MillisecondType, UInt32Type, BRANCHLESS_MAX_4B); branchless_filter_type!(Int64Type, UInt64Type, BRANCHLESS_MAX_8B); branchless_filter_type!(UInt64Type, UInt64Type, BRANCHLESS_MAX_8B); -branchless_filter_type!(Float64Type, UInt64Type, BRANCHLESS_MAX_8B); +branchless_filter_type!( + Float64Type, + UInt64Type, + BRANCHLESS_MAX_8B, + [0.0_f64.to_bits(), (-0.0_f64).to_bits()] +); branchless_filter_type!(Date64Type, UInt64Type, BRANCHLESS_MAX_8B); branchless_filter_type!(Time64MicrosecondType, UInt64Type, BRANCHLESS_MAX_8B); branchless_filter_type!(Time64NanosecondType, UInt64Type, BRANCHLESS_MAX_8B); @@ -218,8 +240,10 @@ where } let all_values = branchless_values::(in_array); - let mut in_list_values = Vec::with_capacity(non_null_count); - + // Float zero may add its other signed encoding. + let mut in_list_values = Vec::with_capacity( + non_null_count + usize::from(T::SIGNED_ZERO_BITS.is_some()), + ); match in_array.nulls() { None => { in_list_values.extend(all_values.iter().copied()); @@ -234,6 +258,18 @@ where } debug_assert_eq!(in_list_values.len(), non_null_count); + + // Add the other signed zero once so per-row lookup stays branch-free. + if let Some([positive_zero, negative_zero]) = T::SIGNED_ZERO_BITS { + match ( + in_list_values.contains(&positive_zero), + in_list_values.contains(&negative_zero), + ) { + (true, false) => in_list_values.push(negative_zero), + (false, true) => in_list_values.push(positive_zero), + _ => {} + } + } let in_list_values = in_list_values.into_boxed_slice(); let check_values = membership_check_for_len::(in_list_values.len()); @@ -290,13 +326,23 @@ where T: BranchlessFilterType, BranchlessNative: Copy + PartialEq, { + /// Selects a fixed-size comparison function for the enclosing `len` and `T`. + /// + /// Arguments must enumerate every length from zero through the logical limit. + /// Their count is therefore one past that limit: the extra length supported + /// for floats when adding the other signed-zero encoding. + /// + /// The function is selected once during filter construction. macro_rules! choose { - ($($n:literal),* $(,)?) => { + ($($n:literal),+ $(,)?) => {{ + const EXTRA_LEN: usize = [$($n),+].len(); match len { - $($n => check_values::, $n>,)* + $($n => check_values::, $n>,)+ + EXTRA_LEN if T::SIGNED_ZERO_BITS.is_some() => + check_values::, EXTRA_LEN>, _ => unreachable!("list length exceeds the configured limit"), } - }; + }}; } // Avoid creating checks for lengths a type does not support. @@ -451,11 +497,11 @@ mod tests { assert_eq!( filter.contains(&needles, false)?, - BooleanArray::from(vec![None, Some(true), Some(true), None, None]) + BooleanArray::from(vec![Some(true), Some(true), Some(true), None, None]) ); assert_eq!( filter.contains(&needles, true)?, - BooleanArray::from(vec![None, Some(false), Some(false), None, None]) + BooleanArray::from(vec![Some(false), Some(false), Some(false), None, None]) ); let wrong_type = UInt16Array::from(vec![Some(0x8000), Some(0x7e01)]); @@ -466,20 +512,26 @@ mod tests { } #[test] - fn branchless_filter_floats_use_bit_equality() -> Result<()> { + fn branchless_filter_floats_use_sql_zero_equality() -> Result<()> { let nan_a = f32::from_bits(0x7fc0_0001); let nan_b = f32::from_bits(0x7fc0_0002); let haystack: ArrayRef = - Arc::new(Float32Array::from(vec![Some(-0.0), Some(nan_a)])); + Arc::new(Float32Array::from(vec![Some(0.0), Some(nan_a)])); let filter = BranchlessFilter::::try_new(&haystack)?; let needles = Float32Array::from(vec![Some(0.0), Some(-0.0), Some(nan_a), Some(nan_b)]); assert_eq!( filter.contains(&needles, false)?, - BooleanArray::from(vec![Some(false), Some(true), Some(true), Some(false)]) + BooleanArray::from(vec![Some(true), Some(true), Some(true), Some(false)]) ); + // A list containing both encodings is not expanded. + let zero_only: ArrayRef = + Arc::new(Float32Array::from(vec![Some(0.0), Some(-0.0)])); + let filter = BranchlessFilter::::try_new(&zero_only)?; + assert_eq!(filter.in_list_values.len(), 2); + let nan_a = f64::from_bits(0x7ff8_0000_0000_0001); let nan_b = f64::from_bits(0x7ff8_0000_0000_0002); let haystack: ArrayRef = @@ -490,7 +542,7 @@ mod tests { assert_eq!( filter.contains(&needles, false)?, - BooleanArray::from(vec![Some(false), Some(true), Some(true), Some(false)]) + BooleanArray::from(vec![Some(true), Some(true), Some(true), Some(false)]) ); Ok(()) diff --git a/datafusion/physical-expr/src/expressions/in_list/primitive_filter.rs b/datafusion/physical-expr/src/expressions/in_list/primitive_filter.rs index ceb9bd525b965..a3273c3eccc20 100644 --- a/datafusion/physical-expr/src/expressions/in_list/primitive_filter.rs +++ b/datafusion/physical-expr/src/expressions/in_list/primitive_filter.rs @@ -27,6 +27,7 @@ use arrow::array::{Array, ArrayRef, AsArray, BooleanArray}; use arrow::datatypes::*; use arrow::util::bit_iterator::BitIndexIterator; use datafusion_common::{HashSet, Result, exec_datafusion_err}; +use half::f16; use super::branchless_filter::{BranchlessFilter, BranchlessFilterType}; use super::result::build_in_list_result; @@ -199,6 +200,9 @@ trait BitmapFilterType: ArrowPrimitiveType + Send + Sync + 'static { /// Returns the index in the bitmap to check for this value. fn index(value: Self::Native) -> usize; + + /// Bitmap indices for the two signed-zero encodings of a float type. + const SIGNED_ZERO_INDICES: Option<[usize; 2]> = None; } /// `Int8` has 256 possible bit patterns, so four `u64` words cover the full domain. @@ -254,6 +258,11 @@ impl BitmapFilterType for Float16Type { fn index(value: Self::Native) -> usize { value.to_bits() as usize } + + const SIGNED_ZERO_INDICES: Option<[usize; 2]> = Some([ + f16::ZERO.to_bits() as usize, + f16::NEG_ZERO.to_bits() as usize, + ]); } /// `IN` filter backed by one bit per possible value. @@ -291,6 +300,13 @@ where } } } + // Store both signed zeros so per-row lookup stays branch-free. + if let Some([positive_zero, negative_zero]) = T::SIGNED_ZERO_INDICES + && (bits.get_bit(positive_zero) || bits.get_bit(negative_zero)) + { + bits.set_bit(positive_zero); + bits.set_bit(negative_zero); + } Ok(Self { null_count: prim_array.null_count(), bits, @@ -332,6 +348,13 @@ where } } +/// Hash keys and the two signed-zero encodings for float types. +trait HashSetKey: From + Eq + Hash + Sized { + const SIGNED_ZERO_KEYS: Option<[Self; 2]> = None; +} + +impl HashSetKey for T where T: Copy + Eq + Hash {} + /// Wrapper for f32 that implements Hash and Eq using bit comparison. /// This treats NaN values as equal to each other when they have the same bit pattern. #[derive(Clone, Copy)] @@ -357,6 +380,10 @@ impl From for OrderedFloat32 { } } +impl HashSetKey for OrderedFloat32 { + const SIGNED_ZERO_KEYS: Option<[Self; 2]> = Some([Self(0.0), Self(-0.0)]); +} + /// Wrapper for f64 that implements Hash and Eq using bit comparison. /// This treats NaN values as equal to each other when they have the same bit pattern. #[derive(Clone, Copy)] @@ -382,6 +409,10 @@ impl From for OrderedFloat64 { } } +impl HashSetKey for OrderedFloat64 { + const SIGNED_ZERO_KEYS: Option<[Self; 2]> = Some([Self(0.0), Self(-0.0)]); +} + /// Hash-set membership for primitive types. /// /// `K` defaults to the Arrow type's native value. Floats use an ordered wrapper @@ -399,7 +430,7 @@ impl PrimitiveHashSetFilter where T: ArrowPrimitiveType, T::Native: Copy, - K: From + Eq + Hash, + K: HashSetKey, { fn try_new(in_array: &ArrayRef) -> Result { let in_array = in_array.as_primitive_opt::().ok_or_else(|| { @@ -413,6 +444,13 @@ where values.insert(K::from(value)); } + // Store both signed zeros so per-row lookup stays branch-free. + if let Some([positive_zero, negative_zero]) = K::SIGNED_ZERO_KEYS + && (values.contains(&positive_zero) || values.contains(&negative_zero)) + { + values.insert(positive_zero); + values.insert(negative_zero); + } Ok(Self { null_count: in_array.null_count(), values, @@ -459,9 +497,8 @@ mod tests { use arrow::array::{ DictionaryArray, Float16Array, Float32Array, Float64Array, Int8Array, Int16Array, - UInt8Array, UInt16Array, UInt32Array, + PrimitiveArray, UInt8Array, UInt16Array, UInt32Array, }; - use half::f16; use super::super::dictionary_filter::DictionaryFilter; @@ -507,6 +544,28 @@ mod tests { Ok(()) } + #[test] + fn branchless_float_zero_expansion_handles_max_list_len() -> Result<()> { + fn assert_routed_filter( + negative_zero: T::Native, + positive_zero: T::Native, + ) -> Result<()> { + // The mirror zero is appended after the full logical list. + let haystack: ArrayRef = Arc::new(PrimitiveArray::::from_value( + negative_zero, + T::MAX_LIST_LEN, + )); + let filter = instantiate_branchless_filter(&haystack)? + .expect("a full float list still uses a branchless filter"); + let needles = PrimitiveArray::::from_value(positive_zero, 1); + assert_contains(filter.as_ref(), &needles, vec![Some(true)]) + } + + assert_routed_filter::(f16::NEG_ZERO, f16::ZERO)?; + assert_routed_filter::(-0.0, 0.0)?; + assert_routed_filter::(-0.0, 0.0) + } + #[test] fn primitive_hash_filter_handles_float_keys() -> Result<()> { let nan32 = f32::NAN; @@ -524,14 +583,14 @@ mod tests { assert_contains( &filter, &needles, - vec![Some(true), Some(false), Some(true), Some(false), None], + vec![Some(true), Some(true), Some(true), Some(false), None], )?; let nan64 = f64::NAN; - let haystack: ArrayRef = Arc::new(Float64Array::from(vec![1.0, nan64])); + let haystack: ArrayRef = Arc::new(Float64Array::from(vec![-0.0, nan64])); let filter = PrimitiveHashSetFilter::::try_new(&haystack)?; - let needles = Float64Array::from(vec![Some(1.0), Some(nan64), Some(2.0)]); + let needles = Float64Array::from(vec![Some(0.0), Some(nan64), Some(2.0)]); assert_contains(&filter, &needles, vec![Some(true), Some(true), Some(false)]) } @@ -660,8 +719,8 @@ mod tests { ); let filter = BitmapFilter::::try_new(&haystack)?; let needles = Float16Array::from(vec![ - Some(f16::from_f32(0.0)), Some(f16::from_f32(-0.0)), + Some(f16::from_f32(0.0)), Some(nan_a), Some(nan_b), None, diff --git a/datafusion/physical-expr/src/expressions/in_list/strategy.rs b/datafusion/physical-expr/src/expressions/in_list/strategy.rs index 5cbdf8cd7486a..af13debb5559b 100644 --- a/datafusion/physical-expr/src/expressions/in_list/strategy.rs +++ b/datafusion/physical-expr/src/expressions/in_list/strategy.rs @@ -63,7 +63,7 @@ fn view_types_match(needle_type: &DataType, list_type: &DataType) -> bool { && dictionary_value_type(needle_type) == list_type } -fn dictionary_value_type(mut data_type: &DataType) -> &DataType { +pub(super) fn dictionary_value_type(mut data_type: &DataType) -> &DataType { while let DataType::Dictionary(_, value_type) = data_type { data_type = value_type; } diff --git a/datafusion/sqllogictest/test_files/negative_zero.slt b/datafusion/sqllogictest/test_files/negative_zero.slt index fe4df57e179cc..99b28c9a9ba41 100644 --- a/datafusion/sqllogictest/test_files/negative_zero.slt +++ b/datafusion/sqllogictest/test_files/negative_zero.slt @@ -47,6 +47,52 @@ SELECT 0.0 IS DISTINCT FROM -0.0 AS is_distinct; ---- false +##### +## IN / NOT IN predicates +##### + +statement ok +CREATE TABLE negative_zero_in_list AS +SELECT arrow_cast(-0.0, 'Float16') AS negative_f16, + arrow_cast(0.0, 'Float32') AS positive_f32, + 0.0 AS positive_f64, + -0.0 AS negative_f64, + arrow_cast(-0.0, 'Dictionary(Int32, Float64)') AS negative_dict_f64; + +# Three-item lists become comparisons; four-item lists remain InList. +# Check IN and NOT IN across float widths and both zero signs. +query BBBBB +SELECT + positive_f64 IN (-0.0, 1, 2), + negative_f16 IN (arrow_cast(0.0, 'Float16'), 1, 2, 3), + positive_f32 IN (arrow_cast(-0.0, 'Float32'), 1, 2, 3), + positive_f64 IN (-0.0, 1, 2, 3), + positive_f64 NOT IN (-0.0, 1, 2, 3) +FROM negative_zero_in_list; +---- +true true true true false + +# Dictionary encoding must preserve equality on both sides of the rewrite threshold. +query BB +SELECT + negative_dict_f64 IN (0.0), + negative_dict_f64 IN (0.0, 1, 2, 3) +FROM negative_zero_in_list; +---- +true true + +# A column-valued list item forces non-static evaluation. Check both directions. +query BB +SELECT + positive_f64 IN (negative_f64, 1, 2, 3), + negative_f64 IN (positive_f64, 1, 2, 3) +FROM negative_zero_in_list; +---- +true true + +statement ok +DROP TABLE negative_zero_in_list; + ##### ## SELECT DISTINCT with +0.0 / -0.0 (Float64) ##### @@ -244,13 +290,74 @@ true false true statement ok CREATE TABLE nested_zeros(id INT, a DOUBLE[]) AS VALUES (1, [0.0]), (2, [-0.0]); -query II -SELECT l.id, r.id FROM nested_zeros l JOIN nested_zeros r ON l.a = r.a ORDER BY l.id, r.id; +# Column-valued list items exercise dynamic array comparisons in both directions. +query IIBB +SELECT l.id, r.id, + l.a IN (r.a, [1.0], [2.0], [3.0]), + l.a NOT IN (r.a, [1.0], [2.0], [3.0]) +FROM nested_zeros l JOIN nested_zeros r ON l.a = r.a +ORDER BY l.id, r.id; +---- +1 1 true false +1 2 true false +2 1 true false +2 2 true false + +# Short lists become equality comparisons; four-item constant lists use a static filter. +query IBBBB +SELECT id, + a IN ([0.0]), + a IN ([0.0], [1.0], [2.0], [3.0]), + a IN ([-0.0], [1.0], [2.0], [3.0]), + a NOT IN ([-0.0], [1.0], [2.0], [3.0]) +FROM nested_zeros ORDER BY id; +---- +1 true true true false +2 true true true false + +# The nonmatching column-derived item keeps scalar RHS comparisons dynamic. +# Also check scalar LHS values against array-valued list items. +query IBBBB +SELECT id, + a IN ([-0.0], [1.0], [2.0], [CAST(id AS DOUBLE)]), + a NOT IN ([0.0], [1.0], [2.0], [CAST(id AS DOUBLE)]), + [-0.0] IN (a, [1.0], [2.0], [3.0]), + [0.0] NOT IN (a, [1.0], [2.0], [3.0]) +FROM nested_zeros ORDER BY id; +---- +1 true false true false +2 true false true false + +# Structural literal equality must not drive IN intersection or difference for nested floats. +query IBB +SELECT id, + a IN ([0.0], [1.0], [2.0], [3.0]) + AND a IN ([-0.0], [4.0], [5.0], [6.0]), + a IN ([0.0], [1.0], [2.0], [3.0]) + AND a NOT IN ([-0.0], [4.0], [5.0], [6.0]) +FROM nested_zeros ORDER BY id; +---- +1 true false +2 true false + +statement ok +INSERT INTO nested_zeros VALUES (3, NULL), (4, [NULL]), (5, [9.0]); + +# Matches win over an outer NULL list item; nested NULL elements remain comparable. +# A missing match or a NULL input produces NULL, for both static and dynamic lists. +query IBBBB +SELECT id, + a IN ([-0.0], [NULL], [1.0], NULL), + a NOT IN ([-0.0], [NULL], [1.0], NULL), + a IN ([-0.0], [NULL], [CAST(id AS DOUBLE)], NULL), + a NOT IN ([-0.0], [NULL], [CAST(id AS DOUBLE)], NULL) +FROM nested_zeros ORDER BY id; ---- -1 1 -1 2 -2 1 -2 2 +1 true false true false +2 true false true false +3 NULL NULL NULL NULL +4 true false true false +5 NULL NULL NULL NULL statement ok DROP TABLE nested_zeros;