From e50d5ac8c82de1af76761b5dd8b1f40ae903aeb0 Mon Sep 17 00:00:00 2001 From: pushnanashi2 Date: Mon, 14 Sep 2026 00:58:22 +0900 Subject: [PATCH] fix: preserve floating-point percentile interpolation precision --- .../src/percentile_cont.rs | 97 +++++++++++++++---- .../sqllogictest/test_files/aggregate.slt | 17 +++- 2 files changed, 93 insertions(+), 21 deletions(-) diff --git a/datafusion/functions-aggregate/src/percentile_cont.rs b/datafusion/functions-aggregate/src/percentile_cont.rs index 3dbaf40efa36d..ce20665fe5f33 100644 --- a/datafusion/functions-aggregate/src/percentile_cont.rs +++ b/datafusion/functions-aggregate/src/percentile_cont.rs @@ -67,7 +67,7 @@ use datafusion_functions_aggregate_common::aggregate::groups_accumulator::nulls: use datafusion_functions_aggregate_common::utils::Hashable; use datafusion_macros::user_doc; -/// Precision multiplier for linear interpolation calculations. +/// Precision multiplier for decimal linear interpolation calculations. /// /// This value of 1,000,000 was chosen to balance precision with overflow safety: /// - Provides 6 decimal places of precision for the fractional component @@ -78,9 +78,7 @@ use datafusion_macros::user_doc; /// is computed as: `lower + ((upper - lower) * (fraction * PRECISION)) / PRECISION` /// to avoid floating-point operations on integer types while maintaining precision. /// -/// The interpolation arithmetic for floats is performed in f64 and then cast back to the -/// native type to avoid overflowing Float16 intermediates. -const INTERPOLATION_PRECISION: usize = 1_000_000; +const DECIMAL_INTERPOLATION_PRECISION: i64 = 1_000_000; create_func!(PercentileCont, percentile_cont_udaf); @@ -934,18 +932,19 @@ where upper: T::Native, fraction: f64, ) -> Result { - // Linear interpolation. - // We compute a quantized interpolation weight using `FLOAT_INTERPOLATION_PRECISION` because: - // 1. Both values come from the input data, so (upper - lower) is bounded by the value range - // 2. fraction is between 0 and 1; quantizing it provides stable, predictable results - // 3. The result is guaranteed to be between lower_value and upper_value (modulo cast rounding) - // 4. Arithmetic is performed in f64 and cast back to avoid overflowing Float16 intermediates - let scaled = (fraction * (INTERPOLATION_PRECISION as f64)) as usize; - let weight = scaled as f64 / (INTERPOLATION_PRECISION as f64); - let lower_f: f64 = lower.as_(); let upper_f: f64 = upper.as_(); - let interpolated_f = lower_f + (upper_f - lower_f) * weight; + // The usual `lower + (upper - lower) * fraction` avoids overflowing + // same-sign finite inputs. For opposite signs, use separately weighted + // terms because the subtraction itself can overflow f64. + let interpolated_f = if lower_f.is_finite() + && upper_f.is_finite() + && lower_f.is_sign_negative() != upper_f.is_sign_negative() + { + lower_f * (1.0 - fraction) + upper_f * fraction + } else { + lower_f + (upper_f - lower_f) * fraction + }; Ok(interpolated_f.as_()) } } @@ -962,7 +961,7 @@ where T: DecimalType, T::Native: DecimalCast, { - let den = INTERPOLATION_PRECISION as i64; + let den = DECIMAL_INTERPOLATION_PRECISION; debug_assert!(num >= 0); debug_assert!(num <= den); @@ -1002,8 +1001,8 @@ where debug_assert!((0.0..=1.0).contains(&fraction)); debug_assert!(lower <= upper); - let num = (fraction * INTERPOLATION_PRECISION as f64) as i64; - let den = INTERPOLATION_PRECISION as i64; + let num = (fraction * DECIMAL_INTERPOLATION_PRECISION as f64) as i64; + let den = DECIMAL_INTERPOLATION_PRECISION; // Happy path: `upper - lower` does not overflow // (could be a case for Decimal128 with max precision) @@ -1103,7 +1102,9 @@ fn calculate_percentile>( mod tests { use super::*; use arrow::array::Float64Array; - use arrow::datatypes::{Decimal64Type, Decimal128Type, Float16Type, Float64Type}; + use arrow::datatypes::{ + Decimal64Type, Decimal128Type, Float16Type, Float32Type, Float64Type, + }; use half::f16; #[test] @@ -1181,6 +1182,66 @@ mod tests { ); } + #[test] + fn float_interpolation_preserves_sub_millionth_weights() { + // Regression test for https://github.com/apache/datafusion/issues/25266 + let percentile = 0.00000025; + + let mut f16_values = + vec![f16::from_f32(0.0), f16::from_f32(1.0), f16::from_f32(2.0)]; + let f16_result = calculate_percentile::( + &mut f16_values, + percentile, + ) + .expect("interpolation succeeds") + .expect("non-empty result"); + assert_eq!(f16_result, f16::from_f64(0.0000005)); + + let mut f32_values = vec![0.0_f32, 1.0, 2.0]; + let f32_result = calculate_percentile::( + &mut f32_values, + percentile, + ) + .expect("interpolation succeeds") + .expect("non-empty result"); + assert_eq!(f32_result, 0.0000005_f32); + + let mut f64_values = vec![0.0_f64, 1.0, 2.0]; + let f64_result = calculate_percentile::( + &mut f64_values, + percentile, + ) + .expect("interpolation succeeds") + .expect("non-empty result"); + assert_eq!(f64_result, 0.0000005_f64); + } + + #[test] + fn f64_interpolation_does_not_overflow_for_opposite_extremes() { + let midpoint = + >::interpolate( + -f64::MAX, + f64::MAX, + 0.5, + ) + .expect("interpolation succeeds"); + assert_eq!(midpoint, 0.0); + + let lower_quartile = + >::interpolate( + -f64::MAX, + f64::MAX, + 0.25, + ) + .expect("interpolation succeeds"); + let expected = -f64::MAX / 2.0; + assert!(lower_quartile.is_finite()); + assert!( + (lower_quartile - expected).abs() <= f64::EPSILON * f64::MAX, + "expected {expected}, got {lower_quartile}" + ); + } + #[test] fn percentile_cont_decimal64() { // Test values: [100.00, 200.00, 300.00, 400.00, 500.00] diff --git a/datafusion/sqllogictest/test_files/aggregate.slt b/datafusion/sqllogictest/test_files/aggregate.slt index a4cd3b0ccc1f0..61f55b46c0389 100644 --- a/datafusion/sqllogictest/test_files/aggregate.slt +++ b/datafusion/sqllogictest/test_files/aggregate.slt @@ -1183,6 +1183,17 @@ DROP TABLE distinct_pct_nulls; # percentile_cont for shorter floats +# Regression test for https://github.com/apache/datafusion/issues/25266 +# Float interpolation must not quantize weights to six decimal places. +query RRR +SELECT + percentile_cont(arrow_cast(x, 'Float16'), arrow_cast(0.00000025, 'Float64')), + percentile_cont(arrow_cast(x, 'Float32'), arrow_cast(0.00000025, 'Float64')), + percentile_cont(arrow_cast(x, 'Float64'), arrow_cast(0.00000025, 'Float64')) +FROM (VALUES (0.0), (1.0), (2.0)) AS t(x); +---- +0.000000476837 0.0000005 0.0000005 + query RT select percentile_cont(arrow_cast(col_f32, 'Float16'), 0.5), arrow_typeof(percentile_cont(arrow_cast(col_f32, 'Float16'), 0.5)) from median_table; ---- @@ -4361,7 +4372,7 @@ SELECT percentile_cont(0.95) WITHIN GROUP (ORDER BY c3 DESC) FROM aggregate_test query R SELECT percentile_cont(0.05) WITHIN GROUP (ORDER BY c3 DESC) FROM aggregate_test_100 ---- -118.099998 +118.1 # Test with GROUP BY query TR @@ -4379,7 +4390,7 @@ SELECT c1, percentile_cont(0.95) WITHIN GROUP (ORDER BY c3) FROM aggregate_test_ a 65 b 68 c 118 -d 123.299998 +d 123.3 e 112 # Test with NULLs @@ -4601,7 +4612,7 @@ SELECT percentile_cont(0.6) WITHIN GROUP (ORDER BY v DESC) FROM (VALUES (1), (2) query R SELECT percentile_cont(0.3) WITHIN GROUP (ORDER BY v) FROM (VALUES (10), (20), (30), (40), (50)) as t (v) ---- -21.99999 +22 # Show 0.7 descending (should be same as 0.3 ascending) query R