0lai0 commented on code in PR #5367:
URL: https://github.com/apache/datafusion-comet/pull/5367#discussion_r3882005061
##########
native/spark-expr/src/math_funcs/pow.rs:
##########
@@ -42,86 +45,55 @@ fn spark_powf(base: f64, exp: f64) -> f64 {
/// Unlike DataFusion's `power`, `pow(0, -1)` returns `Infinity` rather than
erroring. Only null
/// inputs produce null; otherwise every result is the `spark_powf` value.
pub fn spark_pow(args: &[ColumnarValue]) -> Result<ColumnarValue,
DataFusionError> {
- if args.len() != 2 {
- return Err(DataFusionError::Internal(format!(
- "spark_pow requires 2 arguments, got {}",
- args.len()
- )));
- }
+ let [base, exp] = take_function_args("spark_pow", args)?;
+ apply(base, exp, spark_pow_kernel)
+}
- fn as_f64_array(
- value: &Arc<dyn arrow::array::Array>,
- ) -> Result<&Float64Array, DataFusionError> {
- value
- .as_any()
- .downcast_ref::<Float64Array>()
- .ok_or_else(|| {
- DataFusionError::Internal(format!(
- "spark_pow expected Float64, got {:?}",
- value.data_type()
- ))
- })
- }
+fn as_f64_array(array: &dyn Array) -> Result<&Float64Array, ArrowError> {
+ array
+ .as_any()
+ .downcast_ref::<Float64Array>()
+ .ok_or_else(|| {
+ ArrowError::ComputeError(format!(
+ "spark_pow expected Float64, got {:?}",
+ array.data_type()
+ ))
+ })
+}
- fn as_f64_scalar(scalar: &ScalarValue) -> Result<Option<f64>,
DataFusionError> {
- match scalar {
- ScalarValue::Float64(v) => Ok(*v),
- _ => Err(DataFusionError::Internal(format!(
- "spark_pow expected Float64 scalar, got {scalar:?}",
- ))),
- }
- }
+/// Array/array uses [`binary`] over [`spark_powf`]. Scalar/array uses
[`unary`] so the
+/// scalar is not broadcast. A null scalar short-circuits to an all-null array.
+fn spark_pow_kernel(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef,
ArrowError> {
+ let (left, left_is_scalar) = lhs.get();
+ let (right, right_is_scalar) = rhs.get();
+ let left = as_f64_array(left)?;
+ let right = as_f64_array(right)?;
- match (&args[0], &args[1]) {
- (ColumnarValue::Array(base_arr), ColumnarValue::Array(exp_arr)) => {
- let bases = as_f64_array(base_arr)?;
- let exps = as_f64_array(exp_arr)?;
- let result: Float64Array = bases
- .iter()
- .zip(exps.iter())
- .map(|(b, e)| match (b, e) {
- (Some(base), Some(exp)) => Some(spark_powf(base, exp)),
- _ => None,
- })
- .collect();
- Ok(ColumnarValue::Array(Arc::new(result)))
- }
- (ColumnarValue::Scalar(base_scalar), ColumnarValue::Array(exp_arr)) =>
{
- let exps = as_f64_array(exp_arr)?;
- let result: Float64Array = match as_f64_scalar(base_scalar)? {
- Some(base) => exps
- .iter()
- .map(|e| e.map(|exp| spark_powf(base, exp)))
- .collect(),
- None => Float64Array::new_null(exp_arr.len()),
- };
- Ok(ColumnarValue::Array(Arc::new(result)))
- }
- (ColumnarValue::Array(base_arr), ColumnarValue::Scalar(exp_scalar)) =>
{
- let bases = as_f64_array(base_arr)?;
- let result: Float64Array = match as_f64_scalar(exp_scalar)? {
- Some(exp) => bases
- .iter()
- .map(|b| b.map(|base| spark_powf(base, exp)))
- .collect(),
- None => Float64Array::new_null(base_arr.len()),
- };
- Ok(ColumnarValue::Array(Arc::new(result)))
+ let result = match (left_is_scalar, right_is_scalar) {
+ (true, false) => {
+ if left.is_null(0) {
+ Float64Array::new_null(right.len())
+ } else {
+ unary(right, |exp| spark_powf(left.value(0), exp))
+ }
}
- (ColumnarValue::Scalar(base_scalar),
ColumnarValue::Scalar(exp_scalar)) => {
- let result = match (as_f64_scalar(base_scalar)?,
as_f64_scalar(exp_scalar)?) {
- (Some(base), Some(exp)) =>
ScalarValue::Float64(Some(spark_powf(base, exp))),
- _ => ScalarValue::Float64(None),
- };
- Ok(ColumnarValue::Scalar(result))
+ (false, true) => {
+ if right.is_null(0) {
+ Float64Array::new_null(left.len())
+ } else {
+ unary(left, |base| spark_powf(base, right.value(0)))
+ }
}
- }
+ _ => binary(left, right, spark_powf)?,
Review Comment:
Thanks for reviews !!
Fixed via an adaptive null-aware dispatcher that only pays the null-skip
overhead when it wins. Crossover is `null_count > 3 * len / 4` (75%).
Before vs. after @ 8192 rows (composed nulls, real payload in null slots,
like `pow(a + 2.5D, b)`):
| Shape | 90% nulls | 99% nulls |
|-------|-----------|-----------|
| Regression reported before | 47.9µs (1.84x slow) | 47.0µs (1.97x slow) |
| After null-aware dispatch, array/array | 5.56µs | 1.83µs |
| After, pipeline `pow(a+2.5D, 3)` incl. add | 6.35µs | 3.37µs |
Threshold behaviour is visible in the sweep: array/array composed 70% =
36.9µs (raw-buffer kernel), 80% = 9.1µs (null-skip kicks in). No-null /
sparse-null shapes are unchanged (27.4µs / 26.7µs), so the null-oblivious fast
path is preserved.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]