sunchao commented on code in PR #4798:
URL: https://github.com/apache/datafusion-comet/pull/4798#discussion_r4125067113
##########
spark/src/main/scala/org/apache/comet/serde/aggregates.scala:
##########
@@ -1116,6 +1116,66 @@ object CometApproxCountDistinct extends
CometAggregateExpressionSerde[HyperLogLo
}
}
+object CometPivotFirst extends CometAggregateExpressionSerde[PivotFirst] {
+
+ // Delegate to Spark's own PivotFirst.supportsDataType so the two lists
cannot drift if a
+ // future Spark version adds a value type to the fast-path gate.
+ private def unsupportedValueTypeReason(dt: DataType): String =
+ s"Unsupported value data type: $dt"
+
+ private val emptyPivotValuesReason = "Pivot values list is empty"
+
+ override def getUnsupportedReasons(): Seq[String] = Seq(
+ "Value data type outside PivotFirst.supportsDataType " +
+ "(Boolean, Byte, Short, Int, Long, Float, Double, Decimal)",
+ emptyPivotValuesReason)
+
+ override def getSupportLevel(expr: PivotFirst): SupportLevel = {
+ if (!PivotFirst.supportsDataType(expr.valueDataType)) {
+ Unsupported(Some(unsupportedValueTypeReason(expr.valueDataType)))
+ } else if (expr.pivotColumnValues.isEmpty) {
+ Unsupported(Some(emptyPivotValuesReason))
+ } else {
+ Compatible()
Review Comment:
[P2] Gate binary pivot columns before returning `Compatible()`. For a
Parquet row `(g=1, k=X'61', v=10)`, `SELECT * FROM t PIVOT (sum(v) FOR k IN
(X'61', X'62'))` returns `(1,NULL,NULL)` in Spark, but the new native aggregate
populates the first slot with 10. Spark’s atomic-key `HashMap[Any, Int]`
compares these byte arrays by identity, while `ScalarValue::Binary` compares
their contents. Enabling this path changes query results by default. Preserve
Spark matching, or make binary pivot columns fall back.
Evidence: Ran the Parquet-backed reference query on Spark 4.1.3 and obtained
`[Row(g=1, a=None, b=None)]`, with `pivotfirst` in its physical plan. The
exact-source native probe using pivot literal
`ScalarValue::Binary(Some(vec![97]))` and an independently constructed binary
input containing `b"a"` produced `[Int32(10)]`. The checked Spark
implementations select `HashMap` for atomic pivot-column types, and the new
serde has no binary-key restriction.
##########
native/spark-expr/src/agg_funcs/pivot_first.rs:
##########
@@ -0,0 +1,497 @@
+// 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.
+
+//! Spark's `PivotFirst` aggregate. Used only by the second phase of the
optimized pivot plan
+//! generated by `PivotTransformer`. For each group, `PivotFirst` maintains an
array of
+//! `pivot_values.len()` slots; on each input row it evaluates the pivot
column, looks up its
+//! index in `pivot_values`, and writes the value column into that slot when a
match is found
+//! and the value is non-null. Rows with unmatched pivot values are ignored;
matched rows with
+//! a null value column leave the slot unchanged (matches Spark).
+//!
+//! State layout is one column per pivot slot, matching Spark's
`aggBufferAttributes` (which
+//! declares `indexSize` `AttributeReference`s, one per pivot value). This
keeps the shuffle
+//! schema between Partial and Final consistent with what Spark catalyst
declared; otherwise
+//! the shuffle exchange rejects the batch. `evaluate()` reassembles the slots
into a
+//! `ListArray` matching `PivotFirst.dataType = ArrayType(value_type)`.
+
+use arrow::array::{Array, ArrayRef};
+use arrow::datatypes::{DataType, Field, FieldRef};
+use datafusion::common::utils::SingleRowListArrayBuilder;
+use datafusion::common::{DataFusionError, Result as DFResult, ScalarValue};
+use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
+use datafusion::logical_expr::Volatility::Immutable;
+use datafusion::logical_expr::{Accumulator, AggregateUDFImpl, Signature};
+use datafusion::physical_expr::expressions::format_state_name;
+use std::collections::HashMap;
+use std::sync::Arc;
+
+/// UDAF implementation of Spark's `PivotFirst`.
+///
+/// `pivot_values` is a fixed, plan-time list of the pivot column values that
occupy each
+/// output slot; `pivot_index[v] = i` means an input row whose pivot column
equals `v` writes
+/// into slot `i`. Both the vector and the map are wrapped in `Arc` because
`accumulator()`
+/// fires once per group in a grouped aggregate and we want that path to bump
a refcount
+/// rather than deep-clone.
+#[derive(Debug)]
+pub struct SparkPivotFirst {
+ signature: Signature,
+ value_type: DataType,
+ // Kept for `PartialEq`/`Hash` (identity of the aggregate for plan
comparison) and for the
+ // deterministic slot ordering `state_fields` needs. `HashMap` alone would
give us the map
+ // but not a stable order or a `Hash` impl.
+ pivot_values: Arc<Vec<ScalarValue>>,
+ pivot_index: Arc<HashMap<ScalarValue, usize>>,
+}
+
+impl PartialEq for SparkPivotFirst {
+ fn eq(&self, other: &Self) -> bool {
+ self.value_type == other.value_type && self.pivot_values ==
other.pivot_values
+ }
+}
+
+impl Eq for SparkPivotFirst {}
+
+impl std::hash::Hash for SparkPivotFirst {
+ fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
+ self.value_type.hash(state);
+ self.pivot_values.hash(state);
+ }
+}
+
+impl SparkPivotFirst {
+ pub fn new(value_type: DataType, pivot_values: Vec<ScalarValue>) -> Self {
+ let mut pivot_index = HashMap::with_capacity(pivot_values.len());
+ // Spark's PivotFirst uses the FIRST occurrence's index
(HashMap/TreeMap semantics), so
+ // when duplicates are somehow present we mirror that by only
inserting the first one.
+ // `pivot_key` can fold two distinct pivot values (`0.0` and `-0.0`)
onto one key, and can
+ // drop one entirely (NaN), so the index is not necessarily the same
length as the slot
+ // vector - the slot count is always `pivot_values.len()`.
+ for (i, v) in pivot_values.iter().enumerate() {
+ if let Some(key) = pivot_key(v.clone()) {
+ pivot_index.entry(key).or_insert(i);
+ }
+ }
+ Self {
+ signature: Signature::user_defined(Immutable),
+ value_type,
+ pivot_values: Arc::new(pivot_values),
+ pivot_index: Arc::new(pivot_index),
+ }
+ }
+}
+
+/// Rewrite a pivot column value into the key Spark would match it on, or
`None` when Spark can
+/// never match it.
+///
+/// Spark's `PivotFirst` looks pivot values up in a Scala `HashMap[Any, Int]`,
so matching goes
+/// through `BoxesRunTime.equals` / `Statics.anyHash` on the boxed Catalyst
value rather than
+/// through `ScalarValue`'s own equality. The two disagree on floats in
opposite directions:
+///
+/// * `-0.0` and `0.0` are one key for Spark (`-0.0 == 0.0` numerically, and
`doubleHash` folds
+/// both onto the hash of `0L`), while `ScalarValue` keeps them apart.
+/// * `NaN` matches nothing for Spark, not even another `NaN`, because Scala's
`==` on `Double`
+/// is IEEE. `ScalarValue` treats `NaN` as equal to itself.
+///
+/// Nulls are left alone: a null pivot column value does match a null entry in
the pivot list,
+/// which is what Spark's `pivotIndex.getOrElse(null, -1)` does.
+fn pivot_key(v: ScalarValue) -> Option<ScalarValue> {
+ match v {
+ ScalarValue::Float32(Some(f)) => {
+ if f.is_nan() {
+ None
+ } else if f == 0.0 {
+ Some(ScalarValue::Float32(Some(0.0)))
+ } else {
+ Some(ScalarValue::Float32(Some(f)))
+ }
+ }
+ ScalarValue::Float64(Some(f)) => {
+ if f.is_nan() {
+ None
+ } else if f == 0.0 {
+ Some(ScalarValue::Float64(Some(0.0)))
+ } else {
+ Some(ScalarValue::Float64(Some(f)))
+ }
+ }
+ other => Some(other),
+ }
+}
+
+impl AggregateUDFImpl for SparkPivotFirst {
+ fn name(&self) -> &str {
+ "pivot_first"
+ }
+
+ fn signature(&self) -> &Signature {
+ &self.signature
+ }
+
+ fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
+ Ok(DataType::List(Arc::new(Field::new_list_field(
+ self.value_type.clone(),
+ true,
+ ))))
+ }
+
+ fn state_fields(&self, args: StateFieldsArgs) -> DFResult<Vec<FieldRef>> {
+ // One field per pivot slot, matching Spark's aggBufferAttributes so
the shuffle
+ // exchange sees the same schema catalyst declared.
`format_state_name` is the same
+ // helper other aggregates in this crate use (see `avg.rs`,
`stddev.rs`).
+ Ok((0..self.pivot_values.len())
+ .map(|i| {
+ Arc::new(Field::new(
+ format_state_name(args.name, &i.to_string()),
+ self.value_type.clone(),
+ true,
+ ))
+ })
+ .collect())
+ }
+
+ fn accumulator(&self, _acc_args: AccumulatorArgs) -> DFResult<Box<dyn
Accumulator>> {
+ Ok(Box::new(PivotFirstAccumulator::new(
+ self.value_type.clone(),
+ Arc::clone(&self.pivot_index),
+ self.pivot_values.len(),
+ )))
+ }
+}
+
+/// Per-group state: `slots[i]` holds the latest non-null value assigned to
pivot slot `i`, or
+/// `None` when nothing has written to that slot yet.
+#[derive(Debug)]
+struct PivotFirstAccumulator {
+ value_type: DataType,
+ pivot_index: Arc<HashMap<ScalarValue, usize>>,
+ slots: Vec<Option<ScalarValue>>,
+}
+
+impl PivotFirstAccumulator {
+ /// `num_slots` is the pivot list's length, which is what `state_fields`
declares. It can
+ /// exceed `pivot_index.len()` when pivot values collide under `pivot_key`
or are unmatchable
+ /// (NaN); those slots exist in the output and stay null.
+ fn new(
+ value_type: DataType,
+ pivot_index: Arc<HashMap<ScalarValue, usize>>,
+ num_slots: usize,
+ ) -> Self {
+ let slots = vec![None; num_slots];
+ Self {
+ value_type,
+ pivot_index,
+ slots,
+ }
+ }
+
+ /// Turn slot `i` into a `ScalarValue`, substituting a typed null when the
slot is empty.
+ fn slot_or_null(&self, i: usize) -> DFResult<ScalarValue> {
+ Ok(match &self.slots[i] {
+ Some(v) => v.clone(),
+ None => ScalarValue::try_from(&self.value_type)?,
+ })
+ }
+}
+
+impl Accumulator for PivotFirstAccumulator {
+ fn update_batch(&mut self, values: &[ArrayRef]) -> DFResult<()> {
+ if values.len() != 2 {
+ return Err(DataFusionError::Internal(format!(
+ "PivotFirst expects 2 inputs (pivot, value); got {}",
+ values.len()
+ )));
+ }
+ let pivot_arr = &values[0];
+ let value_arr = &values[1];
+ if pivot_arr.len() != value_arr.len() {
+ return Err(DataFusionError::Internal(
+ "PivotFirst pivot and value arrays have different
lengths".into(),
+ ));
+ }
+ for row in 0..pivot_arr.len() {
+ // Spark ignores the row entirely if either the pivot value is
unmatched (index<0)
+ // or the value is null. Matching Spark exactly here is important
because
+ // `PivotFirst.update` never writes for a null value, so a
mid-batch null does not
+ // clobber an earlier non-null.
+ let pivot_scalar = ScalarValue::try_from_array(pivot_arr, row)?;
+ let Some(key) = pivot_key(pivot_scalar) else {
+ // A NaN pivot value matches no slot in Spark, so the row is
ignored.
+ continue;
+ };
+ if let Some(&slot_idx) = self.pivot_index.get(&key) {
Review Comment:
[P1] Match array pivot keys using Spark value semantics. `ScalarValue::List`
equality includes Arrow field metadata, whereas Spark’s `TreeMap` uses
interpreted value ordering. Pivot literals have nullable list children, but
input arrays can have non-nullable children or different field names. Equal
values therefore miss this lookup and silently produce null totals. This
already breaks `FOR a IN (array(1, 1), array(2, 2))` in both Spark SQL CI
shards. Nested signed zeros also fail to match. Please implement
Spark-compatible complex-key comparison or make these pivot-column types fall
back until supported.
Evidence: Head-associated CI jobs
https://github.com/apache/datafusion-comet/actions/runs/34362701066/job/102532509588
and
https://github.com/apache/datafusion-comet/actions/runs/34362701066/job/102530993096
report pivot.sql query #25 expecting `(2012,35000,NULL)` and
`(2013,NULL,78000)`, but receiving nulls in every pivot column. An exact-source
Rust probe with literal `[1,1]` and identical input values returned
`Int32(NULL)` when child nullability or field name differed, and `Int32(10)`
when metadata matched. A `[0.0]` input also missed a `[-0.0]` pivot key, while
Spark 4.1.3 returned 10.
##########
native/spark-expr/src/agg_funcs/pivot_first.rs:
##########
@@ -0,0 +1,497 @@
+// 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.
+
+//! Spark's `PivotFirst` aggregate. Used only by the second phase of the
optimized pivot plan
+//! generated by `PivotTransformer`. For each group, `PivotFirst` maintains an
array of
+//! `pivot_values.len()` slots; on each input row it evaluates the pivot
column, looks up its
+//! index in `pivot_values`, and writes the value column into that slot when a
match is found
+//! and the value is non-null. Rows with unmatched pivot values are ignored;
matched rows with
+//! a null value column leave the slot unchanged (matches Spark).
+//!
+//! State layout is one column per pivot slot, matching Spark's
`aggBufferAttributes` (which
+//! declares `indexSize` `AttributeReference`s, one per pivot value). This
keeps the shuffle
+//! schema between Partial and Final consistent with what Spark catalyst
declared; otherwise
+//! the shuffle exchange rejects the batch. `evaluate()` reassembles the slots
into a
+//! `ListArray` matching `PivotFirst.dataType = ArrayType(value_type)`.
+
+use arrow::array::{Array, ArrayRef};
+use arrow::datatypes::{DataType, Field, FieldRef};
+use datafusion::common::utils::SingleRowListArrayBuilder;
+use datafusion::common::{DataFusionError, Result as DFResult, ScalarValue};
+use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
+use datafusion::logical_expr::Volatility::Immutable;
+use datafusion::logical_expr::{Accumulator, AggregateUDFImpl, Signature};
+use datafusion::physical_expr::expressions::format_state_name;
+use std::collections::HashMap;
+use std::sync::Arc;
+
+/// UDAF implementation of Spark's `PivotFirst`.
+///
+/// `pivot_values` is a fixed, plan-time list of the pivot column values that
occupy each
+/// output slot; `pivot_index[v] = i` means an input row whose pivot column
equals `v` writes
+/// into slot `i`. Both the vector and the map are wrapped in `Arc` because
`accumulator()`
+/// fires once per group in a grouped aggregate and we want that path to bump
a refcount
+/// rather than deep-clone.
+#[derive(Debug)]
+pub struct SparkPivotFirst {
+ signature: Signature,
+ value_type: DataType,
+ // Kept for `PartialEq`/`Hash` (identity of the aggregate for plan
comparison) and for the
+ // deterministic slot ordering `state_fields` needs. `HashMap` alone would
give us the map
+ // but not a stable order or a `Hash` impl.
+ pivot_values: Arc<Vec<ScalarValue>>,
+ pivot_index: Arc<HashMap<ScalarValue, usize>>,
+}
+
+impl PartialEq for SparkPivotFirst {
+ fn eq(&self, other: &Self) -> bool {
+ self.value_type == other.value_type && self.pivot_values ==
other.pivot_values
+ }
+}
+
+impl Eq for SparkPivotFirst {}
+
+impl std::hash::Hash for SparkPivotFirst {
+ fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
+ self.value_type.hash(state);
+ self.pivot_values.hash(state);
+ }
+}
+
+impl SparkPivotFirst {
+ pub fn new(value_type: DataType, pivot_values: Vec<ScalarValue>) -> Self {
+ let mut pivot_index = HashMap::with_capacity(pivot_values.len());
+ // Spark's PivotFirst uses the FIRST occurrence's index
(HashMap/TreeMap semantics), so
+ // when duplicates are somehow present we mirror that by only
inserting the first one.
+ // `pivot_key` can fold two distinct pivot values (`0.0` and `-0.0`)
onto one key, and can
+ // drop one entirely (NaN), so the index is not necessarily the same
length as the slot
+ // vector - the slot count is always `pivot_values.len()`.
+ for (i, v) in pivot_values.iter().enumerate() {
+ if let Some(key) = pivot_key(v.clone()) {
+ pivot_index.entry(key).or_insert(i);
+ }
+ }
+ Self {
+ signature: Signature::user_defined(Immutable),
+ value_type,
+ pivot_values: Arc::new(pivot_values),
+ pivot_index: Arc::new(pivot_index),
+ }
+ }
+}
+
+/// Rewrite a pivot column value into the key Spark would match it on, or
`None` when Spark can
+/// never match it.
+///
+/// Spark's `PivotFirst` looks pivot values up in a Scala `HashMap[Any, Int]`,
so matching goes
+/// through `BoxesRunTime.equals` / `Statics.anyHash` on the boxed Catalyst
value rather than
+/// through `ScalarValue`'s own equality. The two disagree on floats in
opposite directions:
+///
+/// * `-0.0` and `0.0` are one key for Spark (`-0.0 == 0.0` numerically, and
`doubleHash` folds
+/// both onto the hash of `0L`), while `ScalarValue` keeps them apart.
+/// * `NaN` matches nothing for Spark, not even another `NaN`, because Scala's
`==` on `Double`
+/// is IEEE. `ScalarValue` treats `NaN` as equal to itself.
+///
+/// Nulls are left alone: a null pivot column value does match a null entry in
the pivot list,
+/// which is what Spark's `pivotIndex.getOrElse(null, -1)` does.
+fn pivot_key(v: ScalarValue) -> Option<ScalarValue> {
+ match v {
+ ScalarValue::Float32(Some(f)) => {
+ if f.is_nan() {
+ None
+ } else if f == 0.0 {
+ Some(ScalarValue::Float32(Some(0.0)))
+ } else {
+ Some(ScalarValue::Float32(Some(f)))
+ }
+ }
+ ScalarValue::Float64(Some(f)) => {
+ if f.is_nan() {
+ None
+ } else if f == 0.0 {
+ Some(ScalarValue::Float64(Some(0.0)))
+ } else {
+ Some(ScalarValue::Float64(Some(f)))
+ }
+ }
+ other => Some(other),
+ }
+}
+
+impl AggregateUDFImpl for SparkPivotFirst {
+ fn name(&self) -> &str {
+ "pivot_first"
+ }
+
+ fn signature(&self) -> &Signature {
+ &self.signature
+ }
+
+ fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
+ Ok(DataType::List(Arc::new(Field::new_list_field(
+ self.value_type.clone(),
+ true,
+ ))))
+ }
+
+ fn state_fields(&self, args: StateFieldsArgs) -> DFResult<Vec<FieldRef>> {
+ // One field per pivot slot, matching Spark's aggBufferAttributes so
the shuffle
+ // exchange sees the same schema catalyst declared.
`format_state_name` is the same
+ // helper other aggregates in this crate use (see `avg.rs`,
`stddev.rs`).
+ Ok((0..self.pivot_values.len())
Review Comment:
[P2] Preserve Catalyst’s buffer width for duplicate pivot values. Spark
sizes `aggBufferAttributes` from `pivotIndex.size`, while this code declares
one state field per original list entry. With ANSI disabled, pivoting a row
`(1,'a',10)` over `IN ('a','b','b')` succeeds in Spark with `(1,10,NULL,NULL)`
and two aggregate-buffer fields. Native execution produces three buffer fields,
making the partial aggregate incompatible with Spark’s declared shuffle/FFI
schema and triggering column-count checks. Please fall back for duplicate lists
until both the buffer layout and Spark’s last-occurrence index mapping are
reproduced.
Evidence: Spark 4.1.3 inspection reported `list 3`, `HashMap(b -> 2, a ->
0)`, and `buffer_fields 2`, and returned `(1,10,NULL,NULL)`. The exact-source
Rust probe built the aggregate expression and observed three state fields with
`[Int64(10), Int64(NULL), Int64(NULL)]`. Combining its output with Catalyst’s
schema failed with `number of columns(4) must match number of fields(3) in
schema`. Comet emits raw partial buffers, and both `export_batch` and shuffle
decoding enforce column counts.
--
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]