This is an automated email from the ASF dual-hosted git repository.
zhztheplayer pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/gluten.git
The following commit(s) were added to refs/heads/main by this push:
new 3df21cad0c [CORE][VL] Add configuration for maximum input partitions
in V2 batch scans (#12589)
3df21cad0c is described below
commit 3df21cad0c2713214e72f6a0626c9790ed61c1a9
Author: Hongze Zhang <[email protected]>
AuthorDate: Wed Jul 22 13:19:25 2026 +0100
[CORE][VL] Add configuration for maximum input partitions in V2 batch scans
(#12589)
---
.../apache/gluten/execution/VeloxScanSuite.scala | 24 +++++++
docs/Configuration.md | 1 +
.../org/apache/gluten/config/GlutenConfig.scala | 10 +++
.../execution/BatchScanExecTransformer.scala | 34 +++++++--
.../execution/GlutenWholeStageColumnarRDD.scala | 8 +++
.../gluten/execution/WholeStageTransformer.scala | 12 ++--
.../WholeStageTransformerPartitionSuite.scala | 83 ++++++++++++++++++++++
.../org/apache/gluten/integration/Suite.scala | 1 -
8 files changed, 162 insertions(+), 11 deletions(-)
diff --git
a/backends-velox/src/test/scala/org/apache/gluten/execution/VeloxScanSuite.scala
b/backends-velox/src/test/scala/org/apache/gluten/execution/VeloxScanSuite.scala
index 0ba22bf4df..4794437910 100644
---
a/backends-velox/src/test/scala/org/apache/gluten/execution/VeloxScanSuite.scala
+++
b/backends-velox/src/test/scala/org/apache/gluten/execution/VeloxScanSuite.scala
@@ -87,6 +87,30 @@ class VeloxScanSuite extends VeloxWholeStageTransformerSuite
{
}
}
+ test("coalesce v2 batch scan input partitions") {
+ withTempDir {
+ dir =>
+
spark.range(8).repartition(4).write.mode("overwrite").parquet(dir.getCanonicalPath)
+
+ withSQLConf(
+ SQLConf.USE_V1_SOURCE_LIST.key -> "",
+ SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1",
+ GlutenConfig.COLUMNAR_BATCHSCAN_MAX_INPUT_PARTITIONS.key -> "2") {
+ val df = spark.read.parquet(dir.getCanonicalPath)
+ checkAnswer(df, spark.range(8).toDF())
+
+ val scans = getExecutedPlan(df).collect { case scan:
BatchScanExecTransformer => scan }
+ assert(scans.size == 1)
+ val partitions = scans.head.getPartitions
+ assert(partitions.size == 2)
+ assert(
+ partitions
+
.map(_.asInstanceOf[SparkDataSourceRDDPartition].inputPartitions.size)
+ .sum > 2)
+ }
+ }
+ }
+
test("Test file scheme validation") {
withTempPath {
path =>
diff --git a/docs/Configuration.md b/docs/Configuration.md
index 0e4bc90b0c..83f44ec941 100644
--- a/docs/Configuration.md
+++ b/docs/Configuration.md
@@ -46,6 +46,7 @@ nav_order: 15
| spark.gluten.sql.columnar.appendData | 🔄
Dynamic | true | Enable or disable columnar v2 command append
data.
[...]
| spark.gluten.sql.columnar.arrowUdf | 🔄
Dynamic | true | Enable or disable columnar arrow udf.
[...]
| spark.gluten.sql.columnar.batchscan | 🔄
Dynamic | true | Enable or disable columnar batchscan.
[...]
+| spark.gluten.sql.columnar.batchscan.maxInputPartitions | 🔄
Dynamic | 2147483647 | Maximum number of Spark task partitions for
supported DataSource V2 batch scans.
[...]
| spark.gluten.sql.columnar.broadcastExchange | 🔄
Dynamic | true | Enable or disable columnar broadcastExchange.
[...]
| spark.gluten.sql.columnar.broadcastJoin | 🔄
Dynamic | true | Enable or disable columnar broadcastJoin.
[...]
| spark.gluten.sql.columnar.broadcastNestedLoopJoin.enabled | 🔄
Dynamic | true | Enable or disable columnar
broadcastNestedLoopJoin.
[...]
diff --git
a/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
b/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
index 93fb3888e9..4e6dcf6f8f 100644
---
a/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
+++
b/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
@@ -76,6 +76,8 @@ class GlutenConfig(conf: SQLConf) extends
GlutenCoreConfig(conf) {
def enableColumnarBatchScan: Boolean = getConf(COLUMNAR_BATCHSCAN_ENABLED)
+ def batchScanMaxInputPartitions: Int =
getConf(COLUMNAR_BATCHSCAN_MAX_INPUT_PARTITIONS)
+
def enableColumnarFileScan: Boolean = getConf(COLUMNAR_FILESCAN_ENABLED)
def enableColumnarHiveTableScan: Boolean =
getConf(COLUMNAR_HIVETABLESCAN_ENABLED)
@@ -854,6 +856,14 @@ object GlutenConfig extends ConfigRegistry {
.booleanConf
.createWithDefault(true)
+ val COLUMNAR_BATCHSCAN_MAX_INPUT_PARTITIONS =
+ buildConf("spark.gluten.sql.columnar.batchscan.maxInputPartitions")
+ .doc(
+ "Maximum number of Spark task partitions for supported DataSource V2
batch scans. ")
+ .intConf
+ .checkValue(_ > 0, s"must be positive.")
+ .createWithDefault(Int.MaxValue)
+
val COLUMNAR_FILESCAN_ENABLED =
buildConf("spark.gluten.sql.columnar.filescan")
.doc("Enable or disable columnar filescan.")
diff --git
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala
index a0c3bb8757..31fe9898e7 100644
---
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala
+++
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala
@@ -17,6 +17,7 @@
package org.apache.gluten.execution
import org.apache.gluten.backendsapi.BackendsApiManager
+import org.apache.gluten.config.GlutenConfig
import org.apache.gluten.metrics.MetricsUpdater
import org.apache.gluten.sql.shims.SparkShimLoader
import org.apache.gluten.substrait.rel.LocalFilesNode.ReadFileFormat
@@ -26,6 +27,7 @@ import org.apache.spark.Partition
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.catalyst.expressions._
import org.apache.spark.sql.catalyst.plans.QueryPlan
+import org.apache.spark.sql.catalyst.plans.physical.UnknownPartitioning
import org.apache.spark.sql.catalyst.util.truncatedString
import org.apache.spark.sql.connector.catalog.Table
import org.apache.spark.sql.connector.read.Scan
@@ -173,9 +175,9 @@ abstract class BatchScanExecTransformerBase(
override def metricsUpdater(): MetricsUpdater =
BackendsApiManager.getMetricsApiInstance.genBatchScanTransformerMetricsUpdater(metrics)
- @transient protected lazy val finalPartitions: Seq[Partition] =
- SparkShimLoader.getSparkShims
- .orderPartitions(
+ @transient protected lazy val finalPartitions: Seq[Partition] = {
+ val orderedPartitions =
+ SparkShimLoader.getSparkShims.orderPartitions(
this,
scan,
keyGroupedPartitioning,
@@ -184,11 +186,31 @@ abstract class BatchScanExecTransformerBase(
commonPartitionValues,
applyPartialClustering,
replicatePartitions)
- .zipWithIndex
- .map {
- case (inputPartitions, index) => new
SparkDataSourceRDDPartition(index, inputPartitions)
+
+ val target = GlutenConfig.get.batchScanMaxInputPartitions
+ val taskPartitions =
+ if (
+ orderedPartitions.size > target &&
+ // Coalescing changes task boundaries. Only do it when Spark does not
advertise a
+ // distribution whose partition groups must remain aligned, such as
key-grouped
+ // partitioning used by storage-partitioned joins.
+ outputPartitioning.isInstanceOf[UnknownPartitioning]
+ ) {
+ Seq.tabulate(target) {
+ index =>
+ val from = index * orderedPartitions.size / target
+ val until = (index + 1) * orderedPartitions.size / target
+ orderedPartitions.slice(from, until).flatten
+ }
+ } else {
+ orderedPartitions
}
+ taskPartitions.zipWithIndex.map {
+ case (inputPartitions, index) => new SparkDataSourceRDDPartition(index,
inputPartitions)
+ }
+ }
+
@transient override lazy val fileFormat: ReadFileFormat =
BackendsApiManager.getSettings.getSubstraitReadFileFormatV2(scan)
diff --git
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/GlutenWholeStageColumnarRDD.scala
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/GlutenWholeStageColumnarRDD.scala
index afec9cc10f..ce17823a9d 100644
---
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/GlutenWholeStageColumnarRDD.scala
+++
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/GlutenWholeStageColumnarRDD.scala
@@ -101,6 +101,14 @@ class GlutenWholeStageColumnarRDD(
}
override protected def getPartitions: Array[Partition] = {
+ rdds.getPartitionLengthOption.foreach {
+ inputPartitionCount =>
+ require(
+ inputPartitionCount == inputPartitions.size,
+ s"Whole-stage partition count ${inputPartitions.size} does not match
" +
+ s"input RDD partition count $inputPartitionCount"
+ )
+ }
inputPartitions.zipWithIndex
.map {
case (partition, i) => FirstZippedPartitionsPartition(i, partition,
rdds.getPartitions(i))
diff --git
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/WholeStageTransformer.scala
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/WholeStageTransformer.scala
index 1f3ae4d753..71cac1b5f3 100644
---
a/gluten-substrait/src/main/scala/org/apache/gluten/execution/WholeStageTransformer.scala
+++
b/gluten-substrait/src/main/scala/org/apache/gluten/execution/WholeStageTransformer.scala
@@ -489,12 +489,16 @@ class ColumnarInputRDDsWrapper(columnarInputRDDs:
Seq[RDD[ColumnarBatch]]) exten
}
def getPartitionLength: Int = {
- assert(columnarInputRDDs.nonEmpty)
- val nonBroadcastRDD =
columnarInputRDDs.find(!_.isInstanceOf[BroadcastBuildSideRDD])
- assert(nonBroadcastRDD.isDefined)
- nonBroadcastRDD.get.partitions.length
+ getPartitionLengthOption.getOrElse {
+ throw new IllegalStateException("No non-broadcast input RDD is
available")
+ }
}
+ def getPartitionLengthOption: Option[Int] =
+ columnarInputRDDs
+ .find(!_.isInstanceOf[BroadcastBuildSideRDD])
+ .map(_.partitions.length)
+
def getIterators(
inputColumnarRDDPartitions: Seq[Partition],
context: TaskContext): Seq[Iterator[ColumnarBatch]] = {
diff --git
a/gluten-substrait/src/test/scala/org/apache/gluten/execution/WholeStageTransformerPartitionSuite.scala
b/gluten-substrait/src/test/scala/org/apache/gluten/execution/WholeStageTransformerPartitionSuite.scala
new file mode 100644
index 0000000000..a7f1986375
--- /dev/null
+++
b/gluten-substrait/src/test/scala/org/apache/gluten/execution/WholeStageTransformerPartitionSuite.scala
@@ -0,0 +1,83 @@
+/*
+ * 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.gluten.execution
+
+import org.apache.gluten.metrics.IMetrics
+
+import org.apache.spark.{Partition, SparkContext, TaskContext}
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.execution.metric.SQLMetrics
+import org.apache.spark.sql.test.SharedSparkSession
+import org.apache.spark.sql.utils.SparkInputMetricsUtil.InputMetricsWrapper
+import org.apache.spark.sql.vectorized.ColumnarBatch
+
+class WholeStageTransformerPartitionSuite extends SharedSparkSession {
+ test("align whole-stage and input-wrapper partitions by index") {
+ val wholeStageRDD = createWholeStageRDD(nativePartitionCount = 3,
inputPartitionCount = 3)
+
+ val partitions =
wholeStageRDD.partitions.map(_.asInstanceOf[FirstZippedPartitionsPartition])
+ assert(partitions.map(_.index).sameElements(Array(0, 1, 2)))
+ assert(partitions.map(_.inputPartition.index).sameElements(Array(0, 1, 2)))
+ assert(
+ partitions
+ .map(_.inputColumnarRDDPartitions.map(_.index))
+ .sameElements(Array(Seq(0), Seq(1), Seq(2))))
+ }
+
+ test("fail when an input wrapper has fewer partitions than the whole stage")
{
+ val wholeStageRDD = createWholeStageRDD(nativePartitionCount = 3,
inputPartitionCount = 2)
+ val error = intercept[IllegalArgumentException](wholeStageRDD.partitions)
+ assert(error.getMessage.contains("Whole-stage partition count 3"))
+ assert(error.getMessage.contains("input RDD partition count 2"))
+ }
+
+ test("fail when an input wrapper has more partitions than the whole stage") {
+ val wholeStageRDD = createWholeStageRDD(nativePartitionCount = 2,
inputPartitionCount = 3)
+ val error = intercept[IllegalArgumentException](wholeStageRDD.partitions)
+ assert(error.getMessage.contains("Whole-stage partition count 2"))
+ assert(error.getMessage.contains("input RDD partition count 3"))
+ }
+
+ private def createWholeStageRDD(
+ nativePartitionCount: Int,
+ inputPartitionCount: Int): GlutenWholeStageColumnarRDD = {
+ val nativePartitions =
+ (0 until nativePartitionCount).map(index => GlutenPartition(index,
Array.emptyByteArray))
+ val inputRDDs =
+ new ColumnarInputRDDsWrapper(Seq(new PartitionOnlyRDD(sparkContext,
inputPartitionCount)))
+
+ new GlutenWholeStageColumnarRDD(
+ sparkContext,
+ nativePartitions,
+ inputRDDs,
+ SQLMetrics.createTimingMetric(sparkContext, "pipeline time"),
+ (_: InputMetricsWrapper) => (),
+ (_: IMetrics) => ())
+ }
+
+ private class PartitionOnlyRDD(sc: SparkContext, partitionCount: Int)
+ extends RDD[ColumnarBatch](sc, Nil) {
+
+ override protected def getPartitions: Array[Partition] =
+ Array.tabulate(partitionCount)(TestPartition)
+
+ override def compute(split: Partition, context: TaskContext):
Iterator[ColumnarBatch] =
+ throw new UnsupportedOperationException("Partition-only test RDD must
not be executed")
+ }
+
+ private case class TestPartition(index: Int) extends Partition
+}
diff --git
a/tools/gluten-it/common/src/main/scala/org/apache/gluten/integration/Suite.scala
b/tools/gluten-it/common/src/main/scala/org/apache/gluten/integration/Suite.scala
index 16aa00ae4f..f97f543c9d 100644
---
a/tools/gluten-it/common/src/main/scala/org/apache/gluten/integration/Suite.scala
+++
b/tools/gluten-it/common/src/main/scala/org/apache/gluten/integration/Suite.scala
@@ -70,7 +70,6 @@ abstract class Suite(
new SparkSessionSwitcher(appName, masterUrl, logLevel.toString)
// define initial configs
- sessionSwitcher.addDefaultConf("spark.sql.sources.useV1SourceList", "")
sessionSwitcher.addDefaultConf("spark.sql.shuffle.partitions",
s"$shufflePartitions")
sessionSwitcher.addDefaultConf("spark.storage.blockManagerSlaveTimeoutMs",
"3600000")
sessionSwitcher.addDefaultConf("spark.executor.heartbeatInterval", "10s")
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]