diff --git a/datafusion/functions/benches/math_expressions/main.rs b/datafusion/functions/benches/math_expressions/main.rs index 6599ea7fca5e0..b87092b73e87c 100644 --- a/datafusion/functions/benches/math_expressions/main.rs +++ b/datafusion/functions/benches/math_expressions/main.rs @@ -40,6 +40,7 @@ mod round_dense; mod signum; mod trunc; mod trunc_precision; +mod unary_math; criterion_main!( atan2::benches, @@ -58,4 +59,5 @@ criterion_main!( signum::benches, trunc::benches, trunc_precision::benches, + unary_math::benches, ); diff --git a/datafusion/functions/benches/math_expressions/unary_math.rs b/datafusion/functions/benches/math_expressions/unary_math.rs new file mode 100644 index 0000000000000..51cdb69610c77 --- /dev/null +++ b/datafusion/functions/benches/math_expressions/unary_math.rs @@ -0,0 +1,83 @@ +// 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. + +use std::hint::black_box; +use std::sync::Arc; + +use arrow::array::ArrayRef; +use arrow::datatypes::{Field, Float32Type, Float64Type}; +use arrow::util::bench_util::create_primitive_array; +use criterion::{Criterion, criterion_group}; +use datafusion_common::config::ConfigOptions; +use datafusion_expr::{ColumnarValue, ScalarFunctionArgs}; +use datafusion_functions::math::{degrees, sqrt}; + +const BATCH_SIZE: usize = 8192; + +fn criterion_benchmark(c: &mut Criterion) { + let config_options = Arc::new(ConfigOptions::default()); + + // Both functions are cheap enough that overhead in the code they share + // shows up: `sqrt` returns an error for invalid inputs, `degrees` doesn't. + for function in [sqrt(), degrees()] { + let mut group = c.benchmark_group(function.name().to_string()); + + for (nulls, null_density) in [("", 0.0), (" with nulls", 0.1)] { + let arrays: [(&str, ArrayRef); 2] = [ + ( + "f64", + Arc::new(create_primitive_array::( + BATCH_SIZE, + null_density, + )), + ), + ( + "f32", + Arc::new(create_primitive_array::( + BATCH_SIZE, + null_density, + )), + ), + ]; + + for (type_name, array) in arrays { + let field = Arc::new(Field::new("a", array.data_type().clone(), true)); + let args = vec![ColumnarValue::Array(array)]; + + group.bench_function(format!("{type_name}{nulls}"), |b| { + b.iter(|| { + black_box( + function + .invoke_with_args(ScalarFunctionArgs { + args: args.clone(), + arg_fields: vec![Arc::clone(&field)], + number_rows: BATCH_SIZE, + return_field: Arc::clone(&field), + config_options: Arc::clone(&config_options), + }) + .unwrap(), + ) + }) + }); + } + } + + group.finish(); + } +} + +criterion_group!(benches, criterion_benchmark); diff --git a/datafusion/functions/src/macros.rs b/datafusion/functions/src/macros.rs index 8a6607c46b45e..42b33d84ca97f 100644 --- a/datafusion/functions/src/macros.rs +++ b/datafusion/functions/src/macros.rs @@ -209,6 +209,8 @@ macro_rules! downcast_arg { /// $OUTPUT_ORDERING: the output ordering calculation method of the function /// $STRICT: whether the function returns NULL when any argument is NULL /// $GET_DOC: the function to get the documentation of the UDF +/// $INPUT_ERROR (optional): a function that returns an error message for an +/// argument value outside the function's domain, or `None` for a valid value macro_rules! make_math_unary_udf { ($UDF:ident, $NAME:ident, $UNARY_FUNC:ident, $OUTPUT_ORDERING:expr, $EVALUATE_BOUNDS:expr, $STRICT:expr, $GET_DOC:expr) => { make_math_unary_udf!( @@ -219,10 +221,10 @@ macro_rules! make_math_unary_udf { $EVALUATE_BOUNDS, $STRICT, $GET_DOC, - None:: Result<()>> + None:: Option<&'static str>> ); }; - ($UDF:ident, $NAME:ident, $UNARY_FUNC:ident, $OUTPUT_ORDERING:expr, $EVALUATE_BOUNDS:expr, $STRICT:expr, $GET_DOC:expr, $VALIDATOR:expr) => { + ($UDF:ident, $NAME:ident, $UNARY_FUNC:ident, $OUTPUT_ORDERING:expr, $EVALUATE_BOUNDS:expr, $STRICT:expr, $GET_DOC:expr, $INPUT_ERROR:expr) => { $crate::make_udf_function!($NAME::$UDF, $NAME); mod $NAME { @@ -231,7 +233,6 @@ macro_rules! make_math_unary_udf { use arrow::array::{ArrayRef, AsArray}; use arrow::datatypes::{DataType, Float32Type, Float64Type}; - use arrow::error::ArrowError; use datafusion_common::{Result, exec_err}; use datafusion_expr::interval_arithmetic::Interval; use datafusion_expr::sort_properties::{ExprProperties, SortProperties}; @@ -297,36 +298,32 @@ macro_rules! make_math_unary_udf { let args = ColumnarValue::values_to_arrays(&args.args)?; let arr: ArrayRef = match args[0].data_type() { DataType::Float64 => { - let values = args[0] - .as_primitive::() - .try_unary::<_, Float64Type, _>( - |x: f64| -> std::result::Result { - if let Some(validate) = $VALIDATOR { - validate(x).map_err(|error| { - ArrowError::ComputeError(error.to_string()) - })?; - } - - Ok(f64::$UNARY_FUNC(x)) - }, - )?; - Arc::new(values) as ArrayRef + let array = args[0].as_primitive::(); + let result = match $INPUT_ERROR { + Some(input_error) => { + $crate::math::common::unary_with_input_check( + array, + f64::$UNARY_FUNC, + input_error, + )? + } + None => array.unary::<_, Float64Type>(f64::$UNARY_FUNC), + }; + Arc::new(result) as ArrayRef } DataType::Float32 => { - let values = args[0] - .as_primitive::() - .try_unary::<_, Float32Type, _>( - |x: f32| -> std::result::Result { - if let Some(validate) = $VALIDATOR { - validate(x as f64).map_err(|error| { - ArrowError::ComputeError(error.to_string()) - })?; - } - - Ok(f32::$UNARY_FUNC(x)) - }, - )?; - Arc::new(values) as ArrayRef + let array = args[0].as_primitive::(); + let result = match $INPUT_ERROR { + Some(input_error) => { + $crate::math::common::unary_with_input_check( + array, + f32::$UNARY_FUNC, + |x: f32| input_error(x as f64), + )? + } + None => array.unary::<_, Float32Type>(f32::$UNARY_FUNC), + }; + Arc::new(result) as ArrayRef } other => { return exec_err!( diff --git a/datafusion/functions/src/math/common.rs b/datafusion/functions/src/math/common.rs index e9f7c08235a2d..105ca576d9859 100644 --- a/datafusion/functions/src/math/common.rs +++ b/datafusion/functions/src/math/common.rs @@ -15,8 +15,10 @@ // specific language governing permissions and limitations // under the License. -use arrow::array::ArrowNativeTypeOp; +use arrow::array::{Array, ArrowNativeTypeOp, ArrowPrimitiveType, PrimitiveArray}; +use arrow::buffer::BooleanBuffer; use arrow::error::ArrowError; +use datafusion_common::{Result, exec_err}; use num_traits::{CheckedMul, CheckedNeg, Signed}; use std::fmt::Display; use std::mem::swap; @@ -150,10 +152,60 @@ pub(crate) fn lcm_signed_int(x: i64, y: i64) -> Result { }) } +/// An alternative to `try_unary` that lets the compiler vectorize both the +/// input check and `op`, for functions that return an error for some argument +/// values, such as `sqrt`, which returns an error for negative numbers. +/// +/// `try_unary` can return early on any value, which keeps the compiler from +/// vectorizing its loop. Instead, this applies `op` to every value in `array`, +/// like `unary`, and calls `input_error` on every value, including those in +/// null slots, in the same loop. Only if some value fails is the array searched +/// again, ignoring null slots, for an error to report. `input_error` should +/// therefore be a cheap check, such as a comparison. +/// +/// For cheap functions like `sqrt`, this is several times faster than +/// `try_unary`. But because it also does work for null slots, `try_unary` can +/// be faster on arrays with many nulls, especially when `op` is expensive and +/// can't be vectorized anyway. +pub(crate) fn unary_with_input_check( + array: &PrimitiveArray, + op: impl Fn(T::Native) -> T::Native, + input_error: impl Fn(T::Native) -> Option<&'static str>, +) -> Result> { + let mut any_invalid = false; + let values: Vec = array + .values() + .iter() + .map(|&x| { + any_invalid |= input_error(x).is_some(); + op(x) + }) + .collect(); + + // The check above also ran on null slots, which can hold any value, so the + // failure may be spurious. Re-check every value into a bitmap and mask out + // the null slots, which is faster than checking only the non-null values one + // at a time. + if any_invalid { + let input = array.values(); + let mut failed = + BooleanBuffer::collect_bool(input.len(), |i| input_error(input[i]).is_some()); + if let Some(nulls) = array.nulls() { + failed = &failed & nulls.inner(); + } + if let Some(message) = failed.set_indices().find_map(|i| input_error(input[i])) { + return exec_err!("{message}"); + } + } + + Ok(PrimitiveArray::new(values.into(), array.nulls().cloned())) +} + #[cfg(test)] mod tests { use super::*; - use arrow_buffer::i256; + use arrow::array::Float64Array; + use arrow_buffer::{NullBuffer, i256}; const GCD_COMMON_TEST_CASES: [(i64, i64, i64); 18] = [ // Basic cases @@ -317,4 +369,21 @@ mod tests { ); } } + + #[test] + fn test_unary_with_input_check() { + let input_error = |x: f64| (x < 0.0).then_some("negative input"); + + // -1.0 is in a null slot, so it is not an error. + let array = Float64Array::new( + vec![4.0, -1.0, 9.0].into(), + Some(NullBuffer::from(vec![true, false, true])), + ); + let result = unary_with_input_check(&array, f64::sqrt, input_error).unwrap(); + assert_eq!(result, Float64Array::from(vec![Some(2.0), None, Some(3.0)])); + + let array = Float64Array::from(vec![Some(4.0), None, Some(-1.0)]); + let error = unary_with_input_check(&array, f64::sqrt, input_error).unwrap_err(); + assert_eq!(error.strip_backtrace(), "Execution error: negative input"); + } } diff --git a/datafusion/functions/src/math/mod.rs b/datafusion/functions/src/math/mod.rs index 4b79866895d84..78031c62794d1 100644 --- a/datafusion/functions/src/math/mod.rs +++ b/datafusion/functions/src/math/mod.rs @@ -18,7 +18,6 @@ //! "math" DataFusion functions use crate::math::monotonicity::*; -use datafusion_common::{Result, exec_err}; use datafusion_expr::ScalarUDF; use std::sync::Arc; @@ -44,12 +43,10 @@ pub mod round; pub mod signum; pub mod trunc; -fn validate_sqrt_input(value: f64) -> Result<()> { - if value < 0.0 { - exec_err!("cannot take square root of a negative number") - } else { - Ok(()) - } +/// `f64::sqrt` returns NaN for negative numbers; like PostgreSQL, `sqrt` +/// returns an error instead. +fn sqrt_input_error(value: f64) -> Option<&'static str> { + (value < 0.0).then_some("cannot take square root of a negative number") } // Create UDFs @@ -238,7 +235,7 @@ make_math_unary_udf!( super::bounds::sqrt_bounds, true, super::get_sqrt_doc, - Some(super::validate_sqrt_input) + Some(super::sqrt_input_error) ); make_math_unary_udf!( TanFunc, @@ -264,7 +261,7 @@ make_udf_function!(trunc::TruncFunc, trunc); mod strict_tests { use super::*; use arrow::datatypes::Field; - use datafusion_common::ScalarValue; + use datafusion_common::{Result, ScalarValue}; use datafusion_expr::{ ColumnarValue, ReturnFieldArgs, ScalarFunctionArgs, ScalarUDF, }; diff --git a/datafusion/sqllogictest/test_files/scalar.slt b/datafusion/sqllogictest/test_files/scalar.slt index 35705a12a9622..ebb7e78aa4e76 100644 --- a/datafusion/sqllogictest/test_files/scalar.slt +++ b/datafusion/sqllogictest/test_files/scalar.slt @@ -1262,6 +1262,14 @@ select sqrt(-1); query error cannot take square root of a negative number select sqrt((-1.0)::float8); +# sqrt scalar negative float4 +query error cannot take square root of a negative number +select sqrt(arrow_cast(-1, 'Float32')); + +# sqrt negative float4 column +query error cannot take square root of a negative number +select sqrt(arrow_cast(a, 'Float32')) from signed_integers; + ## tan # tan scalar function