From b9b2c2d0d614e2514212dd0aa387057324b125da Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Thu, 17 Sep 2026 20:53:58 +0000 Subject: [PATCH 1/5] fix: route wide decimal hash shuffle keys through Spark-compatible hashing --- .../user-guide/latest/tuning/shuffle.md | 7 ++ .../shuffle/CometShuffleExchangeExec.scala | 10 +- .../approved-plans-v1_4/q49/extended.txt | 2 +- .../q70a/extended.txt | 2 +- .../q14a/extended.txt | 2 +- .../approved-plans-v2_7/q14a/extended.txt | 2 +- .../approved-plans-v2_7/q36a/extended.txt | 2 +- .../approved-plans-v2_7/q49/extended.txt | 2 +- .../approved-plans-v2_7/q5a/extended.txt | 2 +- .../approved-plans-v2_7/q70a/extended.txt | 2 +- .../approved-plans-v2_7/q77a/extended.txt | 2 +- .../approved-plans-v2_7/q80a/extended.txt | 2 +- .../approved-plans-v2_7/q86a/extended.txt | 2 +- .../comet/CometFuzzAggregateSuite.scala | 28 +++++- .../org/apache/comet/CometFuzzTestSuite.scala | 24 ++++- .../CometSparkSessionExtensionsSuite.scala | 66 ++++++++++++- .../comet/exec/CometNativeShuffleSuite.scala | 90 +++++++++++++++++- .../CometWideDecimalShuffleBenchmark.scala | 95 +++++++++++++++++++ 18 files changed, 312 insertions(+), 30 deletions(-) create mode 100644 spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala diff --git a/docs/source/user-guide/latest/tuning/shuffle.md b/docs/source/user-guide/latest/tuning/shuffle.md index 71c9c431ab0..720fd119770 100644 --- a/docs/source/user-guide/latest/tuning/shuffle.md +++ b/docs/source/user-guide/latest/tuning/shuffle.md @@ -53,6 +53,13 @@ scalar types. Hash partitioning keys must be scalar types unless and later) map keys. That setting is disabled by default until the performance of the nested hashing paths has been measured. Columns that are not partitioning keys may contain complex types like maps, structs, and arrays. +Hash partitioning on decimal keys with precision greater than 18 falls back because native hashing does not match +Spark's partition assignments. This can affect decimal aggregate overflow behavior, including `AVG(DISTINCT ...)`. +With `spark.comet.shuffle.mode=auto`, Comet uses Columnar Shuffle when eligible; with `native`, it uses Spark shuffle. +The restriction applies recursively to hash partitioning keys, including decimals inside structs, arrays, and maps +when nested hash partitioning is enabled. Wider decimals remain supported as payload columns, range partitioning +keys, and in single-partition shuffles. + ### Columnar (JVM) Shuffle Comet Columnar shuffle is JVM-based and supports `HashPartitioning`, `RoundRobinPartitioning`, `RangePartitioning`, and diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index 7a0d231a141..90b01fee5a0 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -555,12 +555,10 @@ object CometShuffleExchangeExec _: FloatType | _: DoubleType | _: StringType | _: BinaryType | _: TimestampType | _: TimestampNTZType | _: DateType => true - case _: DecimalType => - // TODO enforce this check - // https://github.com/apache/datafusion-comet/issues/3079 - // Decimals with precision > 18 require Java BigDecimal conversion before hashing - // d.precision <= 18 - true + case d: DecimalType => + // Spark hashes wider decimals through BigInteger bytes, which native hashing does not + // match. Different partition assignments can change decimal aggregate overflow behavior. + d.precision <= 18 case dt if isTimeType(dt) => true case StructType(fields) if nestedHashPartitioningEnabled => diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q49/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q49/extended.txt index 6deb75bea15..1f9d12fc934 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q49/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q49/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometProject diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt index 185ce8a5247..09918cf843c 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt @@ -5,7 +5,7 @@ CometColumnarToRow +- CometSort +- CometExchange +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark4_1/q14a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark4_1/q14a/extended.txt index 9a11b4895a7..2a3100957f1 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark4_1/q14a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark4_1/q14a/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q14a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q14a/extended.txt index a4ef99efa8e..8f5b9835de5 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q14a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q14a/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q36a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q36a/extended.txt index f11c15706a4..04405d9b84e 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q36a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q36a/extended.txt @@ -5,7 +5,7 @@ CometColumnarToRow +- CometSort +- CometExchange +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q49/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q49/extended.txt index 6deb75bea15..1f9d12fc934 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q49/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q49/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometProject diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q5a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q5a/extended.txt index f42b9795c47..4b41ce8a173 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q5a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q5a/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q70a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q70a/extended.txt index 594ae116008..6b6725c38f5 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q70a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q70a/extended.txt @@ -5,7 +5,7 @@ CometColumnarToRow +- CometSort +- CometExchange +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q77a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q77a/extended.txt index daccf86c3f7..a3fbdd2c329 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q77a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q77a/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q80a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q80a/extended.txt index f4db0be1188..ea972a9b91d 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q80a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q80a/extended.txt @@ -1,7 +1,7 @@ CometColumnarToRow +- CometTakeOrderedAndProject +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q86a/extended.txt b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q86a/extended.txt index 6269f5ddc51..b2fe83f4b36 100644 --- a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q86a/extended.txt +++ b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q86a/extended.txt @@ -5,7 +5,7 @@ CometColumnarToRow +- CometSort +- CometExchange +- CometHashAggregate - +- CometExchange + +- CometColumnarExchange +- CometHashAggregate +- CometUnion :- CometHashAggregate diff --git a/spark/src/test/scala/org/apache/comet/CometFuzzAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/CometFuzzAggregateSuite.scala index 3674887f80f..32b21cb92ad 100644 --- a/spark/src/test/scala/org/apache/comet/CometFuzzAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometFuzzAggregateSuite.scala @@ -19,6 +19,10 @@ package org.apache.comet +import org.apache.spark.sql.execution.aggregate.HashAggregateExec +import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec +import org.apache.spark.sql.types.DecimalType + import org.apache.comet.DataTypeSupport.isComplexType class CometFuzzAggregateSuite extends CometFuzzTestBase { @@ -31,7 +35,16 @@ class CometFuzzAggregateSuite extends CometFuzzTestBase { val (_, cometPlan) = checkSparkAnswer(sql) assert(1 == collectNativeScans(cometPlan).length) - checkSparkAnswerAndOperator(sql) + val hasWideDecimalKey = df.schema(col).dataType match { + case d: DecimalType => d.precision > 18 + case _ => false + } + // Wide decimal hash keys require Spark shuffle when columnar shuffle is disabled. + if (CometConf.COMET_SHUFFLE_MODE.get() == "native" && hasWideDecimalKey) { + checkSparkAnswerAndOperator(sql, classOf[HashAggregateExec], classOf[ShuffleExchangeExec]) + } else { + checkSparkAnswerAndOperator(sql) + } } } @@ -55,7 +68,18 @@ class CometFuzzAggregateSuite extends CometFuzzTestBase { val (_, cometPlan) = checkSparkAnswer(sql) assert(1 == collectNativeScans(cometPlan).length) - checkSparkAnswerAndOperator(sql) + val hasWideDecimalKey = Seq("c1", "c2", "c3", col).exists { key => + df.schema(key).dataType match { + case d: DecimalType => d.precision > 18 + case _ => false + } + } + // Check both GROUP BY and DISTINCT keys, not unrelated decimal payload columns. + if (CometConf.COMET_SHUFFLE_MODE.get() == "native" && hasWideDecimalKey) { + checkSparkAnswerAndOperator(sql, classOf[HashAggregateExec], classOf[ShuffleExchangeExec]) + } else { + checkSparkAnswerAndOperator(sql) + } } } diff --git a/spark/src/test/scala/org/apache/comet/CometFuzzTestSuite.scala b/spark/src/test/scala/org/apache/comet/CometFuzzTestSuite.scala index 55fb6010c99..91bd098c1f3 100644 --- a/spark/src/test/scala/org/apache/comet/CometFuzzTestSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometFuzzTestSuite.scala @@ -160,6 +160,17 @@ class CometFuzzTestSuite extends CometFuzzTestBase { } test("distribute by single column (complex types)") { + // Inspect only the key schema: any wide decimal leaf requires Spark's hash partition + // assignments, even when native hashing of nested keys is enabled. + def hasWideDecimal(dataType: DataType): Boolean = dataType match { + case decimal: DecimalType => decimal.precision > 18 + case StructType(fields) => fields.exists(field => hasWideDecimal(field.dataType)) + case ArrayType(elementType, _) => hasWideDecimal(elementType) + case MapType(keyType, valueType, _) => + hasWideDecimal(keyType) || hasWideDecimal(valueType) + case _ => false + } + val df = spark.read.parquet(filename) df.createOrReplaceTempView("t1") val columns = df.schema.fields.filter(f => isComplexType(f.dataType)).map(_.name) @@ -180,15 +191,20 @@ class CometFuzzTestSuite extends CometFuzzTestBase { } assert(cometShuffleExchanges.length == expectedNumCometShuffles) - // With the config enabled these keys do run through native shuffle. This is the widest - // nested-type coverage in the repo, so it is worth asserting that they are admitted rather - // than only that they fall back. + // Enabling nested keys admits supported types, but wide decimal leaves still require + // Spark's hash partition assignments. JVM shuffle supports both kinds of key. withSQLConf(CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_NESTED_ENABLED.key -> "true") { val enabledDf = spark.sql(sql) enabledDf.collect() val enabledPlan = enabledDf.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec].executedPlan - assert(collectCometShuffleExchanges(enabledPlan).length == 1) + val expectedEnabledShuffles = + if (CometConf.COMET_SHUFFLE_MODE.get() == "native" && + hasWideDecimal(df.schema(col).dataType)) 0 + else 1 + assert( + collectCometShuffleExchanges(enabledPlan).length == expectedEnabledShuffles, + s"Unexpected shuffle for ${df.schema(col)}") } } } diff --git a/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala b/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala index 92bac5dba86..3464765aba3 100644 --- a/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala @@ -20,11 +20,11 @@ package org.apache.comet import org.apache.spark.sql._ -import org.apache.spark.sql.catalyst.expressions.AttributeReference +import org.apache.spark.sql.catalyst.expressions.{Ascending, AttributeReference, SortOrder} import org.apache.spark.sql.catalyst.plans.logical.LocalRelation -import org.apache.spark.sql.catalyst.plans.physical.{RoundRobinPartitioning, SinglePartition} -import org.apache.spark.sql.comet.CometScanWrapper -import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometColumnarShuffle, CometShuffleExchangeExec} +import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning, RoundRobinPartitioning, SinglePartition} +import org.apache.spark.sql.comet.{CometScanWrapper, CometSinkPlaceHolder} +import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec, CometShuffleManager} import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.LongType @@ -130,6 +130,64 @@ class CometSparkSessionExtensionsSuite extends CometTestBase { assert(after eq before, s"reload unpacked another copy of the native library: $after") } + test("wide decimal hash keys use Spark-compatible shuffle partitioning") { + // Check the native precision boundary, then routing of unsupported keys in each mode. + for (mode <- Seq("native", "auto", "jvm"); precision <- Seq(18, 19, 38)) { + withSQLConf( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> mode, + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_NATIVE_RANGE_PARTITIONING_ENABLED.key -> "true", + "spark.shuffle.manager" -> classOf[CometShuffleManager].getName) { + val originalChild = spark + .range(1) + .selectExpr(s"CAST(id AS DECIMAL($precision, 0)) AS d", "id") + .queryExecution + .executedPlan + val child = CometSinkPlaceHolder( + OperatorOuterClass.Operator.getDefaultInstance, + originalChild, + originalChild) + val shuffle = ShuffleExchangeExec(HashPartitioning(Seq(child.output.head), 2), child) + val expected = if (mode == "jvm" || (mode == "auto" && precision > 18)) { + Some(CometColumnarShuffle) + } else if (precision <= 18) { + Some(CometNativeShuffle) + } else { + None + } + + withClue(s"mode=$mode, precision=$precision: ") { + assert(CometShuffleExchangeExec.shuffleSupported(shuffle) == expected) + if (expected.isEmpty) { + assert( + shuffle + .getTagValue(CometExplainInfo.FALLBACK_REASONS) + .getOrElse(Set.empty[String]) + .exists(_.contains("unsupported hash partitioning data type for native shuffle"))) + } else { + // A native failure must not tag an exchange that can use the columnar path. + assert(shuffle.getTagValue(CometExplainInfo.FALLBACK_REASONS).isEmpty) + } + + if (mode == "native" && precision > 18) { + // Wide decimals remain supported as payloads, range keys, and in a single partition. + Seq( + HashPartitioning(Seq(child.output(1)), 2), + RangePartitioning(Seq(SortOrder(child.output.head, Ascending)), 2), + SinglePartition).foreach { partitioning => + val supported = ShuffleExchangeExec(partitioning, child) + assert( + CometShuffleExchangeExec.shuffleSupported(supported).contains(CometNativeShuffle)) + } + } + } + } + } + } + test("Arrow properties") { NativeBase.setLoaded(false) NativeBase.load() diff --git a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala index b817ed7a503..f4654b32c37 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala @@ -47,9 +47,9 @@ import org.apache.spark.sql.comet.execution.shuffle.{CometNativeShuffle, CometSh import org.apache.spark.sql.execution.{LocalTableScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, AQEShuffleReadExec, ShuffleQueryStageExec} import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} -import org.apache.spark.sql.functions.{broadcast, col, count, countDistinct, sum} +import org.apache.spark.sql.functions.{broadcast, col, count, countDistinct, spark_partition_id, sum} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{ArrayType, DataType, LongType, MapType, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, DataType, DecimalType, LongType, MapType, StructField, StructType} import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.apache.comet.{CometConf, CometExecIterator, CometExplainInfo, CometShuffleBlockIterator, CometShuffleSizeLimitException, Native} @@ -440,7 +440,91 @@ class CometNativeShuffleSuite extends CometTestBase with AdaptiveSparkPlanHelper val shuffled = df .select($"_1") .repartition(10, col(c)) - checkShuffleAnswer(shuffled, 1, checkNativeOperators = true) + val nativeHashSupported = df.schema(c).dataType match { + case d: DecimalType => d.precision <= 18 + case _ => true + } + checkShuffleAnswer( + shuffled, + if (nativeHashSupported) 1 else 0, + checkNativeOperators = nativeHashSupported) + } + } + } + } + } + } + + for (precision <- Seq(18, 19, 38)) { + test(s"decimal hash shuffle preserves Spark partitions at precision $precision") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withTable("decimal_shuffle") { + sql(s"CREATE TABLE decimal_shuffle(id INT, k DECIMAL($precision, 0)) USING parquet") + val maximum = "9" * precision + sql(s"""INSERT INTO decimal_shuffle VALUES + |(0, NULL), (1, 0), (2, 1), (3, -1), (4, $maximum), (5, -$maximum) + |""".stripMargin) + for (mode <- Seq("native", "auto", "jvm")) { + withSQLConf(CometConf.COMET_SHUFFLE_MODE.key -> mode) { + val shuffled = spark.table("decimal_shuffle").repartition(7, $"k") + val native = precision <= 18 && mode != "jvm" + val sparkShuffle = precision > 18 && mode == "native" + checkCometExchange(shuffled, if (sparkShuffle) 0 else 1, native) + assert(shuffled.queryExecution.executedPlan.collect { case _: ShuffleExchangeExec => + 1 + }.sum == (if (sparkShuffle) 1 else 0)) + // Result equality alone cannot detect a different hash partition assignment. + checkSparkAnswer(shuffled.withColumn("partition", spark_partition_id())) + } + } + } + } + } + } + + test("wide decimals remain supported in shuffle payloads, ranges and single partitions") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_NATIVE_RANGE_PARTITIONING_ENABLED.key -> "true") { + withTable("decimal_shuffle") { + sql("CREATE TABLE decimal_shuffle(id INT, k DECIMAL(38, 0)) USING parquet") + sql( + "INSERT INTO decimal_shuffle VALUES (0, NULL), (1, 1), (2, -1), " + + "(3, 99999999999999999999999999999999999999)") + for (mode <- Seq("native", "auto", "jvm")) { + withSQLConf(CometConf.COMET_SHUFFLE_MODE.key -> mode) { + val input = spark.table("decimal_shuffle") + Seq( + input.repartition(7, $"id"), + input.repartitionByRange(7, $"k"), + input.repartition(1)).foreach { shuffled => + checkCometExchange(shuffled, 1, native = mode != "jvm") + checkSparkAnswer(shuffled) + } + } + } + } + } + } + + test("decimal hash shuffle checks nested keys recursively") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withNestedHashPartitioning { + for (precision <- Seq(18, 38)) { + withTable("decimal_shuffle") { + sql(s"""CREATE TABLE decimal_shuffle( + |id INT, s STRUCT>, + |a ARRAY>) USING parquet + |""".stripMargin) + sql("""INSERT INTO decimal_shuffle VALUES + |(0, NULL, NULL), + |(1, named_struct('a', array(1, -1)), array(named_struct('d', 1))), + |(2, named_struct('a', array(2, NULL)), array(named_struct('d', NULL))) + |""".stripMargin) + for (key <- Seq("s", "a")) { + val shuffled = spark.table("decimal_shuffle").repartition(7, col(key)) + checkCometExchange(shuffled, if (precision <= 18) 1 else 0, native = true) + checkSparkAnswer(shuffled.withColumn("partition", spark_partition_id())) } } } diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala new file mode 100644 index 00000000000..14429f60302 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala @@ -0,0 +1,95 @@ +/* + * 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.spark.sql.benchmark + +import org.apache.spark.benchmark.Benchmark +import org.apache.spark.sql.Row +import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec +import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.CometConf + +/** + * Measures wide-decimal hash shuffle routing with COUNT, independently of decimal AVG support. + * Run this benchmark on both revisions to compare the routing change. Fixture creation, result + * validation and route reporting are untimed. The session uses local[1]; `--reverse` reverses + * case order, and `--validate-only` checks the fixture and results without collecting timings. + */ +object CometWideDecimalShuffleBenchmark extends CometBenchmarkBase { + override def runCometBenchmark(mainArgs: Array[String]): Unit = { + val rows = 1024 * 1024 + val groups = 10000 + val partitions = 4 + val filePartitionBytes = 16 * 1024 * 1024 + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> partitions.toString, + // Keep each file in one split, and prevent Spark from combining files into one task. + SQLConf.FILES_MAX_PARTITION_BYTES.key -> filePartitionBytes.toString, + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> filePartitionBytes.toString, + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { + withTempPath { dir => + withTempTable("parquetV1Table") { + val query = "SELECT k, COUNT(v) FROM parquetV1Table GROUP BY k" + var expected = Seq.empty[Row] + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + prepareTable( + dir, + spark + .range(0L, rows.toLong, 1L, partitions) + .selectExpr( + s"CAST(id % $groups AS DECIMAL(38, 2)) AS k", + "CAST(id % 97 AS INT) AS v")) + assert(spark.table("parquetV1Table").rdd.getNumPartitions == partitions) + expected = spark.sql(query).collect().toSeq.sortBy(_.getDecimal(0)) + } + val benchmark = new Benchmark("wide_decimal_hash_shuffle", rows, output = output) + val modes = Seq("Spark", "native", "auto", "jvm") + for (mode <- (if (mainArgs.contains("--reverse")) modes.reverse else modes)) { + val configs = Seq( + CometConf.COMET_ENABLED.key -> (mode != "Spark").toString, + CometConf.COMET_SHUFFLE_MODE.key -> (if (mode == "Spark") "native" else mode)) + withSQLConf(configs: _*) { + val df = spark.sql(query) + assert(df.collect().toSeq.sortBy(_.getDecimal(0)) == expected) + val plan = df.queryExecution.executedPlan + val routes = plan.collect { + case exchange: CometShuffleExchangeExec => exchange.shuffleType.toString + case _: ShuffleExchangeExec => "Spark shuffle" + } + assert(routes.size == 1, plan.treeString) + benchmark.out.println(s"Wide-decimal shuffle mode=$mode: ${routes.head}") + benchmark.out.println(plan.treeString) + } + benchmark.addCase(mode) { _ => + withSQLConf(configs: _*) { + spark.sql(query).collect() + } + } + } + if (!mainArgs.contains("--validate-only")) benchmark.run() + } + } + } + } +} From 9fce8b5d6080b0ac2303260090ca3b32070de131 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Thu, 17 Sep 2026 20:57:17 +0000 Subject: [PATCH 2/5] test: retain safe aggregate buffers after wide decimal shuffle fallback --- .../comet/exec/CometNativeShuffleSuite.scala | 32 ++++++++++++++++++- 1 file changed, 31 insertions(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala index f4654b32c37..8832133a399 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala @@ -41,11 +41,12 @@ import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.catalyst.expressions.aggregate.Final import org.apache.spark.sql.catalyst.plans.logical.LocalRelation import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning -import org.apache.spark.sql.comet.{CometExec, CometHashAggregateExec, CometLocalTableScanExec, CometMetricNode, CometNativeExec, CometScanWrapper, CometSparkToColumnarExec, CometTakeOrderedAndProjectExec} +import org.apache.spark.sql.comet.{CometExec, CometHashAggregateExec, CometLocalTableScanExec, CometMetricNode, CometNativeExec, CometNativeScanExec, CometScanWrapper, CometSparkToColumnarExec, CometTakeOrderedAndProjectExec} import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.comet.execution.shuffle.{CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution.{LocalTableScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, AQEShuffleReadExec, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.aggregate.ObjectHashAggregateExec import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} import org.apache.spark.sql.functions.{broadcast, col, count, countDistinct, spark_partition_id, sum} import org.apache.spark.sql.internal.SQLConf @@ -507,6 +508,35 @@ class CometNativeShuffleSuite extends CometTestBase with AdaptiveSparkPlanHelper } } + test("wide decimal shuffle fallback keeps collection aggregate buffers in Spark") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.USE_OBJECT_HASH_AGG.key -> "true", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true") { + withParquetTable((0 until 100).map(i => (i % 3, i % 7)), "decimal_shuffle") { + for (precision <- Seq(18, 38); function <- Seq("collect_list", "collect_set")) { + val key = s"CAST(_1 AS DECIMAL($precision, 0))" + val df = + sql(s"SELECT $key, sort_array($function(_2)) FROM decimal_shuffle GROUP BY $key") + val plan = df.queryExecution.executedPlan + val nativeExpected = precision <= 18 + assert( + plan.collect { case _: CometHashAggregateExec => 1 }.sum == + (if (nativeExpected) 2 else 0), + plan.treeString) + assert( + plan.collect { case _: ObjectHashAggregateExec => 1 }.sum == + (if (nativeExpected) 0 else 2), + plan.treeString) + // Restoring Spark's aggregate buffers must retain the accelerated input scan. + assert(plan.collect { case _: CometNativeScanExec => 1 }.sum == 1, plan.treeString) + checkCometExchange(df, if (nativeExpected) 1 else 0, native = true) + checkSparkAnswer(df) + } + } + } + } + test("decimal hash shuffle checks nested keys recursively") { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { withNestedHashPartitioning { From 76ea3ec6753a3e68e2b09a17337d9d95fbb1bd9f Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Mon, 28 Sep 2026 21:31:19 +0000 Subject: [PATCH 3/5] Preserve one-partition decimal shuffle and cover mixed joins --- docs/source/contributor-guide/jvm_shuffle.md | 4 +- .../contributor-guide/native_shuffle.md | 5 +- .../user-guide/latest/tuning/shuffle.md | 16 ++-- .../scala/org/apache/comet/serde/hash.scala | 1 + .../shuffle/CometShuffleExchangeExec.scala | 7 +- .../CometSparkSessionExtensionsSuite.scala | 66 +------------- .../comet/exec/CometNativeShuffleSuite.scala | 89 ++++++++++++++----- .../CometCelebornShufflePlanningSuite.scala | 32 ++++--- 8 files changed, 115 insertions(+), 105 deletions(-) diff --git a/docs/source/contributor-guide/jvm_shuffle.md b/docs/source/contributor-guide/jvm_shuffle.md index 11014f34909..8632f28631c 100644 --- a/docs/source/contributor-guide/jvm_shuffle.md +++ b/docs/source/contributor-guide/jvm_shuffle.md @@ -50,7 +50,9 @@ JVM shuffle (`CometColumnarExchange`) is used instead of native shuffle (`CometE range key always falls back here. `HashPartitioning` keys must be primitive only by default: setting `spark.comet.shuffle.native.partitioning.hash.nested.enabled` to `true` keeps struct and array keys, and map keys on Spark 4.0 and later, on the native path. The config defaults to - `false`, so a complex hash key falls back to JVM columnar shuffle unless it is enabled. See + `false`, so a complex hash key falls back to JVM columnar shuffle unless it is enabled. + Decimal hash keys with precision greater than 18, including nested decimal leaves, also use + JVM shuffle when there is more than one partition and the shuffle mode allows it. See [Supported partition key types](native_shuffle.md#when-native-shuffle-is-used) for the exact rules. Complex types are fully supported as data columns in both implementations. diff --git a/docs/source/contributor-guide/native_shuffle.md b/docs/source/contributor-guide/native_shuffle.md index f852dd497a1..bb4dfa56536 100644 --- a/docs/source/contributor-guide/native_shuffle.md +++ b/docs/source/contributor-guide/native_shuffle.md @@ -71,6 +71,9 @@ Native shuffle (`CometExchange`) is selected when all of the following condition Spark's `mapsort` normalization makes physical entry order irrelevant. A collated string at any depth still disqualifies the key. The config defaults to `false` pending measurement of the nested hashing paths, so by default a complex hash key falls back to JVM shuffle. + Decimal leaves with precision greater than 18 also disqualify a hash key when there is more + than one output partition: native decimal hashing differs from Spark's. A one-partition hash + exchange stays native because the writer serializes it as `SinglePartition` without hashing. ## Architecture @@ -510,7 +513,7 @@ independently compressed, allowing parallel decompression during reads. | Input format | Columnar (direct from Comet operators) | Row-based (via ColumnarToRowExec) | | Partitioning logic | Rust implementation | Spark's partitioner | | Supported schemes | Hash, Range, Single, RoundRobin | Hash, Range, Single, RoundRobin | -| Partition key types | Primitives only (Hash, Range) | Any type | +| Partition key types | See the hash and range rules above | Any type | | Performance | Higher (no format conversion) | Lower (columnar→row→columnar) | | Writer variants | Single path | Bypass (hash) and sort-based | diff --git a/docs/source/user-guide/latest/tuning/shuffle.md b/docs/source/user-guide/latest/tuning/shuffle.md index 720fd119770..5995e8cc76d 100644 --- a/docs/source/user-guide/latest/tuning/shuffle.md +++ b/docs/source/user-guide/latest/tuning/shuffle.md @@ -53,12 +53,16 @@ scalar types. Hash partitioning keys must be scalar types unless and later) map keys. That setting is disabled by default until the performance of the nested hashing paths has been measured. Columns that are not partitioning keys may contain complex types like maps, structs, and arrays. -Hash partitioning on decimal keys with precision greater than 18 falls back because native hashing does not match -Spark's partition assignments. This can affect decimal aggregate overflow behavior, including `AVG(DISTINCT ...)`. -With `spark.comet.shuffle.mode=auto`, Comet uses Columnar Shuffle when eligible; with `native`, it uses Spark shuffle. -The restriction applies recursively to hash partitioning keys, including decimals inside structs, arrays, and maps -when nested hash partitioning is enabled. Wider decimals remain supported as payload columns, range partitioning -keys, and in single-partition shuffles. +Hash partitioning into multiple partitions on decimal keys with precision greater than 18 falls back because native +hashing does not match Spark's partition assignments. Mixing the two partitioners can silently lose join rows. +The restriction also applies to decimal leaves in nested hash keys. Wider decimals remain supported as payload +columns, range partitioning keys, and in single-partition shuffles, including one-partition hash exchanges. + +With `spark.comet.shuffle.mode=auto`, Comet uses Columnar Shuffle when eligible, adding a columnar-to-row-to-columnar +conversion. With `native`, it uses Spark shuffle, which can also restore aggregates such as `collect_list` and +`collect_set` to Spark to preserve buffer compatibility. Celeborn uses Spark/Celeborn shuffle for unsupported native +keys because it has no Comet columnar fallback. Spark-compatible native decimal hashing is tracked in +[#5994](https://github.com/apache/datafusion-comet/issues/5994). ### Columnar (JVM) Shuffle diff --git a/spark/src/main/scala/org/apache/comet/serde/hash.scala b/spark/src/main/scala/org/apache/comet/serde/hash.scala index ee3e80059d5..760bb3c2829 100644 --- a/spark/src/main/scala/org/apache/comet/serde/hash.scala +++ b/spark/src/main/scala/org/apache/comet/serde/hash.scala @@ -134,6 +134,7 @@ private object HashUtils { } private def unsupportedReasonFor(dt: DataType): Option[String] = dt match { + // Keep in sync with CometShuffleExchangeExec's hash-key restriction until #5994 is fixed. case d: DecimalType if d.precision > 18 => Some(unsupportedDecimalReason) case s: StructType => s.fields.iterator.flatMap(f => unsupportedReasonFor(f.dataType).iterator).toSeq.headOption diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index 90b01fee5a0..480de5d60a9 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -556,9 +556,10 @@ object CometShuffleExchangeExec _: TimestampNTZType | _: DateType => true case d: DecimalType => - // Spark hashes wider decimals through BigInteger bytes, which native hashing does not - // match. Different partition assignments can change decimal aggregate overflow behavior. - d.precision <= 18 + // Match the SQL hash restriction in serde/HashUtils until #5994 fixes native encoding. + // Different partition assignments break mixed native/Spark joins. A single partition + // does not hash the key: CometNativeShuffleWriter serializes it as SinglePartition. + d.precision <= 18 || s.outputPartitioning.numPartitions == 1 case dt if isTimeType(dt) => true case StructType(fields) if nestedHashPartitioningEnabled => diff --git a/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala b/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala index 3464765aba3..92bac5dba86 100644 --- a/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometSparkSessionExtensionsSuite.scala @@ -20,11 +20,11 @@ package org.apache.comet import org.apache.spark.sql._ -import org.apache.spark.sql.catalyst.expressions.{Ascending, AttributeReference, SortOrder} +import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.catalyst.plans.logical.LocalRelation -import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning, RoundRobinPartitioning, SinglePartition} -import org.apache.spark.sql.comet.{CometScanWrapper, CometSinkPlaceHolder} -import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec, CometShuffleManager} +import org.apache.spark.sql.catalyst.plans.physical.{RoundRobinPartitioning, SinglePartition} +import org.apache.spark.sql.comet.CometScanWrapper +import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometColumnarShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.LongType @@ -130,64 +130,6 @@ class CometSparkSessionExtensionsSuite extends CometTestBase { assert(after eq before, s"reload unpacked another copy of the native library: $after") } - test("wide decimal hash keys use Spark-compatible shuffle partitioning") { - // Check the native precision boundary, then routing of unsupported keys in each mode. - for (mode <- Seq("native", "auto", "jvm"); precision <- Seq(18, 19, 38)) { - withSQLConf( - CometConf.COMET_ENABLED.key -> "true", - CometConf.COMET_EXEC_ENABLED.key -> "true", - CometConf.COMET_SHUFFLE_ENABLED.key -> "true", - CometConf.COMET_SHUFFLE_MODE.key -> mode, - CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "true", - CometConf.COMET_SHUFFLE_NATIVE_RANGE_PARTITIONING_ENABLED.key -> "true", - "spark.shuffle.manager" -> classOf[CometShuffleManager].getName) { - val originalChild = spark - .range(1) - .selectExpr(s"CAST(id AS DECIMAL($precision, 0)) AS d", "id") - .queryExecution - .executedPlan - val child = CometSinkPlaceHolder( - OperatorOuterClass.Operator.getDefaultInstance, - originalChild, - originalChild) - val shuffle = ShuffleExchangeExec(HashPartitioning(Seq(child.output.head), 2), child) - val expected = if (mode == "jvm" || (mode == "auto" && precision > 18)) { - Some(CometColumnarShuffle) - } else if (precision <= 18) { - Some(CometNativeShuffle) - } else { - None - } - - withClue(s"mode=$mode, precision=$precision: ") { - assert(CometShuffleExchangeExec.shuffleSupported(shuffle) == expected) - if (expected.isEmpty) { - assert( - shuffle - .getTagValue(CometExplainInfo.FALLBACK_REASONS) - .getOrElse(Set.empty[String]) - .exists(_.contains("unsupported hash partitioning data type for native shuffle"))) - } else { - // A native failure must not tag an exchange that can use the columnar path. - assert(shuffle.getTagValue(CometExplainInfo.FALLBACK_REASONS).isEmpty) - } - - if (mode == "native" && precision > 18) { - // Wide decimals remain supported as payloads, range keys, and in a single partition. - Seq( - HashPartitioning(Seq(child.output(1)), 2), - RangePartitioning(Seq(SortOrder(child.output.head, Ascending)), 2), - SinglePartition).foreach { partitioning => - val supported = ShuffleExchangeExec(partitioning, child) - assert( - CometShuffleExchangeExec.shuffleSupported(supported).contains(CometNativeShuffle)) - } - } - } - } - } - } - test("Arrow properties") { NativeBase.setLoaded(false) NativeBase.load() diff --git a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala index 8832133a399..4c1a5eef2fd 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala @@ -43,8 +43,8 @@ import org.apache.spark.sql.catalyst.plans.logical.LocalRelation import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning import org.apache.spark.sql.comet.{CometExec, CometHashAggregateExec, CometLocalTableScanExec, CometMetricNode, CometNativeExec, CometNativeScanExec, CometScanWrapper, CometSparkToColumnarExec, CometTakeOrderedAndProjectExec} import org.apache.spark.sql.comet.execution.arrow.CometArrowStream -import org.apache.spark.sql.comet.execution.shuffle.{CometNativeShuffle, CometShuffleExchangeExec} -import org.apache.spark.sql.execution.{LocalTableScanExec, SparkPlan} +import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec} +import org.apache.spark.sql.execution.{FileSourceScanExec, LocalTableScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, AQEShuffleReadExec, ShuffleQueryStageExec} import org.apache.spark.sql.execution.aggregate.ObjectHashAggregateExec import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} @@ -508,30 +508,77 @@ class CometNativeShuffleSuite extends CometTestBase with AdaptiveSparkPlanHelper } } - test("wide decimal shuffle fallback keeps collection aggregate buffers in Spark") { + test( + "wide decimal hash shuffle keeps native aggregates unless multiple partitions need fallback") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.USE_OBJECT_HASH_AGG.key -> "true", CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true") { withParquetTable((0 until 100).map(i => (i % 3, i % 7)), "decimal_shuffle") { - for (precision <- Seq(18, 38); function <- Seq("collect_list", "collect_set")) { - val key = s"CAST(_1 AS DECIMAL($precision, 0))" - val df = - sql(s"SELECT $key, sort_array($function(_2)) FROM decimal_shuffle GROUP BY $key") - val plan = df.queryExecution.executedPlan - val nativeExpected = precision <= 18 - assert( - plan.collect { case _: CometHashAggregateExec => 1 }.sum == - (if (nativeExpected) 2 else 0), - plan.treeString) - assert( - plan.collect { case _: ObjectHashAggregateExec => 1 }.sum == - (if (nativeExpected) 0 else 2), - plan.treeString) - // Restoring Spark's aggregate buffers must retain the accelerated input scan. - assert(plan.collect { case _: CometNativeScanExec => 1 }.sum == 1, plan.treeString) - checkCometExchange(df, if (nativeExpected) 1 else 0, native = true) - checkSparkAnswer(df) + for ((precision, partitions, mode) <- Seq( + (18, 2, "native"), + (38, 2, "native"), + (38, 1, "native"), + (38, 1, "auto")); + function <- Seq("collect_list", "collect_set")) { + withSQLConf( + SQLConf.SHUFFLE_PARTITIONS.key -> partitions.toString, + CometConf.COMET_SHUFFLE_MODE.key -> mode) { + val key = s"CAST(_1 AS DECIMAL($precision, 0))" + val df = + sql(s"SELECT $key, sort_array($function(_2)) FROM decimal_shuffle GROUP BY $key") + val plan = df.queryExecution.executedPlan + val nativeExpected = precision <= 18 || partitions == 1 + assert( + plan.collect { case _: CometHashAggregateExec => 1 }.sum == + (if (nativeExpected) 2 else 0), + plan.treeString) + assert( + plan.collect { case _: ObjectHashAggregateExec => 1 }.sum == + (if (nativeExpected) 0 else 2), + plan.treeString) + // Restoring Spark's aggregate buffers must retain the accelerated input scan. + assert(plan.collect { case _: CometNativeScanExec => 1 }.sum == 1, plan.treeString) + val exchanges = checkCometExchange(df, if (nativeExpected) 1 else 0, native = true) + // GROUP BY retains HashPartitioning even when there is only one partition. + assert(exchanges.forall(_.outputPartitioning.isInstanceOf[HashPartitioning])) + checkSparkAnswer(df) + } + } + } + } + } + + test("wide decimal join keeps native and Spark inputs copartitioned") { + withSQLConf( + CometConf.COMET_SHUFFLE_MODE.key -> "auto", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + CometConf.COMET_CONVERT_FROM_JSON_ENABLED.key -> "false", + SQLConf.USE_V1_SOURCE_LIST.key -> "parquet,json", + SQLConf.SHUFFLE_PARTITIONS.key -> "7", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "false") { + withTable("decimal_parquet", "decimal_json") { + for ((table, format) <- Seq("decimal_parquet" -> "parquet", "decimal_json" -> "json")) { + sql(s"CREATE TABLE $table(id INT, k DECIMAL(38, 0)) USING $format") + sql(s"""INSERT INTO $table VALUES + |(1, 1), (2, -1), (3, 123456789012345678901234567890), + |(4, 99999999999999999999999999999999999999) + |""".stripMargin) + } + for (adaptive <- Seq(false, true)) { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString) { + val df = sql("""SELECT p.id, j.id FROM decimal_parquet p + |JOIN decimal_json j ON p.k = j.k""".stripMargin) + checkAnswer(df, Seq(Row(1, 1), Row(2, 2), Row(3, 3), Row(4, 4))) + val plan = df.queryExecution.executedPlan + val exchanges = collect(plan) { case e: CometShuffleExchangeExec => e } + assert(exchanges.size == 2, plan.treeString) + assert(exchanges.forall(_.shuffleType == CometColumnarShuffle), plan.treeString) + assert(collect(plan) { case _: CometNativeScanExec => 1 }.sum == 1, plan.treeString) + assert(collect(plan) { case _: FileSourceScanExec => 1 }.sum == 1, plan.treeString) + } } } } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala index 0e2b47e15ba..9e57d560253 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala @@ -652,17 +652,27 @@ class CometCelebornShufflePlanningSuite extends CometTestBase { } } - test(s"unsupported native repartition executes Spark fallback with AQE=$adaptive") { - withSQLConf( - SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, - CometConf.COMET_SHUFFLE_MODE.key -> "native", - CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "false") { - val nativeRegistrations = manager.nativeRegistrations.get() - val query = input.repartition(2) - assertSparkExchange(query.queryExecution.executedPlan) - checkAnswer(query, (1L to 32L).map(Row(_))) - assert(cometExchanges(query.queryExecution.executedPlan).isEmpty) - assert(manager.nativeRegistrations.get() == nativeRegistrations) + for (wideDecimal <- Seq(false, true)) { + test( + s"unsupported repartition executes Spark fallback: wideDecimal=$wideDecimal, " + + s"AQE=$adaptive") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "false") { + val nativeRegistrations = manager.nativeRegistrations.get() + val sparkRegistrations = manager.sparkRegistrations.get() + val query = if (wideDecimal) { + input.repartition(2, col("value").cast("decimal(38, 0)")) + } else { + input.repartition(2) + } + assertSparkExchange(query.queryExecution.executedPlan) + checkAnswer(query, (1L to 32L).map(Row(_))) + assert(cometExchanges(query.queryExecution.executedPlan).isEmpty) + assert(manager.nativeRegistrations.get() == nativeRegistrations) + assert(manager.sparkRegistrations.get() > sparkRegistrations) + } } } From f682d09023de062b9a5dcb55a219e4a56d75af57 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Mon, 28 Sep 2026 21:43:13 +0000 Subject: [PATCH 4/5] docs: fix native shuffle table formatting --- docs/source/contributor-guide/native_shuffle.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/source/contributor-guide/native_shuffle.md b/docs/source/contributor-guide/native_shuffle.md index bb4dfa56536..76badc4b228 100644 --- a/docs/source/contributor-guide/native_shuffle.md +++ b/docs/source/contributor-guide/native_shuffle.md @@ -513,7 +513,7 @@ independently compressed, allowing parallel decompression during reads. | Input format | Columnar (direct from Comet operators) | Row-based (via ColumnarToRowExec) | | Partitioning logic | Rust implementation | Spark's partitioner | | Supported schemes | Hash, Range, Single, RoundRobin | Hash, Range, Single, RoundRobin | -| Partition key types | See the hash and range rules above | Any type | +| Partition key types | See the hash and range rules above | Any type | | Performance | Higher (no format conversion) | Lower (columnar→row→columnar) | | Writer variants | Single path | Bypass (hash) and sort-based | From f79e2ced39fbe30930a70f8370b1e506ed13a873 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sun, 4 Oct 2026 17:58:14 +0000 Subject: [PATCH 5/5] test: avoid implicit widening in wide decimal benchmark --- .../spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala index 14429f60302..6801213425d 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometWideDecimalShuffleBenchmark.scala @@ -63,7 +63,7 @@ object CometWideDecimalShuffleBenchmark extends CometBenchmarkBase { assert(spark.table("parquetV1Table").rdd.getNumPartitions == partitions) expected = spark.sql(query).collect().toSeq.sortBy(_.getDecimal(0)) } - val benchmark = new Benchmark("wide_decimal_hash_shuffle", rows, output = output) + val benchmark = new Benchmark("wide_decimal_hash_shuffle", rows.toLong, output = output) val modes = Seq("Spark", "native", "auto", "jvm") for (mode <- (if (mainArgs.contains("--reverse")) modes.reverse else modes)) { val configs = Seq(