This is an automated email from the ASF dual-hosted git repository.
hongze pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/incubator-gluten.git
The following commit(s) were added to refs/heads/main by this push:
new 13d9a17078 [GLUTEN-8497][CORE] A unified CallInfo API to replace
AdaptiveContext (#8551)
13d9a17078 is described below
commit 13d9a17078dc9aa293c13a01e14ea52fe33bac97
Author: Hongze Zhang <[email protected]>
AuthorDate: Fri Jan 17 12:00:35 2025 +0800
[GLUTEN-8497][CORE] A unified CallInfo API to replace AdaptiveContext
(#8551)
---
.../gluten/backendsapi/clickhouse/CHRuleApi.scala | 5 +-
.../gluten/backendsapi/velox/VeloxRuleApi.scala | 8 +-
.../gluten/extension/caller/CallerInfo.scala | 71 +++++++++++++
.../extension/columnar/ColumnarRuleApplier.scala | 4 +-
.../extension/columnar/cost/GlutenCostModel.scala | 18 +++-
.../columnar/enumerated/EnumeratedApplier.scala | 19 +---
.../columnar/enumerated/EnumeratedTransform.scala | 4 +-
.../columnar/heuristic/HeuristicApplier.scala | 55 ++++------
.../columnar/heuristic/HeuristicTransform.scala | 4 +-
.../gluten/extension/injector/GlutenInjector.scala | 5 +-
.../gluten/extension/util/AdaptiveContext.scala | 86 ----------------
.../sql/execution/FallbackStrategiesSuite.scala | 113 +++++++++++----------
.../sql/execution/FallbackStrategiesSuite.scala | 113 +++++++++++----------
.../sql/execution/FallbackStrategiesSuite.scala | 113 +++++++++++----------
.../sql/execution/FallbackStrategiesSuite.scala | 113 +++++++++++----------
15 files changed, 355 insertions(+), 376 deletions(-)
diff --git
a/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
b/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
index 21ae342a22..c79931fa4e 100644
---
a/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
+++
b/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
@@ -120,11 +120,10 @@ object CHRuleApi {
injector.injectPostTransform(c =>
AddPreProjectionForHashJoin.apply(c.session))
// Gluten columnar: Fallback policies.
- injector.injectFallbackPolicy(
- c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), c.ac.originalPlan()))
+ injector.injectFallbackPolicy(c => p =>
ExpandFallbackPolicy(c.caller.isAqe(), p))
// Gluten columnar: Post rules.
- injector.injectPost(c => RemoveTopmostColumnarToRow(c.session,
c.ac.isAdaptiveContext()))
+ injector.injectPost(c => RemoveTopmostColumnarToRow(c.session,
c.caller.isAqe()))
SparkShimLoader.getSparkShims
.getExtendedColumnarPostRules()
.foreach(each => injector.injectPost(c => intercept(each(c.session))))
diff --git
a/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
b/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
index 0cf6ac6713..9825ae1d31 100644
---
a/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
+++
b/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
@@ -101,11 +101,10 @@ object VeloxRuleApi {
injector.injectPostTransform(c =>
InsertTransitions.create(c.outputsColumnar, VeloxBatch))
// Gluten columnar: Fallback policies.
- injector.injectFallbackPolicy(
- c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), c.ac.originalPlan()))
+ injector.injectFallbackPolicy(c => p =>
ExpandFallbackPolicy(c.caller.isAqe(), p))
// Gluten columnar: Post rules.
- injector.injectPost(c => RemoveTopmostColumnarToRow(c.session,
c.ac.isAdaptiveContext()))
+ injector.injectPost(c => RemoveTopmostColumnarToRow(c.session,
c.caller.isAqe()))
SparkShimLoader.getSparkShims
.getExtendedColumnarPostRules()
.foreach(each => injector.injectPost(c => each(c.session)))
@@ -180,8 +179,7 @@ object VeloxRuleApi {
injector.injectPostTransform(_ => CollapseProjectExecTransformer)
injector.injectPostTransform(c =>
FlushableHashAggregateRule.apply(c.session))
injector.injectPostTransform(c =>
InsertTransitions.create(c.outputsColumnar, VeloxBatch))
- injector.injectPostTransform(
- c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()))
+ injector.injectPostTransform(c => RemoveTopmostColumnarToRow(c.session,
c.caller.isAqe()))
SparkShimLoader.getSparkShims
.getExtendedColumnarPostRules()
.foreach(each => injector.injectPostTransform(c => each(c.session)))
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
new file mode 100644
index 0000000000..d23d172e46
--- /dev/null
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
@@ -0,0 +1,71 @@
+/*
+ * 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.extension.caller
+
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanExec
+import org.apache.spark.sql.execution.columnar.InMemoryRelation
+
+/**
+ * Helper API that stores information about the call site of the columnar
rule. Specific columnar
+ * rules could call the API to check whether this time of rule call was
initiated for certain
+ * purpose. For example, a rule call could be for AQE optimization, or for
cached plan optimization,
+ * or for regular executed plan optimization.
+ */
+trait CallerInfo {
+ def isAqe(): Boolean
+ def isCache(): Boolean
+}
+
+object CallerInfo {
+ private val localStorage: ThreadLocal[Option[CallerInfo]] =
+ new ThreadLocal[Option[CallerInfo]]() {
+ override def initialValue(): Option[CallerInfo] = None
+ }
+
+ private class Impl(override val isAqe: Boolean, override val isCache:
Boolean) extends CallerInfo
+
+ /*
+ * Find the information about the caller that initiated the rule call.
+ */
+ def create(): CallerInfo = {
+ if (localStorage.get().nonEmpty) {
+ return localStorage.get().get
+ }
+ val stack = Thread.currentThread.getStackTrace
+ new Impl(isAqe = inAqeCall(stack), isCache = inCacheCall(stack))
+ }
+
+ private def inAqeCall(stack: Seq[StackTraceElement]): Boolean = {
+ stack.exists(_.getClassName.equals(AdaptiveSparkPlanExec.getClass.getName))
+ }
+
+ private def inCacheCall(stack: Seq[StackTraceElement]): Boolean = {
+ stack.exists(_.getClassName.equals(InMemoryRelation.getClass.getName))
+ }
+
+ /** For testing only. */
+ def withLocalValue[T](isAqe: Boolean, isCache: Boolean)(body: => T): T = {
+ val prevValue = localStorage.get()
+ val newValue = new Impl(isAqe, isCache)
+ localStorage.set(Some(newValue))
+ try {
+ body
+ } finally {
+ localStorage.set(prevValue)
+ }
+ }
+}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
index 3e9bd72c45..7257865507 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
@@ -17,7 +17,7 @@
package org.apache.gluten.extension.columnar
import org.apache.gluten.config.GlutenConfig
-import org.apache.gluten.extension.util.AdaptiveContext
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.execution.SparkPlan
@@ -29,7 +29,7 @@ trait ColumnarRuleApplier {
object ColumnarRuleApplier {
class ColumnarRuleCall(
val session: SparkSession,
- val ac: AdaptiveContext,
+ val caller: CallerInfo,
val outputsColumnar: Boolean) {
val glutenConf: GlutenConfig = {
new GlutenConfig(session.sessionState.conf)
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
index 80edf8919f..bfaf7380c6 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
@@ -21,11 +21,15 @@ import org.apache.gluten.component.Component
import org.apache.spark.internal.Logging
import org.apache.spark.sql.execution.SparkPlan
import org.apache.spark.util.SparkReflectionUtil
-
+// format: off
/**
* The cost model API of Gluten. Used by:
- * 1. RAS planner for cost-based optimization; 2. Transition graph for
choosing transition paths.
+ * <p>
+ * 1. RAS planner for cost-based optimization;
+ * <p>
+ * 2. Transition graph for choosing transition paths.
*/
+// format: on
trait GlutenCostModel {
def costOf(node: SparkPlan): GlutenCost
def costComparator(): Ordering[GlutenCost]
@@ -38,14 +42,18 @@ trait GlutenCostModel {
}
object GlutenCostModel extends Logging {
- def find(aliasOrClass: String): GlutenCostModel = {
- val costModelRegistry = LongCostModel.registry()
+ private val costModelRegistry = {
+ val r = LongCostModel.registry()
// Components should override Backend's costers. Hence, reversed
registration order is applied.
Component
.sorted()
.reverse
.flatMap(_.costers())
- .foreach(coster => costModelRegistry.register(coster))
+ .foreach(coster => r.register(coster))
+ r
+ }
+
+ def find(aliasOrClass: String): GlutenCostModel = {
val costModel = find(costModelRegistry, aliasOrClass)
costModel
}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
index 04cd70656e..2f5c8c4472 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
@@ -16,9 +16,9 @@
*/
package org.apache.gluten.extension.columnar.enumerated
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.gluten.extension.columnar.{ColumnarRuleApplier,
ColumnarRuleExecutor}
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
-import org.apache.gluten.extension.util.AdaptiveContext
import org.apache.gluten.logging.LogLevelUtil
import org.apache.spark.internal.Logging
@@ -39,27 +39,14 @@ class EnumeratedApplier(
extends ColumnarRuleApplier
with Logging
with LogLevelUtil {
- private val adaptiveContext = AdaptiveContext(session)
-
override def apply(plan: SparkPlan, outputsColumnar: Boolean): SparkPlan = {
- val call = new ColumnarRuleCall(session, adaptiveContext, outputsColumnar)
- val finalPlan = maybeAqe {
- apply0(ruleBuilders.map(b => b(call)), plan)
- }
+ val call = new ColumnarRuleCall(session, CallerInfo.create(),
outputsColumnar)
+ val finalPlan = apply0(ruleBuilders.map(b => b(call)), plan)
finalPlan
}
private def apply0(rules: Seq[Rule[SparkPlan]], plan: SparkPlan): SparkPlan =
new ColumnarRuleExecutor("ras", rules).execute(plan)
-
- private def maybeAqe[T](f: => T): T = {
- adaptiveContext.setAdaptiveContext()
- try {
- f
- } finally {
- adaptiveContext.resetAdaptiveContext()
- }
- }
}
object EnumeratedApplier {}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
index 72926407c5..6af9b00134 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
@@ -18,12 +18,12 @@ package org.apache.gluten.extension.columnar.enumerated
import org.apache.gluten.component.Component
import org.apache.gluten.exception.GlutenException
+import org.apache.gluten.extension.caller.CallerInfo
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import org.apache.gluten.extension.columnar.cost.{GlutenCost, GlutenCostModel}
import
org.apache.gluten.extension.columnar.enumerated.planner.GlutenOptimization
import org.apache.gluten.extension.columnar.enumerated.planner.property.Conv
import org.apache.gluten.extension.injector.Injector
-import org.apache.gluten.extension.util.AdaptiveContext
import org.apache.gluten.logging.LogLevelUtil
import org.apache.gluten.ras.{Cost, CostModel}
import org.apache.gluten.ras.property.PropertySet
@@ -81,7 +81,7 @@ object EnumeratedTransform {
val session = SparkSession.getActiveSession.getOrElse(
throw new GlutenException(
"HeuristicTransform#static can only be called when an active Spark
session exists"))
- val call = new ColumnarRuleCall(session, AdaptiveContext(session), false)
+ val call = new ColumnarRuleCall(session, CallerInfo.create(), false)
dummyInjector.gluten.ras.createEnumeratedTransform(call)
}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
index e4825d2eb7..493319d423 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
@@ -16,9 +16,9 @@
*/
package org.apache.gluten.extension.columnar.heuristic
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.gluten.extension.columnar.{ColumnarRuleApplier,
ColumnarRuleExecutor}
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
-import org.apache.gluten.extension.util.AdaptiveContext
import org.apache.gluten.logging.LogLevelUtil
import org.apache.spark.internal.Logging
@@ -33,35 +33,33 @@ import org.apache.spark.sql.execution.SparkPlan
class HeuristicApplier(
session: SparkSession,
transformBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
- fallbackPolicyBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
+ fallbackPolicyBuilders: Seq[ColumnarRuleCall => SparkPlan =>
Rule[SparkPlan]],
postBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
finalBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]])
extends ColumnarRuleApplier
with Logging
with LogLevelUtil {
- private val adaptiveContext = AdaptiveContext(session)
-
override def apply(plan: SparkPlan, outputsColumnar: Boolean): SparkPlan = {
- val call = new ColumnarRuleCall(session, adaptiveContext, outputsColumnar)
+ val call = new ColumnarRuleCall(session, CallerInfo.create(),
outputsColumnar)
makeRule(call).apply(plan)
}
private def makeRule(call: ColumnarRuleCall): Rule[SparkPlan] = {
- plan =>
- prepareFallback(plan) {
- p =>
- val suggestedPlan = transformPlan("transform", transformRules(call),
p)
- val finalPlan = transformPlan("fallback", fallbackPolicies(call),
suggestedPlan) match {
- case FallbackNode(fallbackPlan) =>
- // we should use vanilla c2r rather than native c2r,
- // and there should be no `GlutenPlan` any more,
- // so skip the `postRules()`.
- fallbackPlan
- case plan =>
- transformPlan("post", postRules(call), plan)
- }
- transformPlan("final", finalRules(call), finalPlan)
+ originalPlan =>
+ val suggestedPlan = transformPlan("transform", transformRules(call),
originalPlan)
+ val finalPlan = transformPlan(
+ "fallback",
+ fallbackPolicies(call).map(_(originalPlan)),
+ suggestedPlan) match {
+ case FallbackNode(fallbackPlan) =>
+ // we should use vanilla c2r rather than native c2r,
+ // and there should be no `GlutenPlan` anymore,
+ // so skip the `postRules()`.
+ fallbackPlan
+ case plan =>
+ transformPlan("post", postRules(call), plan)
}
+ transformPlan("final", finalRules(call), finalPlan)
}
private def transformPlan(
@@ -70,17 +68,6 @@ class HeuristicApplier(
plan: SparkPlan): SparkPlan =
new ColumnarRuleExecutor(phase, rules).execute(plan)
- private def prepareFallback[T](p: SparkPlan)(f: SparkPlan => T): T = {
- adaptiveContext.setAdaptiveContext()
- adaptiveContext.setOriginalPlan(p)
- try {
- f(p)
- } finally {
- adaptiveContext.resetOriginalPlan()
- adaptiveContext.resetAdaptiveContext()
- }
- }
-
/**
* Rules to let planner create a suggested Gluten plan being sent to
`fallbackPolicies` in which
* the plan will be breakdown and decided to be fallen back or not.
@@ -93,7 +80,7 @@ class HeuristicApplier(
* Rules to add wrapper `FallbackNode`s on top of the input plan, as hints
to make planner fall
* back the whole input plan to the original vanilla Spark plan.
*/
- private def fallbackPolicies(call: ColumnarRuleCall): Seq[Rule[SparkPlan]] =
{
+ private def fallbackPolicies(call: ColumnarRuleCall): Seq[SparkPlan =>
Rule[SparkPlan]] = {
fallbackPolicyBuilders.map(b => b.apply(call))
}
@@ -112,12 +99,6 @@ class HeuristicApplier(
private def finalRules(call: ColumnarRuleCall): Seq[Rule[SparkPlan]] = {
finalBuilders.map(b => b.apply(call))
}
-
- // Just for test use.
- def enableAdaptiveContext(): HeuristicApplier = {
- adaptiveContext.enableAdaptiveContext()
- this
- }
}
object HeuristicApplier {}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
index 5a0fdfeefe..011453032b 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
@@ -18,12 +18,12 @@ package org.apache.gluten.extension.columnar.heuristic
import org.apache.gluten.component.Component
import org.apache.gluten.exception.GlutenException
+import org.apache.gluten.extension.caller.CallerInfo
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import org.apache.gluten.extension.columnar.offload.OffloadSingleNode
import org.apache.gluten.extension.columnar.rewrite.RewriteSingleNode
import org.apache.gluten.extension.columnar.validator.Validator
import org.apache.gluten.extension.injector.Injector
-import org.apache.gluten.extension.util.AdaptiveContext
import org.apache.gluten.logging.LogLevelUtil
import org.apache.spark.internal.Logging
@@ -126,7 +126,7 @@ object HeuristicTransform {
val session = SparkSession.getActiveSession.getOrElse(
throw new GlutenException(
"HeuristicTransform#static can only be called when an active Spark
session exists"))
- val call = new ColumnarRuleCall(session, AdaptiveContext(session), false)
+ val call = new ColumnarRuleCall(session, CallerInfo.create(), false)
dummyInjector.gluten.legacy.createHeuristicTransform(call)
}
}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
index a208db2c96..fa8704509e 100644
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
+++
b/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
@@ -57,7 +57,8 @@ object GlutenInjector {
private val preTransformBuilders = mutable.Buffer.empty[ColumnarRuleCall
=> Rule[SparkPlan]]
private val transformBuilders = mutable.Buffer.empty[ColumnarRuleCall =>
Rule[SparkPlan]]
private val postTransformBuilders = mutable.Buffer.empty[ColumnarRuleCall
=> Rule[SparkPlan]]
- private val fallbackPolicyBuilders = mutable.Buffer.empty[ColumnarRuleCall
=> Rule[SparkPlan]]
+ private val fallbackPolicyBuilders =
+ mutable.Buffer.empty[ColumnarRuleCall => SparkPlan => Rule[SparkPlan]]
private val postBuilders = mutable.Buffer.empty[ColumnarRuleCall =>
Rule[SparkPlan]]
private val finalBuilders = mutable.Buffer.empty[ColumnarRuleCall =>
Rule[SparkPlan]]
@@ -73,7 +74,7 @@ object GlutenInjector {
postTransformBuilders += builder
}
- def injectFallbackPolicy(builder: ColumnarRuleCall => Rule[SparkPlan]):
Unit = {
+ def injectFallbackPolicy(builder: ColumnarRuleCall => SparkPlan =>
Rule[SparkPlan]): Unit = {
fallbackPolicyBuilders += builder
}
diff --git
a/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
b/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
deleted file mode 100644
index b0f42e7967..0000000000
---
a/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
+++ /dev/null
@@ -1,86 +0,0 @@
-/*
- * 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.extension.util
-
-import org.apache.spark.sql.SparkSession
-import org.apache.spark.sql.execution.SparkPlan
-import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanExec
-
-import scala.collection.mutable.ListBuffer
-
-// Since: https://github.com/apache/incubator-gluten/pull/3294.
-sealed trait AdaptiveContext {
- def enableAdaptiveContext(): Unit
- def isAdaptiveContext(): Boolean
- def setAdaptiveContext(): Unit
- def resetAdaptiveContext(): Unit
- def setOriginalPlan(plan: SparkPlan): Unit
- def originalPlan(): SparkPlan
- def resetOriginalPlan(): Unit
-}
-
-object AdaptiveContext {
- def apply(session: SparkSession): AdaptiveContext =
- new AdaptiveContextImpl(session)
-
- private val GLUTEN_IS_ADAPTIVE_CONTEXT = "gluten.isAdaptiveContext"
-
- // Holds the original plan for possible entire fallback.
- private val localOriginalPlans: ThreadLocal[ListBuffer[SparkPlan]] =
- ThreadLocal.withInitial(() => ListBuffer.empty[SparkPlan])
- private val localIsAdaptiveContextFlags: ThreadLocal[ListBuffer[Boolean]] =
- ThreadLocal.withInitial(() => ListBuffer.empty[Boolean])
-
- private class AdaptiveContextImpl(session: SparkSession) extends
AdaptiveContext {
- // Just for test use.
- override def enableAdaptiveContext(): Unit = {
- session.sparkContext.setLocalProperty(GLUTEN_IS_ADAPTIVE_CONTEXT, "true")
- }
-
- override def isAdaptiveContext(): Boolean =
- Option(session.sparkContext.getLocalProperty(GLUTEN_IS_ADAPTIVE_CONTEXT))
- .getOrElse("false")
- .toBoolean ||
- localIsAdaptiveContextFlags.get().head
-
- override def setAdaptiveContext(): Unit = {
- val traceElements = Thread.currentThread.getStackTrace
- // ApplyColumnarRulesAndInsertTransitions is called by either
QueryExecution or
- // AdaptiveSparkPlanExec. So by checking the stack trace, we can know
whether
- // columnar rule will be applied in adaptive execution context.
- localIsAdaptiveContextFlags
- .get()
- .prepend(
-
traceElements.exists(_.getClassName.equals(AdaptiveSparkPlanExec.getClass.getName)))
- }
-
- override def resetAdaptiveContext(): Unit =
- localIsAdaptiveContextFlags.get().remove(0)
-
- override def setOriginalPlan(plan: SparkPlan): Unit = {
- localOriginalPlans.get().prepend(plan)
- }
-
- override def originalPlan(): SparkPlan = {
- val plan = localOriginalPlans.get().head
- assert(plan != null)
- plan
- }
-
- override def resetOriginalPlan(): Unit = localOriginalPlans.get().remove(0)
- }
-}
diff --git
a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 97a779289d..bbbc913bd8 100644
---
a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++
b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
import org.apache.gluten.config.GlutenConfig
import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
import org.apache.gluten.extension.columnar.RemoveFallbackTagRule
@@ -54,37 +55,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
testGluten("Fall back the whole plan if meeting the configured threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"1")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
testGluten("Don't fall back the whole plan if NOT meeting the configured
threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"4")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -92,19 +95,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
"Fall back the whole plan if meeting the configured threshold (leaf node
is" +
" transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"2")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
@@ -112,19 +116,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait
{
"Don't Fall back the whole plan if NOT meeting the configured threshold ("
+
"leaf node is transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"3")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -169,9 +174,9 @@ private object FallbackStrategiesSuite {
new HeuristicApplier(
spark,
transformBuilders,
- List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(),
c.ac.originalPlan())),
+ List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
List(
- c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+ c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
_ => ColumnarCollapseTransformStages(GlutenConfig.get)
),
List(_ => RemoveFallbackTagRule())
diff --git
a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 8b68c36bb3..f51f2721f4 100644
---
a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++
b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
import org.apache.gluten.config.GlutenConfig
import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.gluten.extension.columnar.{FallbackTags,
RemoveFallbackTagRule}
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -53,37 +54,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
testGluten("Fall back the whole plan if meeting the configured threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"1")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
testGluten("Don't fall back the whole plan if NOT meeting the configured
threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"4")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -91,19 +94,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
"Fall back the whole plan if meeting the configured threshold (leaf node
is" +
" transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"2")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
@@ -111,19 +115,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait
{
"Don't Fall back the whole plan if NOT meeting the configured threshold ("
+
"leaf node is transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"3")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -179,9 +184,9 @@ private object FallbackStrategiesSuite {
new HeuristicApplier(
spark,
transformBuilders,
- List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(),
c.ac.originalPlan())),
+ List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
List(
- c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+ c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
_ => ColumnarCollapseTransformStages(GlutenConfig.get)
),
List(_ => RemoveFallbackTagRule())
diff --git
a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 8b68c36bb3..f51f2721f4 100644
---
a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++
b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
import org.apache.gluten.config.GlutenConfig
import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.gluten.extension.columnar.{FallbackTags,
RemoveFallbackTagRule}
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -53,37 +54,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
testGluten("Fall back the whole plan if meeting the configured threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"1")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
testGluten("Don't fall back the whole plan if NOT meeting the configured
threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"4")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -91,19 +94,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
"Fall back the whole plan if meeting the configured threshold (leaf node
is" +
" transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"2")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
@@ -111,19 +115,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait
{
"Don't Fall back the whole plan if NOT meeting the configured threshold ("
+
"leaf node is transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"3")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -179,9 +184,9 @@ private object FallbackStrategiesSuite {
new HeuristicApplier(
spark,
transformBuilders,
- List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(),
c.ac.originalPlan())),
+ List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
List(
- c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+ c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
_ => ColumnarCollapseTransformStages(GlutenConfig.get)
),
List(_ => RemoveFallbackTagRule())
diff --git
a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 898967d4f3..1dd7eccc21 100644
---
a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++
b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
import org.apache.gluten.config.GlutenConfig
import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
import org.apache.gluten.extension.columnar.{FallbackTags,
RemoveFallbackTagRule}
import
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
import
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -54,37 +55,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
testGluten("Fall back the whole plan if meeting the configured threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"1")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
testGluten("Don't fall back the whole plan if NOT meeting the configured
threshold") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"4")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -92,19 +95,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
"Fall back the whole plan if meeting the configured threshold (leaf node
is" +
" transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"2")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to fall back the entire plan.
- assert(outputPlan == originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to fall back the entire plan.
+ assert(outputPlan == originalPlan)
+ }
}
}
@@ -112,19 +116,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait
{
"Don't Fall back the whole plan if NOT meeting the configured threshold ("
+
"leaf node is transformable)") {
withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold",
"3")) {
- val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
- val rule = newRuleApplier(
- spark,
- List(
- _ =>
- _ => {
-
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
- },
- c => InsertBackendTransitions(c.outputsColumnar)))
- .enableAdaptiveContext()
- val outputPlan = rule.apply(originalPlan, false)
- // Expect to get the plan with columnar rule applied.
- assert(outputPlan != originalPlan)
+ CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+ val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+ val rule = newRuleApplier(
+ spark,
+ List(
+ _ =>
+ _ => {
+
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+ },
+ c => InsertBackendTransitions(c.outputsColumnar)))
+ val outputPlan = rule.apply(originalPlan, false)
+ // Expect to get the plan with columnar rule applied.
+ assert(outputPlan != originalPlan)
+ }
}
}
@@ -180,9 +185,9 @@ private object FallbackStrategiesSuite {
new HeuristicApplier(
spark,
transformBuilders,
- List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(),
c.ac.originalPlan())),
+ List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
List(
- c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+ c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
_ => ColumnarCollapseTransformStages(GlutenConfig.get)
),
List(_ => RemoveFallbackTagRule())
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]