Skip to content

Commit b05887e

Browse files
committed
[SPARK-60135][SQL] Prune unnecessary subquery-expression traversal
Generated-by: Pi Coding Agent 1.1.0
1 parent 215a146 commit b05887e

3 files changed

Lines changed: 214 additions & 5 deletions

File tree

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/DeduplicateRelations.scala‎

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -246,11 +246,16 @@ object DeduplicateRelations extends Rule[LogicalPlan] {
246246
}
247247
}
248248

249-
val planWithNewSubquery = plan.transformExpressions {
250-
case subquery: SubqueryExpression =>
251-
val (renewed, changed) = renewDuplicatedRelations(existingRelations, subquery.plan)
252-
if (changed) planChanged = true
253-
subquery.withNewPlan(renewed)
249+
val planWithNewSubquery = if (plan.containsPattern(PLAN_EXPRESSION)) {
250+
// Do not cache ineffective transformations: existingRelations changes during traversal.
251+
plan.transformExpressionsWithPruning(_.containsPattern(PLAN_EXPRESSION)) {
252+
case subquery: SubqueryExpression =>
253+
val (renewed, changed) = renewDuplicatedRelations(existingRelations, subquery.plan)
254+
if (changed) planChanged = true
255+
subquery.withNewPlan(renewed)
256+
}
257+
} else {
258+
plan
254259
}
255260

256261
if (planChanged) {
Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,67 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one or more
3+
* contributor license agreements. See the NOTICE file distributed with
4+
* this work for additional information regarding copyright ownership.
5+
* The ASF licenses this file to You under the Apache License, Version 2.0
6+
* (the "License"); you may not use this file except in compliance with
7+
* the License. You may obtain a copy of the License at
8+
*
9+
* http://www.apache.org/licenses/LICENSE-2.0
10+
*
11+
* Unless required by applicable law or agreed to in writing, software
12+
* distributed under the License is distributed on an "AS IS" BASIS,
13+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14+
* See the License for the specific language governing permissions and
15+
* limitations under the License.
16+
*/
17+
18+
package org.apache.spark.sql.catalyst.analysis
19+
20+
import org.apache.spark.benchmark.{Benchmark, BenchmarkBase}
21+
import org.apache.spark.sql.catalyst.dsl.expressions._
22+
import org.apache.spark.sql.catalyst.expressions.{Add, Alias, Expression, Literal, ScalarSubquery}
23+
import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, Project}
24+
25+
/**
26+
* Measures relation deduplication on expression-heavy projections, with and without subqueries.
27+
* Plans are reused across rule invocations, so tree pattern bits are warm after the first run.
28+
* To run this benchmark:
29+
* {{{
30+
* build/sbt "catalyst/Test/runMain <this class>"
31+
* }}}
32+
*/
33+
object DeduplicateRelationsBenchmark extends BenchmarkBase {
34+
override def runBenchmarkSuite(mainArgs: Array[String]): Unit = {
35+
val iterations = 100
36+
for (width <- Seq(100, 1000)) {
37+
runBenchmark(s"DeduplicateRelations: $width columns") {
38+
val benchmark = new Benchmark(
39+
s"DeduplicateRelations: $width columns", iterations, output = output)
40+
for (withSubquery <- Seq(false, true)) {
41+
val relation = LocalRelation($"a".int)
42+
val columns = (0 until width).map { i =>
43+
val initial: Expression = if (withSubquery && i == 0) {
44+
ScalarSubquery(relation.newInstance())
45+
} else {
46+
relation.output.head
47+
}
48+
val expression = (0 until 10).foldLeft(initial) { (child, _) =>
49+
Add(child, Literal(1))
50+
}
51+
Alias(expression, s"c$i")()
52+
}
53+
val plan = Project(columns, relation)
54+
val name = if (withSubquery) "One scalar subquery" else "No subqueries"
55+
benchmark.addCase(name) { _ =>
56+
var i = 0
57+
while (i < iterations) {
58+
DeduplicateRelations(plan)
59+
i += 1
60+
}
61+
}
62+
}
63+
benchmark.run()
64+
}
65+
}
66+
}
67+
}
Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one or more
3+
* contributor license agreements. See the NOTICE file distributed with
4+
* this work for additional information regarding copyright ownership.
5+
* The ASF licenses this file to You under the Apache License, Version 2.0
6+
* (the "License"); you may not use this file except in compliance with
7+
* the License. You may obtain a copy of the License at
8+
*
9+
* http://www.apache.org/licenses/LICENSE-2.0
10+
*
11+
* Unless required by applicable law or agreed to in writing, software
12+
* distributed under the License is distributed on an "AS IS" BASIS,
13+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14+
* See the License for the specific language governing permissions and
15+
* limitations under the License.
16+
*/
17+
18+
package org.apache.spark.sql.catalyst.analysis
19+
20+
import org.apache.spark.sql.catalyst.dsl.expressions._
21+
import org.apache.spark.sql.catalyst.expressions._
22+
import org.apache.spark.sql.catalyst.plans.{Inner, PlanTest}
23+
import org.apache.spark.sql.catalyst.plans.logical._
24+
import org.apache.spark.sql.types.DataType
25+
26+
class DeduplicateRelationsSuite extends PlanTest {
27+
import DeduplicateRelationsSuite._
28+
29+
test("SPARK-60135: skip expression mapping for plans without subqueries") {
30+
val relation = LocalRelation($"a".int)
31+
val plan = new Filter(EqualTo(relation.output.head, Literal(1)), relation) {
32+
override def mapExpressions(f: Expression => Expression): this.type = {
33+
fail("Expressions should not be mapped when there are no subqueries")
34+
}
35+
}
36+
37+
assert(DeduplicateRelations(plan) eq plan)
38+
}
39+
40+
test("SPARK-60135: prune expression branches without subqueries") {
41+
val relation = LocalRelation($"a".int)
42+
val unrelated = TraversalCountingExpression(Add(Literal(1), Literal(2)))
43+
val subquery = ScalarSubquery(relation)
44+
val alias = Alias(Add(unrelated, subquery), "s")()
45+
val plan = Project(Seq(alias), relation)
46+
47+
val result = DeduplicateRelations(plan).asInstanceOf[Project]
48+
assert(unrelated.traversals == 0)
49+
assert(result.child eq relation)
50+
assert(result.output.head.exprId == alias.exprId)
51+
val renewed = result.projectList.head.collect { case s: ScalarSubquery => s }.head
52+
assert(renewed.exprId == subquery.exprId)
53+
assert(renewed.plan.outputSet.intersect(relation.outputSet).isEmpty)
54+
}
55+
56+
test("SPARK-60135: renew children even when there are no subqueries") {
57+
val relation = LocalRelation($"a".int)
58+
val attr = relation.output.head
59+
val project = Project(Seq(Alias(Add(attr, Literal(1)), "b")()),
60+
Filter(EqualTo(attr, Literal(1)), relation))
61+
val plan = Join(project, project, Inner, None, JoinHint.NONE)
62+
63+
val result = DeduplicateRelations(plan).asInstanceOf[Join]
64+
assert(result.left eq project)
65+
assert(result.duplicateResolved)
66+
assert(result.right.collect { case r: LocalRelation => r }.head.outputSet
67+
.intersect(relation.outputSet).isEmpty)
68+
assert(!result.right.exists(_.missingInput.nonEmpty))
69+
}
70+
71+
test("SPARK-60135: renew shared nested subqueries in traversal order") {
72+
val inner = LocalRelation($"a".int)
73+
val middle = LocalRelation($"b".int)
74+
val outer = LocalRelation($"c".int)
75+
val nested = Exists(inner)
76+
val shared = Exists(Filter(Not(nested), middle))
77+
val plan = Filter(And(shared, shared), outer)
78+
79+
// The first occurrence does not change, but the second must still be visited with the
80+
// updated set of relations. It must not be cached as an ineffective transformation.
81+
val result = DeduplicateRelations(plan).asInstanceOf[Filter]
82+
val subqueries = result.condition.collect { case s: Exists => s }
83+
assert(subqueries.size == 2)
84+
assert(subqueries.head eq shared)
85+
assert(subqueries.map(_.exprId) == Seq(shared.exprId, shared.exprId))
86+
val renewed = subqueries.last.plan.asInstanceOf[Filter]
87+
assert(renewed.outputSet.intersect(middle.outputSet).isEmpty)
88+
val renewedNested = renewed.condition.collect { case s: Exists => s }.head
89+
assert(renewedNested.exprId == nested.exprId)
90+
assert(renewedNested.plan.outputSet.intersect(inner.outputSet).isEmpty)
91+
assert(DeduplicateRelations(result) eq result)
92+
}
93+
94+
test("SPARK-60135: preserve correlation when renewing a shared plan with a wrapped subquery") {
95+
val outer = LocalRelation($"a".int)
96+
val inner = LocalRelation($"b".int)
97+
val outerAttr = outer.output.head
98+
val innerAttr = inner.output.head
99+
val subquery = Exists(
100+
Filter(EqualTo(innerAttr, OuterReference(outerAttr)), inner), Seq(outerAttr))
101+
val filter = Filter(Not(subquery), outer)
102+
val plan = Join(filter, filter, Inner, None, JoinHint.NONE)
103+
104+
val result = DeduplicateRelations(plan).asInstanceOf[Join]
105+
assert(result.left eq filter)
106+
assert(result.duplicateResolved)
107+
val renewed = result.right.asInstanceOf[Filter]
108+
val renewedSubquery = renewed.condition.collect { case s: Exists => s }.head
109+
val renewedInner = renewedSubquery.plan.asInstanceOf[Filter]
110+
assert(renewedSubquery.exprId == subquery.exprId)
111+
assert(renewedSubquery.outerAttrs == renewed.output)
112+
assert(renewedInner.condition ==
113+
EqualTo(renewedInner.child.output.head, OuterReference(renewed.output.head)))
114+
assert(renewedInner.outputSet.intersect(inner.outputSet).isEmpty)
115+
assert(!result.exists(_.missingInput.nonEmpty))
116+
assert(renewedInner.missingInput.isEmpty)
117+
}
118+
}
119+
120+
object DeduplicateRelationsSuite {
121+
private case class TraversalCountingExpression(child: Expression)
122+
extends Expression with Unevaluable {
123+
var traversals: Int = 0
124+
125+
override def children: Seq[Expression] = Seq(child)
126+
override def dataType: DataType = child.dataType
127+
override def nullable: Boolean = child.nullable
128+
129+
override def mapChildren(f: Expression => Expression): Expression = {
130+
traversals += 1
131+
super.mapChildren(f)
132+
}
133+
134+
override protected def withNewChildrenInternal(
135+
newChildren: IndexedSeq[Expression]): Expression = copy(child = newChildren.head)
136+
}
137+
}

0 commit comments

Comments
 (0)