sunchao commented on code in PR #4727:
URL: https://github.com/apache/datafusion-comet/pull/4727#discussion_r4124906434
##########
spark/src/main/scala/org/apache/spark/sql/comet/operators.scala:
##########
@@ -2349,6 +2366,48 @@ object CometObjectHashAggregateExec
op.child,
SerializedPlan(None))
}
+
+ /**
+ * For intermediate aggregates containing TypedImperativeAggregate functions
(like CollectSet or
+ * CollectList), Spark declares buffer columns as BinaryType because it
serializes the JVM
+ * state. Native Comet keeps the actual state type instead:
ArrayType(elementType) with
+ * containsNull true for CollectSet/CollectList. Rewrite the Spark-side
output attributes for
+ * Partial, PartialMerge, and mixed {Partial, PartialMerge} stages so
shuffle and downstream
+ * native aggregate stages see the schema that native execution really
produces.
+ *
+ * Final aggregates output user-visible values rather than intermediate
state, so their Spark
+ * result schema is left unchanged.
+ */
+ private def adjustOutputForNativeState(op: ObjectHashAggregateExec):
Seq[Attribute] = {
Review Comment:
[P2] Preserve the shared helper’s `Mode` schema adjustment. This new
overload is selected by `CometObjectHashAggregateExec.createExec` instead of
the inherited `CometBaseAggregate` helper, but it omits that helper’s `Mode`
case. Consequently, native `mode` partials advertise `BinaryType` while
emitting a struct of values and counts. With Comet shuffle enabled, ordinary
supported queries such as `SELECT mode(v) FROM mode_int` now abort during
schema reconciliation instead of returning the mode. Please reuse the shared
helper, which already handles `PartialMerge`, or preserve its complete state
mappings.
Evidence: Exact-head CI job
https://github.com/apache/datafusion-comet/actions/runs/36243582628/job/108409634203
reproduces failures in `mode.sql:47` and `mode_within_group.sql:42`:
`CometSchemaAlignExec cannot reconcile ... expected Binary, found
Struct("values": non-null List(Int32), "counts": non-null List(Int64))`. The
first fixture expects `5`. The base revision uses the inherited helper
containing the `Mode` mapping, while this added overload bypasses it.
##########
native/core/src/execution/spark_aggregate_state.rs:
##########
@@ -0,0 +1,688 @@
+// 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.
+
+//! Decoders for Spark JVM aggregate state consumed by native PartialMerge.
+
+use std::{borrow::Cow, sync::Arc};
+
+use arrow::array::{
+ builder::{make_builder, ArrayBuilder, ListBuilder},
+ Array, ArrayRef, BinaryArray, GenericByteArray, LargeBinaryArray,
OffsetSizeTrait,
+};
+use arrow::datatypes::{DataType, FieldRef, GenericBinaryType};
+use datafusion::common::{DataFusionError, Result};
+use datafusion::physical_expr::aggregate::AggregateFunctionExpr;
+use datafusion_comet_shuffle::spark_unsafe::list::{append_to_builder,
SparkUnsafeArray};
+
+/// Decoder used by `MergeAsPartial` before forwarding state to the inner
accumulator.
+#[derive(Clone, Debug)]
+pub(crate) enum PartialMergeStateDecoder {
+ PassThrough,
+ SparkCollect(SparkCollectStateDecoder),
+}
+
+impl PartialMergeStateDecoder {
+ pub(crate) fn try_new(
+ inner_expr: &AggregateFunctionExpr,
+ state_fields: &[FieldRef],
+ ) -> Result<Self> {
+ Ok(
+ match SparkCollectStateDecoder::try_new(inner_expr, state_fields)?
{
+ Some(decoder) => Self::SparkCollect(decoder),
+ None => Self::PassThrough,
+ },
+ )
+ }
+
+ pub(crate) fn decode<'a>(&self, values: &'a [ArrayRef]) -> Result<Cow<'a,
[ArrayRef]>> {
+ match self {
+ Self::PassThrough => Ok(Cow::Borrowed(values)),
+ Self::SparkCollect(decoder) => decoder.decode(values),
+ }
+ }
+}
+
+/// Decodes Spark JVM collect aggregate buffers into DataFusion collect state.
+///
+/// Spark's `CollectList` / `CollectSet` are `TypedImperativeAggregate`s. When
Spark runs the
+/// lower Partial aggregate, each buffer is serialized as a `BinaryType` value
containing a
+/// single-field `UnsafeRow`; field 0 is the `UnsafeArrayData` with the
collected elements.
+/// DataFusion's collect accumulators expect the merge input to be a
list-typed state column, so
+/// mixed Spark-Partial -> Comet-PartialMerge plans must materialize those
unsafe bytes into an
+/// Arrow `ListArray` before calling the inner accumulator's `merge_batch`.
The single-field
+/// `UnsafeRow` and nested `UnsafeArrayData` layouts used here are unchanged
across Spark 3.4-4.2.
+#[derive(Clone, Debug)]
+pub(crate) struct SparkCollectStateDecoder {
+ item_field: FieldRef,
+}
+
+impl SparkCollectStateDecoder {
+ fn try_new(
+ inner_expr: &AggregateFunctionExpr,
+ state_fields: &[FieldRef],
+ ) -> Result<Option<Self>> {
+ if !matches!(inner_expr.fun().name(), "collect_list" | "collect_set") {
+ return Ok(None);
+ }
+
+ let [state_field] = state_fields else {
+ return Err(DataFusionError::Internal(format!(
+ "Spark collect state decoder expected one state field, got {}",
+ state_fields.len()
+ )));
+ };
+ let DataType::List(item_field) = state_field.data_type() else {
+ return Err(DataFusionError::Internal(format!(
+ "Spark collect state decoder expected List state, got {}",
+ state_field.data_type()
+ )));
+ };
+
+ Ok(Some(Self {
+ item_field: Arc::clone(item_field),
+ }))
+ }
+
+ fn decode<'a>(&self, values: &'a [ArrayRef]) -> Result<Cow<'a,
[ArrayRef]>> {
+ if values.len() != 1 {
+ return Err(DataFusionError::Internal(format!(
+ "Spark collect state decoder expected one state column, got
{}",
+ values.len()
+ )));
+ }
+
+ match values[0].data_type() {
+ DataType::Binary => {
+ let array = values[0]
+ .as_any()
+ .downcast_ref::<BinaryArray>()
+ .ok_or_else(|| {
+ Self::decode_error("expected BinaryArray for Binary
collect state")
+ })?;
+ Ok(Cow::Owned(vec![self.decode_binary_array(array)?]))
+ }
+ DataType::LargeBinary => {
+ let array = values[0]
+ .as_any()
+ .downcast_ref::<LargeBinaryArray>()
+ .ok_or_else(|| {
+ Self::decode_error(
+ "expected LargeBinaryArray for LargeBinary collect
state",
+ )
+ })?;
+ Ok(Cow::Owned(vec![self.decode_binary_array(array)?]))
+ }
+ _ => Ok(Cow::Borrowed(values)),
+ }
+ }
+
+ fn decode_binary_array<O: OffsetSizeTrait>(
+ &self,
+ array: &GenericByteArray<GenericBinaryType<O>>,
+ ) -> Result<ArrayRef> {
+ let mut builder = self.new_list_builder(array.len());
+
+ for row_idx in 0..array.len() {
+ if array.is_null(row_idx) {
+ builder.append_null();
+ } else {
+ self.append_unsafe_row_array(array.value(row_idx), &mut
builder)?;
+ }
+ }
+
+ Ok(Arc::new(builder.finish()))
+ }
+
+ fn new_list_builder(&self, capacity: usize) -> ListBuilder<Box<dyn
ArrayBuilder>> {
+ let value_builder = make_builder(self.item_field.data_type(),
capacity);
+ ListBuilder::with_capacity(value_builder,
capacity).with_field(Arc::clone(&self.item_field))
+ }
+
+ fn append_unsafe_row_array(
+ &self,
+ row_bytes: &[u8],
+ builder: &mut ListBuilder<Box<dyn ArrayBuilder>>,
+ ) -> Result<()> {
+ match self.spark_array_from_single_field_unsafe_row(row_bytes)? {
+ Some(array) => {
+ append_to_builder::<true>(self.item_field.data_type(),
builder.values(), &array)
Review Comment:
[P2] Decode binary `CollectSet` elements using Spark’s actual buffer type.
Across the supported Spark versions, `CollectSet(BinaryType)` serializes each
value as `UnsafeArrayData(ArrayType(ByteType))`, then extracts its byte payload
during evaluation. This decoder instead uses the native `Binary` item type and
copies the entire nested array, including its header and padding. With Comet
shuffle enabled and a Spark `LocalTableScan` partial, `SELECT x, count(DISTINCT
y), collect_set(b) FROM VALUES (1,1,X'ABCD'), (1,2,X'ABCD') AS t(x,y,b) GROUP
BY x` therefore returns a corrupted binary value. Preserve the
aggregate-specific serialized element type and unwrap the byte array, or retain
fallback for this boundary.
Evidence: A bounded Rust test compiled from the exact-head decoder expected
`[171, 205]` but returned `[2, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
171, 205, 0, 0, 0, 0, 0, 0]`. Spark 4.1.3 `CollectSet.serialize` produced
exactly the test’s 64-byte buffer, and `CollectSet.eval` returned the original
`ABCD`. Spark’s `bufferElementType`, `convertToBufferElement`, and `eval`
implement this special binary representation in every inspected version. The
four existing decoder tests passed.
--
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]