diff --git a/backends-velox/src/main/scala/org/apache/gluten/extension/PartialFallback.scala b/backends-velox/src/main/scala/org/apache/gluten/extension/PartialFallback.scala index 8a7da72b3e9..4f15ad31e71 100644 --- a/backends-velox/src/main/scala/org/apache/gluten/extension/PartialFallback.scala +++ b/backends-velox/src/main/scala/org/apache/gluten/extension/PartialFallback.scala @@ -16,6 +16,8 @@ */ package org.apache.gluten.extension +import org.apache.gluten.config.GlutenConfig + import org.apache.spark.sql.catalyst.rules.{Rule, RuleExecutor} import org.apache.spark.sql.execution.{GenerateExec, ProjectExec, SparkPlan} import org.apache.spark.sql.internal.SQLConf @@ -37,8 +39,9 @@ case class PartialFallbackRules() extends Rule[SparkPlan] { } object PartialFallback { - def supportPartialFallback(plan: SparkPlan): Boolean = { - plan.isInstanceOf[ProjectExec] || - plan.isInstanceOf[GenerateExec] + def supportPartialFallback(plan: SparkPlan): Boolean = plan match { + case _: ProjectExec => GlutenConfig.get.enableColumnarPartialProject + case _: GenerateExec => GlutenConfig.get.enableColumnarPartialGenerate + case _ => false } } diff --git a/backends-velox/src/test/scala/org/apache/spark/sql/execution/GlutenHiveUDFSuite.scala b/backends-velox/src/test/scala/org/apache/spark/sql/execution/GlutenHiveUDFSuite.scala index 296de381cd6..914300069f4 100644 --- a/backends-velox/src/test/scala/org/apache/spark/sql/execution/GlutenHiveUDFSuite.scala +++ b/backends-velox/src/test/scala/org/apache/spark/sql/execution/GlutenHiveUDFSuite.scala @@ -162,15 +162,34 @@ class GlutenHiveUDFSuite extends GlutenQueryComparisonTest with SQLTestUtils { val plusOne = udf((x: Long) => x + 1) spark.udf.register("plus_one", plusOne) sql(s"CREATE TEMPORARY FUNCTION noInputUDTF AS '${classOf[NoInputUDTF].getName}'") - runQueryAndCompare(""" - |select plus_one(col1) as col2, l_partkey from ( - | select col1, l_partkey from lineitem lateral view noInputUDTF() as col1 - |)""".stripMargin) { - df => - { - checkOperatorMatch[ColumnarPartialProjectExec](df) - checkOperatorMatch[ColumnarPartialGenerateExec](df) + + for { + enablePartialProject <- Seq(true, false) + enablePartialGenerate <- Seq(true, false) + } { + val expectPartialFallback = enablePartialProject && enablePartialGenerate + withSQLConf( + SQLConf.ANSI_ENABLED.key -> "false", + GlutenConfig.ENABLE_COLUMNAR_PARTIAL_PROJECT.key -> enablePartialProject.toString, + GlutenConfig.ENABLE_COLUMNAR_PARTIAL_GENERATE.key -> enablePartialGenerate.toString + ) { + runQueryAndCompare( + """ + |select plus_one(col1) as col2, l_partkey from ( + | select col1, l_partkey from lineitem lateral view noInputUDTF() as col1 + |)""".stripMargin, + noFallBack = expectPartialFallback + ) { + df => + val executedPlan = getExecutedPlan(df) + assert( + executedPlan.exists(_.isInstanceOf[ColumnarPartialProjectExec]) == + expectPartialFallback) + assert( + executedPlan.exists(_.isInstanceOf[ColumnarPartialGenerateExec]) == + expectPartialFallback) } + } } } }