Skip to content

Commit 474474f

Browse files
Treat signed zeros as equal in IN LIST filters
1 parent 28809d8 commit 474474f

5 files changed

Lines changed: 566 additions & 53 deletions

File tree

‎datafusion/common/src/utils/mod.rs‎

Lines changed: 102 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -1407,9 +1407,9 @@ fn fsl_values_row_number(list_size: i32, array_len: usize) -> Result<Int32Array>
14071407
Ok(PrimitiveArray::new(rows_number.into(), None))
14081408
}
14091409

1410-
/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array.
1411-
/// For non-float arrays returns the input unchanged. NaN payloads are
1412-
/// preserved.
1410+
/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array,
1411+
/// including dictionary-wrapped floats. For other arrays, returns the input
1412+
/// unchanged. NaN payloads are preserved.
14131413
///
14141414
/// Arrow's comparison kernels (`arrow::compute::kernels::cmp::eq` etc.) and
14151415
/// row-encoding (`arrow::row::RowConverter`) use IEEE 754 totalOrder
@@ -1430,6 +1430,17 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14301430
const NEG_ZERO_F32_BITS: u32 = (-0.0_f32).to_bits();
14311431
const NEG_ZERO_F64_BITS: u64 = (-0.0_f64).to_bits();
14321432
match array.data_type() {
1433+
DataType::Dictionary(_, value_type)
1434+
if is_float_or_dictionary_float(value_type) =>
1435+
{
1436+
let dictionary = array.as_any_dictionary();
1437+
let values = normalize_float_zero(dictionary.values());
1438+
if Arc::ptr_eq(&values, dictionary.values()) {
1439+
Arc::clone(array)
1440+
} else {
1441+
dictionary.with_values(values)
1442+
}
1443+
}
14331444
DataType::Float32 => {
14341445
let arr: &Float32Array = array.as_primitive::<Float32Type>();
14351446
if !arr
@@ -1439,8 +1450,13 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14391450
{
14401451
return Arc::clone(array);
14411452
}
1442-
let normalized: Float32Array =
1443-
arr.unary(|v| if v.to_bits() << 1 == 0 { 0.0_f32 } else { v });
1453+
let normalized: Float32Array = arr.unary(|v| {
1454+
if v.to_bits() == NEG_ZERO_F32_BITS {
1455+
0.0_f32
1456+
} else {
1457+
v
1458+
}
1459+
});
14441460
Arc::new(normalized)
14451461
}
14461462
DataType::Float64 => {
@@ -1452,8 +1468,13 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14521468
{
14531469
return Arc::clone(array);
14541470
}
1455-
let normalized: Float64Array =
1456-
arr.unary(|v| if v.to_bits() << 1 == 0 { 0.0_f64 } else { v });
1471+
let normalized: Float64Array = arr.unary(|v| {
1472+
if v.to_bits() == NEG_ZERO_F64_BITS {
1473+
0.0_f64
1474+
} else {
1475+
v
1476+
}
1477+
});
14571478
Arc::new(normalized)
14581479
}
14591480
DataType::Float16 => {
@@ -1466,8 +1487,8 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14661487
return Arc::clone(array);
14671488
}
14681489
let normalized: Float16Array = arr.unary(|v| {
1469-
if v.to_bits() << 1 == 0 {
1470-
half::f16::from_bits(0)
1490+
if v.to_bits() == NEG_ZERO_F16_BITS {
1491+
half::f16::ZERO
14711492
} else {
14721493
v
14731494
}
@@ -1478,22 +1499,38 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14781499
}
14791500
}
14801501

1502+
fn is_float_or_dictionary_float(mut data_type: &DataType) -> bool {
1503+
while let DataType::Dictionary(_, value_type) = data_type {
1504+
data_type = value_type;
1505+
}
1506+
data_type.is_floating()
1507+
}
1508+
14811509
/// Replace `-0.0` with `+0.0` in `Float16`, `Float32`, or `Float64` scalar
1482-
/// values. Other variants are returned unchanged. See [`normalize_float_zero`]
1483-
/// for context.
1484-
pub fn normalize_float_zero_scalar(scalar: ScalarValue) -> ScalarValue {
1485-
match scalar {
1486-
ScalarValue::Float32(Some(v)) if v.to_bits() << 1 == 0 => {
1487-
ScalarValue::Float32(Some(0.0))
1510+
/// values, including dictionary-wrapped floats. Other variants are returned
1511+
/// unchanged. See [`normalize_float_zero`] for context.
1512+
pub fn normalize_float_zero_scalar(mut scalar: ScalarValue) -> ScalarValue {
1513+
let mut value = &mut scalar;
1514+
while let ScalarValue::Dictionary(_, dictionary_value) = value {
1515+
value = dictionary_value.as_mut();
1516+
}
1517+
1518+
match value {
1519+
ScalarValue::Float32(Some(value)) if value.to_bits() == (-0.0_f32).to_bits() => {
1520+
*value = 0.0
14881521
}
1489-
ScalarValue::Float64(Some(v)) if v.to_bits() << 1 == 0 => {
1490-
ScalarValue::Float64(Some(0.0))
1522+
ScalarValue::Float64(Some(value)) if value.to_bits() == (-0.0_f64).to_bits() => {
1523+
*value = 0.0
14911524
}
1492-
ScalarValue::Float16(Some(v)) if v.to_bits() << 1 == 0 => {
1493-
ScalarValue::Float16(Some(half::f16::from_bits(0)))
1525+
ScalarValue::Float16(Some(value))
1526+
if value.to_bits() == half::f16::NEG_ZERO.to_bits() =>
1527+
{
1528+
*value = half::f16::ZERO;
14941529
}
1495-
other => other,
1530+
_ => {}
14961531
}
1532+
1533+
scalar
14971534
}
14981535

14991536
#[cfg(test)]
@@ -1503,9 +1540,9 @@ mod tests {
15031540
use super::*;
15041541
use crate::ScalarValue::Null;
15051542
use arrow::{
1506-
array::{Float64Array, Int32Array},
1543+
array::{DictionaryArray, Float64Array, Int8Array, Int32Array},
15071544
buffer::NullBuffer,
1508-
datatypes::Int32Type,
1545+
datatypes::{Float64Type, Int8Type, Int32Type},
15091546
};
15101547
#[cfg(feature = "sql")]
15111548
use sqlparser::ast::Ident;
@@ -1534,6 +1571,49 @@ mod tests {
15341571
}
15351572
}
15361573

1574+
#[test]
1575+
fn normalize_float_zero_in_dictionary_values() -> Result<()> {
1576+
let nan = f64::from_bits(0x7ff8_0000_0000_0001);
1577+
let keys = Int8Array::from(vec![Some(0), Some(1), None, Some(2)]);
1578+
let array: ArrayRef = Arc::new(DictionaryArray::try_new(
1579+
keys.clone(),
1580+
Arc::new(Float64Array::from(vec![-0.0, nan, 1.0])),
1581+
)?);
1582+
1583+
let normalized = normalize_float_zero(&array);
1584+
assert!(!Arc::ptr_eq(&normalized, &array));
1585+
let dictionary = normalized.as_dictionary::<Int8Type>();
1586+
assert_eq!(dictionary.keys(), &keys);
1587+
let values = dictionary.values().as_primitive::<Float64Type>();
1588+
assert_eq!(values.value(0).to_bits(), 0.0_f64.to_bits());
1589+
assert_eq!(values.value(1).to_bits(), nan.to_bits());
1590+
assert_eq!(values.value(2), 1.0);
1591+
1592+
let without_negative_zero: ArrayRef = Arc::new(DictionaryArray::try_new(
1593+
Int8Array::from(vec![0, 1]),
1594+
Arc::new(Float64Array::from(vec![0.0, nan])),
1595+
)?);
1596+
assert!(Arc::ptr_eq(
1597+
&normalize_float_zero(&without_negative_zero),
1598+
&without_negative_zero
1599+
));
1600+
1601+
let scalar = ScalarValue::Dictionary(
1602+
Box::new(DataType::Int8),
1603+
Box::new(ScalarValue::Float64(Some(-0.0))),
1604+
);
1605+
let ScalarValue::Dictionary(_, value) = normalize_float_zero_scalar(scalar)
1606+
else {
1607+
unreachable!()
1608+
};
1609+
let ScalarValue::Float64(Some(value)) = *value else {
1610+
unreachable!()
1611+
};
1612+
assert_eq!(value.to_bits(), 0.0_f64.to_bits());
1613+
1614+
Ok(())
1615+
}
1616+
15371617
#[test]
15381618
fn test_bisect_linear_left_and_right() -> Result<()> {
15391619
let arrays: Vec<ArrayRef> = vec![

‎datafusion/physical-expr/src/expressions/in_list.rs‎

Lines changed: 155 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ use arrow::compute::kernels::boolean::{not, or_kleene};
3131
use arrow::compute::kernels::cmp::eq as arrow_eq;
3232
use arrow::datatypes::*;
3333

34+
use datafusion_common::utils::{normalize_float_zero, normalize_float_zero_scalar};
3435
use datafusion_common::{
3536
DFSchema, Result, ScalarValue, assert_or_internal_err, exec_err,
3637
};
@@ -82,6 +83,27 @@ fn supports_arrow_eq(dt: &DataType) -> bool {
8283
}
8384
}
8485

86+
fn normalize_in_list_float_zero_value(value: ColumnarValue) -> ColumnarValue {
87+
match value {
88+
ColumnarValue::Array(array)
89+
if is_float_or_dictionary_float(array.data_type()) =>
90+
{
91+
ColumnarValue::Array(normalize_float_zero(&array))
92+
}
93+
ColumnarValue::Scalar(scalar) => {
94+
ColumnarValue::Scalar(normalize_float_zero_scalar(scalar))
95+
}
96+
value => value,
97+
}
98+
}
99+
100+
fn is_float_or_dictionary_float(mut data_type: &DataType) -> bool {
101+
while let DataType::Dictionary(_, value_type) = data_type {
102+
data_type = value_type;
103+
}
104+
data_type.is_floating()
105+
}
106+
85107
/// Evaluates the list of expressions into an array, flattening any dictionaries
86108
fn evaluate_list(
87109
list: &[Arc<dyn PhysicalExpr>],
@@ -370,12 +392,15 @@ impl PhysicalExpr for InListExpr {
370392
// Use Arrow's vectorized eq kernel for types it supports (primitive,
371393
// boolean, string, binary, dictionary), falling back to row-by-row
372394
// comparator for unsupported types (nested, RunEndEncoded, etc.).
373-
let value = value.into_array(num_rows)?;
395+
// Normalize the left side once for the whole list. Doing this
396+
// outside `compare_one` avoids rescanning it for every item.
397+
let value =
398+
normalize_in_list_float_zero_value(value).into_array(num_rows)?;
374399
let lhs_supports_arrow_eq = supports_arrow_eq(value.data_type());
375400

376401
// Helper: compare value against a single list expression
377402
let compare_one = |expr: &Arc<dyn PhysicalExpr>| -> Result<BooleanArray> {
378-
match expr.evaluate(batch)? {
403+
match normalize_in_list_float_zero_value(expr.evaluate(batch)?) {
379404
ColumnarValue::Array(array) => {
380405
if lhs_supports_arrow_eq
381406
&& supports_arrow_eq(array.data_type())
@@ -3364,6 +3389,111 @@ mod tests {
33643389
Ok(())
33653390
}
33663391

3392+
#[test]
3393+
fn test_in_list_with_columns_float_signed_zero() -> Result<()> {
3394+
let schema = Schema::new(vec![
3395+
Field::new("a", DataType::Float64, false),
3396+
Field::new("b", DataType::Float64, false),
3397+
]);
3398+
let batch = RecordBatch::try_new(
3399+
Arc::new(schema.clone()),
3400+
vec![
3401+
Arc::new(Float64Array::from(vec![0.0, -0.0, 1.0])),
3402+
Arc::new(Float64Array::from(vec![-0.0, 0.0, 2.0])),
3403+
],
3404+
)?;
3405+
3406+
let expr = make_in_list_with_columns(
3407+
col("a", &schema)?,
3408+
vec![col("b", &schema)?],
3409+
false,
3410+
);
3411+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3412+
assert_eq!(
3413+
as_boolean_array(&result),
3414+
&BooleanArray::from(vec![true, true, false])
3415+
);
3416+
Ok(())
3417+
}
3418+
3419+
#[test]
3420+
fn test_in_list_with_columns_float_scalar_signed_zero() -> Result<()> {
3421+
let schema = Schema::new(vec![Field::new("a", DataType::Float32, false)]);
3422+
let batch = RecordBatch::try_new(
3423+
Arc::new(schema.clone()),
3424+
vec![Arc::new(Float32Array::from(vec![0.0, -0.0, 1.0]))],
3425+
)?;
3426+
let list = vec![lit(ScalarValue::Float32(Some(-0.0)))];
3427+
3428+
for (negated, expected) in [
3429+
(false, BooleanArray::from(vec![true, true, false])),
3430+
(true, BooleanArray::from(vec![false, false, true])),
3431+
] {
3432+
let expr =
3433+
make_in_list_with_columns(col("a", &schema)?, list.clone(), negated);
3434+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3435+
assert_eq!(as_boolean_array(&result), &expected);
3436+
}
3437+
3438+
// A scalar left-hand side is normalized before it is broadcast.
3439+
let expr = make_in_list_with_columns(
3440+
lit(ScalarValue::Float32(Some(-0.0))),
3441+
vec![col("a", &schema)?],
3442+
false,
3443+
);
3444+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3445+
assert_eq!(
3446+
as_boolean_array(&result),
3447+
&BooleanArray::from(vec![true, true, false])
3448+
);
3449+
3450+
Ok(())
3451+
}
3452+
3453+
#[test]
3454+
fn test_in_list_with_columns_dictionary_float_signed_zero() -> Result<()> {
3455+
let left: ArrayRef = Arc::new(DictionaryArray::try_new(
3456+
Int8Array::from(vec![0, 1, 2]),
3457+
Arc::new(Float64Array::from(vec![0.0, -0.0, 1.0])),
3458+
)?);
3459+
let right: ArrayRef = Arc::new(DictionaryArray::try_new(
3460+
Int8Array::from(vec![0, 1, 2]),
3461+
Arc::new(Float64Array::from(vec![-0.0, 0.0, 2.0])),
3462+
)?);
3463+
let data_type = left.data_type().clone();
3464+
let schema = Schema::new(vec![
3465+
Field::new("a", data_type.clone(), false),
3466+
Field::new("b", data_type, false),
3467+
]);
3468+
let batch = RecordBatch::try_new(Arc::new(schema.clone()), vec![left, right])?;
3469+
3470+
for (negated, expected) in [
3471+
(false, BooleanArray::from(vec![true, true, false])),
3472+
(true, BooleanArray::from(vec![false, false, true])),
3473+
] {
3474+
let expr = make_in_list_with_columns(
3475+
col("a", &schema)?,
3476+
vec![col("b", &schema)?],
3477+
negated,
3478+
);
3479+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3480+
assert_eq!(as_boolean_array(&result), &expected);
3481+
}
3482+
3483+
let scalar = lit(ScalarValue::Dictionary(
3484+
Box::new(DataType::Int8),
3485+
Box::new(ScalarValue::Float64(Some(-0.0))),
3486+
));
3487+
let expr = make_in_list_with_columns(col("a", &schema)?, vec![scalar], false);
3488+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3489+
assert_eq!(
3490+
as_boolean_array(&result),
3491+
&BooleanArray::from(vec![true, true, false])
3492+
);
3493+
3494+
Ok(())
3495+
}
3496+
33673497
/// Tests that short-circuit evaluation produces correct results.
33683498
/// When all rows match after the first list item, remaining items
33693499
/// should be skipped without affecting correctness.
@@ -3889,6 +4019,29 @@ mod tests {
38894019
Ok(())
38904020
}
38914021

4022+
#[test]
4023+
fn test_try_new_from_array_dict_haystack_float64_signed_zero() -> Result<()> {
4024+
// One value beyond the branchless limit selects the hash-set strategy.
4025+
let list_len =
4026+
<Float64Type as branchless_filter::BranchlessFilterType>::MAX_LIST_LEN + 1;
4027+
let mut list_values = vec![Some(-0.0)];
4028+
list_values.extend((1..list_len).map(|value| Some(value as f64)));
4029+
let haystack = make_f64_dict_array(list_values);
4030+
let needles: ArrayRef = Arc::new(Float64Array::from(vec![0.0, -0.0, -1.0]));
4031+
let expected = BooleanArray::from(vec![true, true, false]);
4032+
4033+
assert_eq!(
4034+
eval_in_list_from_array(Arc::clone(&needles), Arc::clone(&haystack))?,
4035+
expected
4036+
);
4037+
assert_eq!(
4038+
eval_in_list_from_array(wrap_in_dict(needles), haystack)?,
4039+
expected
4040+
);
4041+
4042+
Ok(())
4043+
}
4044+
38924045
#[test]
38934046
fn test_try_new_from_array_type_mismatch_rejects() -> Result<()> {
38944047
let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]);

0 commit comments

Comments
 (0)