Skip to content

Commit 8471b1a

Browse files
james-willisJamesclaude
authored
feat(pregel): automatically skip second join when dst columns not needed (#795)
* feat(pregel): automatically skip second join when dst columns not needed Implements automatic optimization for Pregel triplet generation that skips the second join (adding destination vertex state) when no message expressions reference dst.* columns. The optimization works by: 1. Analyzing all message expressions before the iteration loop 2. Extracting column prefixes (src, dst, edge) from the expression AST 3. Skipping the dst vertex join if no dst.* columns are referenced AND skipMessagesFromNonActiveVertices is disabled This provides significant performance improvement for algorithms like PageRank, directed LabelPropagation, and DetectingCycles that only need source vertex or edge columns in their message expressions. Closes #790 * fix: only analyze message expressions for dst detection, provide dst.id when skipping join The previous implementation incorrectly checked both the target ID expression and message expression for dst.* references. Since sendMsgToDst uses Pregel.dst(ID) as the target, it would always detect 'dst' as referenced even when the message itself only used src columns. This fix: 1. Only analyzes the message expressions (not target ID) for dst.* references 2. When skipping the join, creates a minimal dst struct with just the id from edge_dst so that sendMsgToDst can still route messages correctly Added test: 'sendMsgToDst with only src columns in message' to verify the optimization works correctly when dst.id is implicitly used for routing. * style: run scalafmt and update SparkShims docs to be implementation-agnostic * feat: skip dst join when only dst.id is referenced - Add extractColumnReferences to SparkShims returning Map[String, Set[String]] to track which specific fields are accessed under each prefix - Handle resolved expressions (AttributeReference, GetStructField) in addition to unresolved ones for more robust column detection - Update Pregel optimization to skip dst join when only dst.id is referenced since dst.id is available from the edge's dst column - Change optimization log message from logInfo to logDebug - Add test for dst.id-only reference case * refactor: address reviewer feedback on dst join optimization - Remove unused extractColumnPrefixes method from SparkShims (both Spark 3/4) - Refactor Pregel.scala to parse expressions once instead of twice - Add documentation for deeply nested struct access fallback behavior * test: add comprehensive tests for extractColumnReferences and dst join optimization - Add SparkShimsSuite with 22 unit tests for column reference extraction - Add 4 integration tests to PregelSuite for complex dst usage patterns - Fix UTF8String handling in UnresolvedExtractValue pattern matching * refactor: remove Pregel references from SparkShims comments Address peer feedback to keep SparkShims implementation-agnostic by removing specific algorithm references from comments while maintaining functional clarity. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude <noreply@anthropic.com> * perf: optimize caching and partitioning when skipping dst join Implement peer feedback suggestions: 1. Move cache/checkpoint logic before expensive operations - persist srcWithEdges when skipping dst join to avoid recomputation 2. Change partitioning to src-only when dst join is skipped since dst partitioning is unnecessary 3. Move dst state detection earlier to enable these optimizations These changes provide additional performance improvements for algorithms like PageRank that only need source vertex data. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude <noreply@anthropic.com> * fix: correct repartition syntax for conditional partitioning Fix Scala syntax error in repartition call that was causing CI build failures. Use proper sequence expansion syntax for multiple column repartitioning. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude <noreply@anthropic.com> * fix: correct repartition syntax for conditional partitioning Fix Scala syntax error in repartition call that was causing CI build failures. Use proper sequence expansion syntax for multiple column repartitioning. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude <noreply@anthropic.com> * style: apply scalafmt formatting to Pregel.scala Fix formatting issues that were causing CI scalafmt checks to fail. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude <noreply@anthropic.com> * address Sem's PR comments * push triplet filtering when no dst join --------- Co-authored-by: James <james@goivio.com> Co-authored-by: Claude <noreply@anthropic.com>
1 parent f67ba04 commit 8471b1a

5 files changed

Lines changed: 690 additions & 7 deletions

File tree

‎core/src/main/scala-spark-3/org/apache/spark/sql/graphframes/SparkShims.scala‎

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,13 +22,86 @@ import org.apache.spark.sql.DataFrame
2222
import org.apache.spark.sql.Dataset
2323
import org.apache.spark.sql.SparkSession
2424
import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute
25+
import org.apache.spark.sql.catalyst.analysis.UnresolvedExtractValue
26+
import org.apache.spark.sql.catalyst.expressions.AttributeReference
2527
import org.apache.spark.sql.catalyst.expressions.Expression
28+
import org.apache.spark.sql.catalyst.expressions.GetStructField
29+
import org.apache.spark.sql.catalyst.expressions.Literal
2630
import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
2731

2832
import scala.annotation.nowarn
33+
import scala.collection.mutable
2934

3035
object SparkShims {
3136

37+
/**
38+
* Extracts all column references from a Column expression, returning a map from top-level
39+
* prefix to the set of nested field names accessed under that prefix.
40+
*
41+
* For nested column references like "src.id" or "edge.weight", this returns Map("src" ->
42+
* Set("id"), "edge" -> Set("weight")). For top-level references like "src" (the whole struct),
43+
* it returns Map("src" -> Set()).
44+
*
45+
* This handles both unresolved expressions (UnresolvedAttribute, UnresolvedExtractValue) and
46+
* resolved expressions (AttributeReference, GetStructField).
47+
*
48+
* Note: Deeply nested struct access (e.g., "dst.location.city") is not fully parsed. In such
49+
* cases, the prefix is recorded with an empty field set, which causes callers to conservatively
50+
* assume the entire struct is needed. This is the safe/correct fallback behavior.
51+
*
52+
* @param spark
53+
* the SparkSession (unused in Spark 3, included for API compatibility with Spark 4)
54+
* @param expr
55+
* the Column expression to analyze
56+
* @return
57+
* a Map from column prefix to the set of nested field names accessed
58+
*/
59+
@nowarn
60+
def extractColumnReferences(spark: SparkSession, expr: Column): Map[String, Set[String]] = {
61+
val refs = mutable.Map.empty[String, mutable.Set[String]]
62+
63+
def addRef(prefix: String, field: Option[String]): Unit = {
64+
val fields = refs.getOrElseUpdate(prefix, mutable.Set.empty[String])
65+
field.foreach(fields += _)
66+
}
67+
68+
expr.expr.foreach {
69+
// Unresolved: col("src.id") -> UnresolvedAttribute(Seq("src", "id"))
70+
case UnresolvedAttribute(nameParts) if nameParts.nonEmpty =>
71+
addRef(nameParts.head, nameParts.lift(1))
72+
73+
// Unresolved: col("src")("id") -> UnresolvedExtractValue
74+
case UnresolvedExtractValue(child, extraction) =>
75+
child match {
76+
case UnresolvedAttribute(nameParts) if nameParts.nonEmpty =>
77+
extraction match {
78+
case Literal(fieldName: String, _) => addRef(nameParts.head, Some(fieldName))
79+
case Literal(fieldName, _) if fieldName != null =>
80+
// Handle UTF8String (Spark's internal string representation)
81+
addRef(nameParts.head, Some(fieldName.toString))
82+
case _ => addRef(nameParts.head, None) // Unknown field access
83+
}
84+
case _ => // Nested extraction we can't easily parse - conservative fallback
85+
}
86+
87+
// Resolved: AttributeReference for top-level columns
88+
case attr: AttributeReference =>
89+
addRef(attr.name, None)
90+
91+
// Resolved: GetStructField for nested field access like struct.field
92+
// Note: Only handles single-level nesting; deeper nesting falls through to default case
93+
case GetStructField(child, _, Some(fieldName)) =>
94+
child match {
95+
case attr: AttributeReference => addRef(attr.name, Some(fieldName))
96+
case _ => // Deeply nested struct access - conservative fallback (join will be used)
97+
}
98+
99+
case _ => // ignore other expression types
100+
}
101+
102+
refs.map { case (k, v) => k -> v.toSet }.toMap
103+
}
104+
32105
/**
33106
* Apply the given SQL expression (such as `id = 3`) to the field in a column, rather than to
34107
* the column itself.

‎core/src/main/scala-spark-4/org/apache/spark/sql/graphframes/SparkShims.scala‎

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,16 +21,90 @@ import org.apache.spark.sql.Column
2121
import org.apache.spark.sql.DataFrame
2222
import org.apache.spark.sql.SparkSession
2323
import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute
24+
import org.apache.spark.sql.catalyst.analysis.UnresolvedExtractValue
25+
import org.apache.spark.sql.catalyst.expressions.AttributeReference
2426
import org.apache.spark.sql.catalyst.expressions.Expression
27+
import org.apache.spark.sql.catalyst.expressions.GetStructField
28+
import org.apache.spark.sql.catalyst.expressions.Literal
2529
import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
2630
import org.apache.spark.sql.classic.ClassicConversions.*
2731
import org.apache.spark.sql.classic.DataFrame as ClassicDataFrame
2832
import org.apache.spark.sql.classic.Dataset
2933
import org.apache.spark.sql.classic.ExpressionUtils
3034
import org.apache.spark.sql.classic.SparkSession as ClassicSparkSession
3135

36+
import scala.collection.mutable
37+
3238
object SparkShims {
3339

40+
/**
41+
* Extracts all column references from a Column expression, returning a map from top-level
42+
* prefix to the set of nested field names accessed under that prefix.
43+
*
44+
* For nested column references like "src.id" or "edge.weight", this returns Map("src" ->
45+
* Set("id"), "edge" -> Set("weight")). For top-level references like "src" (the whole struct),
46+
* it returns Map("src" -> Set()).
47+
*
48+
* This handles both unresolved expressions (UnresolvedAttribute, UnresolvedExtractValue) and
49+
* resolved expressions (AttributeReference, GetStructField).
50+
*
51+
* Note: Deeply nested struct access (e.g., "dst.location.city") is not fully parsed. In such
52+
* cases, the prefix is recorded with an empty field set, which causes callers to conservatively
53+
* assume the entire struct is needed. This is the safe/correct fallback behavior.
54+
*
55+
* @param spark
56+
* the SparkSession (needed for expression conversion in Spark 4)
57+
* @param expr
58+
* the Column expression to analyze
59+
* @return
60+
* a Map from column prefix to the set of nested field names accessed
61+
*/
62+
def extractColumnReferences(spark: SparkSession, expr: Column): Map[String, Set[String]] = {
63+
val refs = mutable.Map.empty[String, mutable.Set[String]]
64+
65+
def addRef(prefix: String, field: Option[String]): Unit = {
66+
val fields = refs.getOrElseUpdate(prefix, mutable.Set.empty[String])
67+
field.foreach(fields += _)
68+
}
69+
70+
val converted = spark.asInstanceOf[ClassicSparkSession].converter(expr.node)
71+
converted.foreach {
72+
// Unresolved: col("src.id") -> UnresolvedAttribute(Seq("src", "id"))
73+
case UnresolvedAttribute(nameParts) if nameParts.nonEmpty =>
74+
addRef(nameParts.head, nameParts.lift(1))
75+
76+
// Unresolved: col("src")("id") -> UnresolvedExtractValue
77+
case UnresolvedExtractValue(child, extraction) =>
78+
child match {
79+
case UnresolvedAttribute(nameParts) if nameParts.nonEmpty =>
80+
extraction match {
81+
case Literal(fieldName: String, _) => addRef(nameParts.head, Some(fieldName))
82+
case Literal(fieldName, _) if fieldName != null =>
83+
// Handle UTF8String (Spark's internal string representation)
84+
addRef(nameParts.head, Some(fieldName.toString))
85+
case _ => addRef(nameParts.head, None) // Unknown field access
86+
}
87+
case _ => // Nested extraction we can't easily parse - conservative fallback
88+
}
89+
90+
// Resolved: AttributeReference for top-level columns
91+
case attr: AttributeReference =>
92+
addRef(attr.name, None)
93+
94+
// Resolved: GetStructField for nested field access like struct.field
95+
// Note: Only handles single-level nesting; deeper nesting falls through to default case
96+
case GetStructField(child, _, Some(fieldName)) =>
97+
child match {
98+
case attr: AttributeReference => addRef(attr.name, Some(fieldName))
99+
case _ => // Deeply nested struct access - conservative fallback (join will be used)
100+
}
101+
102+
case _ => // ignore other expression types
103+
}
104+
105+
refs.map { case (k, v) => k -> v.toSet }.toMap
106+
}
107+
34108
/**
35109
* Apply the given SQL expression (such as `id = 3`) to the field in a column, rather than to
36110
* the column itself.

‎core/src/main/scala/org/graphframes/lib/Pregel.scala‎

Lines changed: 49 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import org.apache.spark.sql.functions.col
2424
import org.apache.spark.sql.functions.explode
2525
import org.apache.spark.sql.functions.lit
2626
import org.apache.spark.sql.functions.struct
27+
import org.apache.spark.sql.graphframes.SparkShims
2728
import org.graphframes.GraphFrame
2829
import org.graphframes.GraphFrame.*
2930
import org.graphframes.Logging
@@ -397,9 +398,30 @@ class Pregel(val graph: GraphFrame)
397398
((initialAttributes :+ initialActiveVertexExpression.alias(
398399
Pregel.ACTIVE_FLAG_COL)) ++ initVertexCols): _*)
399400

401+
// Automatic optimization: detect if destination vertex state is needed by analyzing
402+
// the MESSAGE expressions only (not the target ID expressions, since dst.id is always
403+
// available from the edge). If no message expression references dst.* columns,
404+
// we can skip the second join entirely.
405+
// Additionally, if the only dst field referenced is "id", we can still skip since
406+
// dst.id is available from the edge's dst column.
407+
val messageExpressions = sendMsgs.toList.map { case (_, msgExpr) => msgExpr }
408+
val allDstRefs = messageExpressions.flatMap { expr =>
409+
SparkShims.extractColumnReferences(graph.spark, expr).get(DST)
410+
}
411+
val dstPrefixReferenced = allDstRefs.nonEmpty
412+
val dstFieldsReferenced = allDstRefs.flatten.toSet
413+
414+
// We need the dst join if dst is referenced AND fields other than just "id" are accessed
415+
val needsDstState =
416+
dstPrefixReferenced && (dstFieldsReferenced.isEmpty || dstFieldsReferenced != Set(ID))
417+
if (!needsDstState) {
418+
logDebug(
419+
"Optimization: skipping second join (dst state not required by message expressions)")
420+
}
421+
400422
val edges = graph.edges
401423
.select(col(SRC).alias("edge_src"), col(DST).alias("edge_dst"), struct(col("*")).as(EDGE))
402-
.repartition(col("edge_src"), col("edge_dst"))
424+
.repartition(col("edge_src"))
403425
.persist(intermediateStorageLevel)
404426

405427
var iteration = 1
@@ -431,15 +453,35 @@ class Pregel(val graph: GraphFrame)
431453
val currRoundPersistent = scala.collection.mutable.Queue[DataFrame]()
432454
currRoundPersistent.enqueue(currentVertices.persist(intermediateStorageLevel))
433455

434-
var tripletsDF = currentVertices
456+
// Prune non-active vertices early if skipMessagesFromNonActiveVertices
457+
// is enabled and we don't need the dst state.
458+
val srcVertices =
459+
if (!needsDstState && skipMessagesFromNonActiveVertices)
460+
currentVertices.filter(col(Pregel.ACTIVE_FLAG_COL))
461+
else currentVertices
462+
463+
// Build triplets: start with src vertex state joined with edges
464+
val srcWithEdges = srcVertices
435465
.select(struct(srcCols: _*).as(SRC))
436466
.join(edges, Pregel.src(ID) === col("edge_src"))
437-
.join(
438-
currentVertices.select(struct(dstCols: _*).as(DST)),
439-
col("edge_dst") === Pregel.dst(ID))
440-
.drop(col("edge_src"), col("edge_dst"))
441467

442-
if (skipMessagesFromNonActiveVertices) {
468+
// Only perform the second join (adding dst vertex state) if needed
469+
var tripletsDF = if (needsDstState) {
470+
srcWithEdges
471+
.join(
472+
currentVertices.select(struct(dstCols: _*).as(DST)),
473+
col("edge_dst") === Pregel.dst(ID))
474+
.drop(col("edge_src"), col("edge_dst"))
475+
} else {
476+
// Skip second join - dst state not needed by any message expression.
477+
// Create a minimal dst struct with just the id from edge_dst for sendMsgToDst to work.
478+
srcWithEdges
479+
.withColumn(DST, struct(col("edge_dst").as(ID)))
480+
.drop(col("edge_src"), col("edge_dst"))
481+
}
482+
483+
// Only prune here if we didn't prune above.
484+
if (needsDstState && skipMessagesFromNonActiveVertices) {
443485
tripletsDF = tripletsDF.filter(
444486
Pregel.src(Pregel.ACTIVE_FLAG_COL) || Pregel.dst(Pregel.ACTIVE_FLAG_COL))
445487
}

0 commit comments

Comments
 (0)