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..76badc4b228 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 71c9c431ab0..5995e8cc76d 100644 --- a/docs/source/user-guide/latest/tuning/shuffle.md +++ b/docs/source/user-guide/latest/tuning/shuffle.md @@ -53,6 +53,17 @@ 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 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 Comet Columnar shuffle is JVM-based and supports `HashPartitioning`, `RoundRobinPartitioning`, `RangePartitioning`, and 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 7a0d231a141..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 @@ -555,12 +555,11 @@ 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 => + // 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/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/exec/CometNativeShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala index b817ed7a503..4c1a5eef2fd 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala @@ -41,15 +41,16 @@ 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.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} -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 +441,167 @@ 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( + "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, 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) + } + } + } + } + } + + 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..6801213425d --- /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.toLong, 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() + } + } + } + } +} 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) + } } }