sunchao commented on code in PR #6170: URL: https://github.com/apache/datafusion-comet/pull/6170#discussion_r4126147168
########## native/spark-expr/src/array_funcs/array_extrema.rs: ########## @@ -0,0 +1,440 @@ +// 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::cmp::Ordering; +use std::sync::Arc; + +use arrow::array::{ + make_array, make_comparator, new_empty_array, Array, ArrayAccessor, ArrayRef, AsArray, + DynComparator, ListArray, MutableArrayData, PrimitiveArray, PrimitiveBuilder, StringArrayType, + StructArray, UInt32Array, +}; +use arrow::buffer::NullBuffer; +use arrow::compute::{take, SortOptions}; +use arrow::datatypes::{ArrowPrimitiveType, DataType, Float32Type, Float64Type}; +use datafusion::common::{exec_err, Result, ScalarValue}; +use datafusion::functions_nested::min_max::{array_max_udf, array_min_udf}; +use datafusion::logical_expr::{ + ColumnarValue, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature, +}; +use num::Float; + +#[derive(Debug, Clone, Copy, Hash, Eq, PartialEq)] +enum Utf8Collation { + Binary, + BinaryRtrim, + Lcase, + LcaseRtrim, +} + +impl Utf8Collation { + fn compare(self, mut left: &str, mut right: &str, unicode_version: u32) -> Ordering { + if matches!(self, Self::BinaryRtrim | Self::LcaseRtrim) { + // Spark RTRIM ignores trailing U+0020, not arbitrary Unicode whitespace. + left = left.trim_end_matches(' '); + right = right.trim_end_matches(' '); + } + if matches!(self, Self::Binary | Self::BinaryRtrim) { + return left.cmp(right); + } + fn lower(value: &str, unicode_version: u32) -> impl Iterator<Item = u32> + '_ { + value.chars().flat_map(move |c| { + let cp = c as u32; + let [first, second] = match cp { + // Spark treats final sigma as ordinary sigma. + 0x3c2 => [0x3c3, 0], + // Unicode 17 adds these mappings to the library's Unicode 16 data. + 0xa7ce | 0xa7d2 | 0xa7d4 if unicode_version == 17 => [cp + 1, 0], + 0x16ea0..=0x16eb8 if unicode_version == 17 => [cp + 0x1b, 0], + _ => unicode_case_mapping::to_lowercase(c), + }; + // The library uses zero for an unchanged code point or absent second value. + std::iter::once(if first == 0 { cp } else { first }) + .chain((second != 0).then_some(second)) + }) + } + // Compare ASCII prefixes lazily so a retained long winner is not rescanned. + for (offset, (l, r)) in left.bytes().zip(right.bytes()).enumerate() { + if !l.is_ascii() || !r.is_ascii() { + return lower(&left[offset..], unicode_version) + .cmp(lower(&right[offset..], unicode_version)); + } + let ordering = l.to_ascii_lowercase().cmp(&r.to_ascii_lowercase()); + if ordering != Ordering::Equal { + return ordering; + } + } + left.len().cmp(&right.len()) + } +} + +/// Spark's array_min/array_max retain the first non-null value on an ordering tie. +/// In particular, signed zeros compare equal and all NaNs compare equal and greater +/// than non-NaNs. Nested arrays and structs use the same ordering, with nulls first. +#[derive(Debug, Hash, Eq, PartialEq)] +pub struct SparkArrayExtrema { + is_min: bool, + datafusion_udf: Arc<ScalarUDF>, + string_collations: Vec<Utf8Collation>, + unicode_version: u32, +} + +impl SparkArrayExtrema { + pub fn new(is_min: bool) -> Self { + Self { + is_min, + // Capture the original implementation, not a registry lookup: these UDFs + // replace the DataFusion names in Comet's function registry. + datafusion_udf: if is_min { + array_min_udf() + } else { + array_max_udf() + }, + string_collations: Vec::new(), + unicode_version: 0, + } + } + + /// Scala validates the Unicode version and lists string-leaf collations depth-first. + pub fn with_collations( + is_min: bool, + collations: &[String], + unicode_version: u32, + ) -> Result<Self> { + let string_collations = collations + .iter() + .map(|name| match name.as_str() { + "UTF8_BINARY" => Ok(Utf8Collation::Binary), + "UTF8_BINARY_RTRIM" => Ok(Utf8Collation::BinaryRtrim), + "UTF8_LCASE" => Ok(Utf8Collation::Lcase), + "UTF8_LCASE_RTRIM" => Ok(Utf8Collation::LcaseRtrim), + _ => exec_err!("Unsupported array extrema collation: {name}"), + }) + .collect::<Result<Vec<_>>>()?; + Ok(Self { + string_collations, + unicode_version, + ..Self::new(is_min) + }) + } +} + +impl ScalarUDFImpl for SparkArrayExtrema { + fn name(&self) -> &str { + if self.is_min { + "array_min" + } else { + "array_max" + } + } + + fn signature(&self) -> &Signature { + self.datafusion_udf.signature() + } + + fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> { + self.datafusion_udf.return_type(arg_types) + } + + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> { + let [input] = args.args.as_slice() else { + return exec_err!("{} takes exactly one argument", self.name()); + }; + let element_type = self.return_type(&[input.data_type()])?; + + // DataFusion's non-primitive path reconstructs an array from scalars, which + // cannot infer a type from an empty iterator. Keep the declared element type. + if matches!(input, ColumnarValue::Array(array) if array.is_empty()) { + return Ok(ColumnarValue::Array(new_empty_array(&element_type))); + } + if self.string_collations.is_empty() + && !matches!( + element_type, + DataType::Float32 | DataType::Float64 | DataType::List(_) | DataType::Struct(_) + ) + { + return self.datafusion_udf.invoke_with_args(args); Review Comment: Agreed. In 341174464, plain `Utf8`, `LargeUtf8`, and `Utf8View` now use the existing comparator/take path. I kept delegation unchanged for unrelated element types to keep this scoped. Added default-binary coverage for the three string layouts, null/empty/sliced inputs and scalar results. Selected StringView results also compact their backing buffers, with a regression guard against retaining a large losing value. The new release Criterion comparison covers eight shapes, both min/max, and 64 arrays of eight strings. All outputs matched DataFusion; DataFusion/Comet mean-time ratios were 1.38-4.61x across the 16 comparisons (32 timings). The full Spark benchmark's plain-binary controls improved 1.06-1.45x against #5403. The PR description now separates these kernel and end-to-end measurements. -- 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]
