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]

Reply via email to