andygrove commented on code in PR #5318:
URL: https://github.com/apache/datafusion-comet/pull/5318#discussion_r4178476399
##########
spark/src/main/scala/org/apache/comet/CometConf.scala:
##########
@@ -277,6 +277,13 @@ object CometConf extends ShimCometConf {
createExecEnabledConfig("emptyRelation", defaultValue = true)
val COMET_EXEC_SAMPLE_ENABLED: ConfigEntry[Boolean] =
createExecEnabledConfig("sample", defaultValue = true)
+ val COMET_EXEC_MERGE_ROWS_ENABLED: ConfigEntry[Boolean] =
+ createExecEnabledConfig(
+ "mergeRows",
+ defaultValue = false,
+ notes = Some(
+ "Ignored on Spark 4.1 and later, where MergeRowsExec remains on Spark
so V2 writers " +
Review Comment:
Since 4.1 is the default build, this flag does nothing there, and nothing
tracks the work to change that. Could you open an issue for native
`MergeRowsExec` on 4.1+ that covers the `MergeSummary` contract, and link it
from this note and the compatibility page? Could the note also say the flag
only takes effect on Spark 3.5 and 4.0, since 3.4 ignores it too? The comment
at `CometExecRule.scala:118` has the same gap.
##########
native/core/src/execution/operators/merge_rows.rs:
##########
@@ -0,0 +1,1486 @@
+// 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 arrow::array::{Array, ArrayRef, BooleanArray, Int64Array, RecordBatch};
+use arrow::compute::kernels::boolean::{and, and_not, not};
+use arrow::compute::{filter_record_batch, prep_null_mask_filter};
+use arrow::datatypes::{DataType, SchemaRef};
+use datafusion::common::tree_node::TreeNodeRecursion;
+use datafusion::common::utils::memory::estimate_memory_size;
+use datafusion::common::{DataFusionError, HashSet, ScalarValue};
+use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation};
+use datafusion::logical_expr::ColumnarValue;
+use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr};
+use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
+use datafusion::physical_plan::metrics::{BaselineMetrics,
ExecutionPlanMetricsSet, MetricsSet};
+use datafusion::{
+ execution::TaskContext,
+ physical_plan::{
+ apply_expression_roots, ChildrenPropertiesMode, DisplayAs,
DisplayFormatType,
+ ExecutionPlan, Partitioning, PlanProperties, RecordBatchStream,
ReplaceChildrenOptions,
+ SendableRecordBatchStream,
+ },
+};
+use datafusion_comet_common::{cast_and_stamp_schema, SparkError};
+use futures::{Stream, StreamExt};
+use std::{
+ pin::Pin,
+ sync::Arc,
+ task::{Context, Poll},
+};
+
+/// A MergeRows instruction: condition plus zero (Discard), one (Keep), or two
(Split)
+/// output row projections.
+#[derive(Debug, Clone)]
+pub struct MergeInstructionExec {
+ pub condition: Arc<dyn PhysicalExpr>,
+ pub outputs: Vec<Vec<Arc<dyn PhysicalExpr>>>,
+}
+
+#[derive(Debug)]
+struct MergeConfig {
+ is_source_row_present: Arc<dyn PhysicalExpr>,
+ is_target_row_present: Arc<dyn PhysicalExpr>,
+ matched_instructions: Vec<MergeInstructionExec>,
+ not_matched_instructions: Vec<MergeInstructionExec>,
+ not_matched_by_source_instructions: Vec<MergeInstructionExec>,
+ row_id_ordinal: Option<usize>,
+}
+
+impl MergeConfig {
+ fn validate(
+ &self,
+ child: &Arc<dyn ExecutionPlan>,
+ output_schema: &SchemaRef,
+ ) -> Result<(), DataFusionError> {
+ if let Some(ordinal) = self.row_id_ordinal {
+ let child_schema = child.schema();
+ let child_fields = child_schema.fields().len();
+ if ordinal >= child_fields {
+ return Err(DataFusionError::Internal(format!(
+ "MergeRows: row id ordinal {ordinal} is out of range for a
child with \
+ {child_fields} columns"
+ )));
+ }
+ let data_type = child_schema.field(ordinal).data_type();
+ if data_type != &DataType::Int64 {
+ return Err(DataFusionError::Internal(format!(
+ "MergeRows: row id column at ordinal {ordinal} must be
Int64, got {data_type}"
+ )));
+ }
+ }
+
+ let output_width = output_schema.fields().len();
+ for (group, instructions) in [
+ ("matched", &self.matched_instructions),
+ ("not matched", &self.not_matched_instructions),
+ (
+ "not matched by source",
+ &self.not_matched_by_source_instructions,
+ ),
+ ] {
+ for (instruction_index, instruction) in
instructions.iter().enumerate() {
+ if instruction.outputs.len() > 2 {
+ return Err(DataFusionError::Internal(format!(
+ "MergeRows: {group} instruction {instruction_index}
has {} output rows; expected at most 2",
+ instruction.outputs.len()
+ )));
+ }
+ for (output_index, output) in
instruction.outputs.iter().enumerate() {
+ if output.len() != output_width {
+ return Err(DataFusionError::Internal(format!(
+ "MergeRows: {group} instruction
{instruction_index} output {output_index} has {} expressions; expected
{output_width}",
+ output.len()
+ )));
+ }
+ }
+ }
+ }
+ Ok(())
+ }
+}
+
+#[derive(Debug)]
+pub struct MergeRowsExec {
+ config: Arc<MergeConfig>,
+ child: Arc<dyn ExecutionPlan>,
+ schema: SchemaRef,
+ cache: Arc<PlanProperties>,
+ metrics: ExecutionPlanMetricsSet,
+}
+
+impl MergeRowsExec {
+ #[allow(clippy::too_many_arguments)]
+ pub fn try_new(
+ is_source_row_present: Arc<dyn PhysicalExpr>,
+ is_target_row_present: Arc<dyn PhysicalExpr>,
+ matched_instructions: Vec<MergeInstructionExec>,
+ not_matched_instructions: Vec<MergeInstructionExec>,
+ not_matched_by_source_instructions: Vec<MergeInstructionExec>,
+ row_id_ordinal: Option<usize>,
+ child: Arc<dyn ExecutionPlan>,
+ schema: SchemaRef,
+ ) -> Result<Self, DataFusionError> {
+ let config = Arc::new(MergeConfig {
+ is_source_row_present,
+ is_target_row_present,
+ matched_instructions,
+ not_matched_instructions,
+ not_matched_by_source_instructions,
+ row_id_ordinal,
+ });
+ config.validate(&child, &schema)?;
+
+ let cache = Arc::new(PlanProperties::new(
+ EquivalenceProperties::new(Arc::clone(&schema)),
+ Partitioning::UnknownPartitioning(1),
+ EmissionType::Incremental,
+ Boundedness::Bounded,
+ ));
+
+ Ok(Self {
+ config,
+ child,
+ schema,
+ cache,
+ metrics: ExecutionPlanMetricsSet::new(),
+ })
+ }
+}
+
+impl DisplayAs for MergeRowsExec {
+ fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) ->
std::fmt::Result {
+ match t {
+ DisplayFormatType::Default | DisplayFormatType::Verbose => {
+ write!(f, "CometMergeRowsExec")
+ }
+ DisplayFormatType::TreeRender => unimplemented!(),
+ }
+ }
+}
+
+impl ExecutionPlan for MergeRowsExec {
+ fn schema(&self) -> SchemaRef {
+ Arc::clone(&self.schema)
+ }
+
+ fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
+ vec![&self.child]
+ }
+
+ fn replace_children(
+ self: Arc<Self>,
+ children: Vec<Arc<dyn ExecutionPlan>>,
+ _options: ReplaceChildrenOptions,
+ ) -> datafusion::common::Result<Arc<dyn ExecutionPlan>> {
+ let [child] = children.as_slice() else {
+ return Err(DataFusionError::Internal(format!(
+ "MergeRows expects exactly one child, got {}",
+ children.len()
+ )));
+ };
+ let child = Arc::clone(child);
+ self.config.validate(&child, &self.schema)?;
+ Ok(Arc::new(MergeRowsExec {
+ config: Arc::clone(&self.config),
+ child,
+ schema: Arc::clone(&self.schema),
+ cache: Arc::clone(&self.cache),
+ metrics: self.metrics.clone(),
+ }))
+ }
+
+ fn apply_expressions(
+ &self,
+ f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) ->
datafusion::common::Result<TreeNodeRecursion>,
+ ) -> datafusion::common::Result<TreeNodeRecursion> {
+ let instructions = self
+ .config
+ .matched_instructions
+ .iter()
+ .chain(self.config.not_matched_instructions.iter())
+ .chain(self.config.not_matched_by_source_instructions.iter());
+ let instruction_expressions = instructions.flat_map(|instruction| {
+ std::iter::once(&instruction.condition)
+ .chain(instruction.outputs.iter().flat_map(|output|
output.iter()))
+ });
+
+ apply_expression_roots(
+ [
+ &self.config.is_source_row_present,
+ &self.config.is_target_row_present,
+ ]
+ .into_iter()
+ .chain(instruction_expressions),
+ f,
+ )
+ }
+
+ fn with_new_children(
+ self: Arc<Self>,
+ children: Vec<Arc<dyn ExecutionPlan>>,
+ ) -> datafusion::common::Result<Arc<dyn ExecutionPlan>> {
+ self.replace_children(
+ children,
+ ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
+ )
+ }
+
+ fn execute(
+ &self,
+ partition: usize,
+ context: Arc<TaskContext>,
+ ) -> datafusion::common::Result<SendableRecordBatchStream> {
+ let reservation = self.config.row_id_ordinal.map(|_| {
+ MemoryConsumer::new(format!("CometMergeRowsExec[{partition}]"))
+ .register(&context.runtime_env().memory_pool)
+ });
+ let child_stream = self.child.execute(partition,
Arc::clone(&context))?;
+ Ok(Box::pin(MergeRowsStream {
+ config: Arc::clone(&self.config),
+ child_stream,
+ schema: Arc::clone(&self.schema),
+ seen: HashSet::new(),
+ reservation,
+ baseline: BaselineMetrics::new(&self.metrics, partition),
+ }))
+ }
+
+ fn properties(&self) -> &Arc<PlanProperties> {
+ &self.cache
+ }
+
+ fn metrics(&self) -> Option<MetricsSet> {
+ Some(self.metrics.clone_inner())
+ }
+
+ fn name(&self) -> &str {
+ "CometMergeRowsExec"
+ }
+}
+
+pub struct MergeRowsStream {
+ config: Arc<MergeConfig>,
+ child_stream: SendableRecordBatchStream,
+ schema: SchemaRef,
+ // Partition-scoped so duplicate matches across Arrow batches are still
detected.
+ seen: HashSet<i64>,
Review Comment:
Spark's `BitmapCardinalityValidator` keeps matched row ids in a
`Roaring64Bitmap`, and this uses a `HashSet<i64>`. I measured both on 5M row
ids. The set takes 75 MB, about 15 bytes per id. A roaring treemap takes 0.6 MB
when the ids are dense, as in a co-partitioned join, and 5 to 17 MB when they
are a shuffled share of 8 to 200 target partitions. Because this reservation
can't spill, a MERGE with many matched rows per task can fail with
`ResourcesExhausted` where Spark succeeds. `roaring` is already compiled in
through iceberg-rust, and `RoaringBitmap::statistics()` reports capacity bytes
that could size the reservation once per batch. Could you switch to
`RoaringTreemap` here? If you'd rather keep that out of this PR, could you open
an issue for it and add a line about the memory cost to the compatibility page?
##########
spark/src/test/scala/org/apache/comet/exec/CometMergeRowsNativeSuiteBase.scala:
##########
@@ -0,0 +1,499 @@
+/*
+ * 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.
+ */
+
+package org.apache.comet.exec
+
+import org.apache.spark.{CometListenerBusUtils, SparkConf}
+import org.apache.spark.sql.CometTestBase
+import org.apache.spark.sql.catalyst.expressions.SubqueryExpression
+import org.apache.spark.sql.comet.CometMergeRowsExec
+import
org.apache.spark.sql.connector.catalog.InMemoryRowLevelOperationTableCatalog
+import org.apache.spark.sql.execution.{QueryExecution, ScalarSubquery}
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
+import org.apache.spark.sql.util.QueryExecutionListener
+
+import org.apache.comet.CometConf
+import org.apache.comet.CometSparkSessionExtensions.isSpark35Plus
+
+/**
+ * `CometMergeRowsExec` converts Spark's `MergeRowsExec` and the `MergeRows`
logical node that
+ * feeds it, both defined in Spark core (`execution.datasources.v2` /
`catalyst.plans.logical`),
+ * not in any connector module. Spark plans a `MergeRowsExec` for MERGE INTO
against any
+ * `SupportsRowLevelOperations` V2 table using group-based (copy-on-write)
planning, independent
+ * of which connector implements the table.
+ *
+ * This suite pins that contract against Spark's own
`InMemoryRowLevelOperationTableCatalog` test
+ * catalog rather than Iceberg. `InMemoryRowLevelOperationTable` selects its
write shape via the
+ * `supports-deltas` table property: unset (default `false`) plans through
group-based
+ * `MergeRows`; `supports-deltas=true` plans a JVM `WriteDelta` whose child is
*also* a
+ * `MergeRowsExec` (and, with `split-updates=true`, emits `Split` for
update-as-delete+insert).
+ * Comet's bottom-up conversion makes that child eligible for native execution
under the JVM
+ * write, so both shapes are exercised here (`deltaMergeCase`) and checked
against a pure-Spark
+ * baseline. See `CometIcebergWriteActionSuite` for MERGE INTO coverage
against real Iceberg
+ * tables.
+ */
+abstract class CometMergeRowsNativeSuiteBase extends CometTestBase with
AdaptiveSparkPlanHelper {
+
+ private val catalog = "generic_rowlevel"
+
+ override protected def sparkConf: SparkConf = {
+ super.sparkConf
+ .set(s"spark.sql.catalog.$catalog",
classOf[InMemoryRowLevelOperationTableCatalog].getName)
+ // The test catalog's BatchScan is not a Comet-native scan. A broadcast
join over these tiny
+ // fixtures therefore leaves MergeRowsExec on the JVM instead of
exercising this suite's
+ // operator. Force the supported shuffle boundary while keeping AQE
enabled.
+ .set("spark.sql.autoBroadcastJoinThreshold", "-1")
+ .set("spark.sql.adaptive.autoBroadcastJoinThreshold", "-1")
+ .set("spark.sql.shuffle.partitions", "4")
+ }
+
+ private def assumeMerge(): Unit = assume(isSpark35Plus, "MergeRowsExec
requires Spark 3.5+")
+
+ test("MERGE all row-routing groups engage CometMergeRowsExec and match
Spark") {
+ assumeMerge()
+ val target = s"$catalog.default.rowlevel_target"
+ val source = s"$catalog.default.rowlevel_source"
+
+ def resetTables(): Unit = {
+ sql(s"DROP TABLE IF EXISTS $target")
+ sql(s"DROP TABLE IF EXISTS $source")
+ sql(s"CREATE TABLE $target (id INT, region STRING, amount DOUBLE) USING
parquet")
+ sql(s"CREATE TABLE $source (id INT, region STRING, amount DOUBLE) USING
parquet")
+ sql(
+ s"INSERT INTO $target VALUES " +
+ (0 until 20).map(i => s"($i, 'r${i % 3}', ${i * 1.5})").mkString(",
"))
+ sql(
+ s"INSERT INTO $source VALUES " +
+ (10 until 30).map(i => s"($i, 's${i % 3}', ${i * 2.0})").mkString(",
"))
+ }
+
+ // Exercise all three MergeRows instruction groups in one end-to-end
query: 10-19 are
+ // MATCHED, 20-29 are NOT MATCHED, and 0-9 are NOT MATCHED BY SOURCE. The
predicate on the
+ // target-only clause also leaves 5-9 to Spark's generated catch-all Keep
instruction.
+ val mergeSql =
+ s"""MERGE INTO $target t USING $source s ON t.id = s.id
+ |WHEN MATCHED THEN UPDATE SET t.amount = s.amount, t.region = s.region
+ |WHEN NOT MATCHED THEN INSERT (id, region, amount) VALUES (s.id,
s.region, s.amount)
+ |WHEN NOT MATCHED BY SOURCE AND t.id < 5 THEN UPDATE SET t.amount =
t.amount + 1000.0
+ |""".stripMargin
+
+ val captured = scala.collection.mutable.ArrayBuffer[QueryExecution]()
+ val listener = new QueryExecutionListener {
+ override def onSuccess(funcName: String, qe: QueryExecution, durationNs:
Long): Unit =
+ captured += qe
+ override def onFailure(funcName: String, qe: QueryExecution, exception:
Exception): Unit =
+ ()
+ }
+ spark.listenerManager.register(listener)
+ try {
+ resetTables()
+ captured.clear()
+ withSQLConf(
+ CometConf.COMET_ENABLED.key -> "true",
Review Comment:
`CometTestBase` already enables Comet, so the `CometConf.COMET_ENABLED.key
-> "true"` entries here, at lines 171 and 248, and in the 4.1+
`CometMergeRowsSuite` at line 74 can go. This was in my first review too.
--
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]