Skip to content

Commit ce59df6

Browse files
committed
[SPARK-60067][SQL] Preserve CSE through scoped field updates
Let subexpression elimination inspect generated With bodies without moving unbound references, lazy definitions, or guarded values outside their scopes. Describe UpdateFields evaluation guards so shared values remain eligible without evaluating null-struct updates eagerly. Validated 363 tests across 11 SQL, Catalyst, codegen, and optimizer suites using recompiled sources and the cached SBT classpath.
1 parent ca280d6 commit ce59df6

5 files changed

Lines changed: 157 additions & 13 deletions

File tree

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/EquivalentExpressions.scala‎

Lines changed: 24 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ import scala.collection.mutable
2323

2424
import org.apache.spark.SparkException
2525
import org.apache.spark.sql.catalyst.expressions.codegen.CodegenFallback
26-
import org.apache.spark.sql.catalyst.trees.TreePattern.{LAMBDA_VARIABLE, PLAN_EXPRESSION}
26+
import org.apache.spark.sql.catalyst.trees.TreePattern.{COMMON_EXPR_REF, LAMBDA_VARIABLE, PLAN_EXPRESSION}
2727
import org.apache.spark.sql.internal.SQLConf
2828
import org.apache.spark.util.Utils
2929

@@ -62,7 +62,7 @@ class EquivalentExpressions(
6262
expr: Expression,
6363
map: mutable.HashMap[ExpressionEquals, ExpressionStats],
6464
useCount: Int = 1): Boolean = {
65-
if (expr.deterministic) {
65+
if (expr.deterministic && supportedExpression(expr)) {
6666
val wrapper = ExpressionEquals(expr)
6767
map.get(wrapper) match {
6868
case Some(stats) =>
@@ -159,12 +159,10 @@ class EquivalentExpressions(
159159
// that the peel lands on rather than about the `And`/`Or` chain above it.
160160
private def childrenToRecurse(expr: Expression): Seq[Expression] = expr match {
161161
case _: CodegenFallback => Nil
162-
// A `CommonExpressionRef` cannot be evaluated ahead of the `With` that binds it, for the same
163-
// reason a `LambdaVariable` cannot be evaluated ahead of its loop. Do not descend, so that no
164-
// subtree holding a reference -- nor a `CommonExpressionDef`, which is unevaluable -- becomes
165-
// a candidate. The `With` itself may still be deduplicated as a whole, which is safe: it
166-
// carries its own definitions and brings their slots into scope wherever it is generated.
167-
case _: With => Nil
162+
// Definitions are lazy. Only the generated body exposes always-evaluated subexpressions;
163+
// supportedExpression excludes candidates whose references would escape their binding scope.
164+
case withExpr: With =>
165+
if (withExpr.refUnderCodegenFallback) Nil else Seq(withExpr.child)
168166
// Peeling here is redundant for safety: `updateExprTree` peels before it descends, so this
169167
// method sees the peeled node either way. It stays because dropping it would pass
170168
// `updateExprTree` the input itself rather than the operand its peel lands on.
@@ -186,11 +184,29 @@ class EquivalentExpressions(
186184
// `LambdaVariable` is usually used as a loop variable, which can't be evaluated ahead of the
187185
// loop. So we can't evaluate sub-expressions containing `LambdaVariable` at the beginning.
188186
!(e.containsPattern(LAMBDA_VARIABLE) ||
187+
e.isInstanceOf[CommonExpressionDef] ||
188+
hasUnboundCommonExpressionRef(e) ||
189189
// `PlanExpression` wraps query plan. To compare query plans of `PlanExpression` on executor,
190190
// can cause error like NPE.
191191
(e.containsPattern(PLAN_EXPRESSION) && Utils.isInRunningSparkTask))
192192
}
193193

194+
private def hasUnboundCommonExpressionRef(
195+
expression: Expression,
196+
boundIds: Set[CommonExpressionId] = Set.empty): Boolean = {
197+
if (!expression.containsPattern(COMMON_EXPR_REF)) {
198+
false
199+
} else {
200+
expression match {
201+
case ref: CommonExpressionRef => !boundIds.contains(ref.id)
202+
case withExpr: With =>
203+
val innerIds = boundIds ++ withExpr.defs.map(_.id)
204+
withExpr.children.exists(hasUnboundCommonExpressionRef(_, innerIds))
205+
case other => other.children.exists(hasUnboundCommonExpressionRef(_, boundIds))
206+
}
207+
}
208+
}
209+
194210
/**
195211
* Adds the expression to this data structure recursively. Stops if a matching expression
196212
* is found. That is, if `expr` has already been added, its children are not added.

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/With.scala‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -165,7 +165,7 @@ case class With(child: Expression, defs: Seq[CommonExpressionDef])
165165
* falls back (`With` is not a `CodegenFallback`, so it has to be named rather than matched). Such
166166
* a reference needs its cell bound and cleared; the generated code clears codegen flags instead.
167167
*/
168-
@transient private lazy val refUnderCodegenFallback: Boolean = {
168+
@transient private[expressions] lazy val refUnderCodegenFallback: Boolean = {
169169
val ids = defs.map(_.id).toSet
170170
def holdsMyRef(e: Expression): Boolean = e.exists {
171171
case r: CommonExpressionRef => ids.contains(r.id)

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeCreator.scala‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -771,7 +771,7 @@ private case class UpdateFieldsExpression(
771771
structExpr: Expression,
772772
updatedValues: Seq[Expression],
773773
outputOrdinals: Seq[Int],
774-
outputFields: Seq[StructField]) extends Expression {
774+
outputFields: Seq[StructField]) extends ConditionalExpression {
775775

776776
override def children: Seq[Expression] = structExpr +: updatedValues
777777

@@ -781,6 +781,17 @@ private case class UpdateFieldsExpression(
781781
private lazy val indexedUpdatedValues = updatedValues.toArray
782782
private lazy val needsSource = nullable || indexedOrdinals.exists(_ >= 0)
783783

784+
override def alwaysEvaluatedInputs: Seq[Expression] =
785+
(if (needsSource) Seq(structExpr) else Nil) ++ (if (nullable) Nil else updatedValues)
786+
787+
override def branchGroups: Seq[Seq[Expression]] = Nil
788+
789+
override def withNewAlwaysEvaluatedInputs(
790+
inputs: Seq[Expression]): ConditionalExpression =
791+
copy(
792+
structExpr = if (needsSource) inputs.head else structExpr,
793+
updatedValues = if (nullable) updatedValues else if (needsSource) inputs.tail else inputs)
794+
784795
override lazy val dataType: StructType = {
785796
val sourceFields = structExpr.dataType.asInstanceOf[StructType]
786797
var updatedIndex = 0

‎sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/SubexpressionEliminationSuite.scala‎

Lines changed: 56 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -546,12 +546,65 @@ class SubexpressionEliminationSuite extends SparkFunSuite with ExpressionEvalHel
546546
assert(equivalence1.getCommonSubexpressions.size == 1)
547547
}
548548

549-
test("SPARK-58902: no candidate below a With, with or without the short-circuit peel") {
549+
test("With exposes independent subexpressions but not unbound references") {
550+
val input = AttributeReference("input", IntegerType)()
551+
val shared = Add(input, Literal(1))
552+
val expression = With(input) { case Seq(ref) =>
553+
Add(ref, Multiply(shared, shared))
554+
}
555+
val equivalence = new EquivalentExpressions
556+
equivalence.addExprTree(expression)
557+
assert(equivalence.getCommonSubexpressions == Seq(shared))
558+
assert(equivalence.getAllExprStates().forall { state =>
559+
state.expr == expression || !state.expr.exists(_.isInstanceOf[CommonExpressionRef])
560+
})
561+
}
562+
563+
test("With candidates must bind references in their lexical scope") {
564+
val input = AttributeReference("input", IntegerType)()
565+
val shared = Add(input, Literal(1))
566+
var inner: Expression = null
567+
val outer = With(input) { case Seq(outerRef) =>
568+
inner = With(input) { case Seq(innerRef) =>
569+
Add(Add(outerRef, innerRef), shared)
570+
}
571+
Add(inner, shared)
572+
}
573+
val equivalence = new EquivalentExpressions
574+
equivalence.addExprTree(outer)
575+
assert(equivalence.getExprState(outer).nonEmpty)
576+
assert(equivalence.getExprState(inner).isEmpty)
577+
assert(equivalence.getCommonSubexpressions == Seq(shared))
578+
assert(!equivalence.addExpr(inner))
579+
val closed = With(input) { case Seq(ref) => Add(ref, Literal(1)) }
580+
val sibling = new CommonExpressionRef(closed.defs.head)
581+
assert(!equivalence.addExpr(Add(closed, sibling)))
582+
}
583+
584+
test("With does not expose lazy definitions or interpreted children") {
585+
val input = AttributeReference("input", IntegerType)()
586+
val shared = Add(input, Literal(1))
587+
val lazyDefinition = With(Multiply(shared, shared)) { case Seq(ref) =>
588+
If(Literal(false), ref, input)
589+
}
590+
val fallback = With(input) { case Seq(ref) =>
591+
Add(CodegenFallbackExpression(ref), Multiply(shared, shared))
592+
}
593+
val guarded = With(input) { case Seq(ref) =>
594+
If(GreaterThan(ref, Literal(0)), Multiply(shared, shared), Literal(0))
595+
}
596+
Seq(lazyDefinition, fallback, guarded).foreach { expression =>
597+
val equivalence = new EquivalentExpressions
598+
equivalence.addExprTree(expression)
599+
assert(equivalence.getCommonSubexpressions.isEmpty)
600+
}
601+
}
602+
603+
test("SPARK-58902: candidates below a With cannot hold unbound references") {
550604
// A `CommonExpressionRef` can only be evaluated inside the `With` that binds it: the codegen
551605
// slots exist only while that `With` is being generated, and `getCommonExpr` throws otherwise.
552606
// Subexpression elimination generates its candidates outside every `With` scope, so a candidate
553-
// holding a reference would fail the query. `childrenToRecurse` therefore stops at a `With`,
554-
// including after `skipForShortcut` has peeled an `And`/`Or` down to one.
607+
// holding an unbound reference would fail the query.
555608
val a = AttributeReference("a", IntegerType)()
556609
val memoized = With(Add(a, a)) { case Seq(ref) =>
557610
And(GreaterThan(ref, Literal(0)), LessThan(ref, Literal(10)))

‎sql/core/src/test/scala/org/apache/spark/sql/ColumnExpressionSuite.scala‎

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -504,6 +504,70 @@ class ColumnExpressionSuite extends SharedSparkSession {
504504
}
505505
}
506506

507+
test("field updates should preserve shared value subexpression elimination") {
508+
withSQLConf(SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key -> "true",
509+
SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "true") {
510+
for {
511+
nested <- Seq(false, true)
512+
external <- Seq(false, true)
513+
opaque <- Seq(false, true)
514+
} {
515+
val calls = spark.sparkContext.longAccumulator
516+
val sourceCalls = spark.sparkContext.longAccumulator
517+
val value = udf((id: Long) => {
518+
calls.add(1)
519+
id + 1
520+
}).apply($"id")
521+
val source = if (opaque) {
522+
udf((id: Long) => {
523+
sourceCalls.add(1)
524+
((id, id), id)
525+
}).asNonNullable().asNondeterministic()($"id")
526+
} else {
527+
struct(struct($"id", $"id").as("_1"), $"id".as("_2"))
528+
}
529+
val prefix = if (nested) "_1." else ""
530+
val updated = source.withField(s"${prefix}first", value)
531+
.withField(s"${prefix}second", value)
532+
val columns = if (external) Seq(updated, value) else Seq(updated)
533+
val rows = spark.range(0, 10, 1, 1).select(columns: _*).collect()
534+
rows.zipWithIndex.foreach { case (row, index) =>
535+
val result = if (nested) row.getStruct(0).getStruct(0) else row.getStruct(0)
536+
assert(result.getLong(2) == index + 1L)
537+
assert(result.getLong(3) == index + 1L)
538+
if (external) assert(row.getLong(1) == index + 1L)
539+
}
540+
val expectedCalls = if (nested && opaque) 10 * (if (external) 3 else 2) else 10
541+
assert(calls.value == expectedCalls,
542+
s"nested=$nested external=$external opaque=$opaque calls=${calls.value}")
543+
if (opaque) assert(sourceCalls.value == 10)
544+
}
545+
}
546+
}
547+
548+
test("field updates should not hoist guarded shared values") {
549+
onEachEvalPath {
550+
withSQLConf(SQLConf.ANSI_ENABLED.key -> "true",
551+
SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key -> "true") {
552+
val schema = StructType(Seq(StructField("a", IntegerType, nullable = false)))
553+
val source = when($"id" === 0, lit(null).cast(schema))
554+
.otherwise(struct(lit(1).as("a")))
555+
val shared = lit(1L) / $"id"
556+
checkAnswer(
557+
spark.range(0, 2, 1, 1).select(
558+
source.withField("b", shared).withField("c", shared)),
559+
Seq(Row(null), Row(Row(1, 1.0, 1.0))))
560+
val throwingSource = udf((id: Long) => {
561+
if (id >= 0) throw new IllegalStateException(s"evaluated for $id") else (id, id)
562+
}).asNonNullable()
563+
checkAnswer(
564+
spark.range(1).select(
565+
throwingSource($"id").withField("_1", lit(1L)).withField("_2", lit(2L))),
566+
Row(Row(1L, 2L)))
567+
}
568+
}
569+
}
570+
507571
test("nested field updates on unresolved columns should remain linear") {
508572
val df = Seq((((1, 2), 3), 0)).toDF("s", "z")
509573
Seq(4, 8, 12).foreach { count =>

0 commit comments

Comments
 (0)