Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
push triplet filtering when no dst join
  • Loading branch information
james-willis committed Mar 12, 2026
commit 1be371ca6bd30b9323edd955d2906b6628944074
23 changes: 14 additions & 9 deletions core/src/main/scala/org/graphframes/lib/Pregel.scala
Original file line number Diff line number Diff line change
Expand Up @@ -402,7 +402,6 @@ class Pregel(val graph: GraphFrame)
// the MESSAGE expressions only (not the target ID expressions, since dst.id is always
// available from the edge). If no message expression references dst.* columns,
// we can skip the second join entirely.
// However, if skipMessagesFromNonActiveVertices is enabled, we need dst._pregel_is_active.
// Additionally, if the only dst field referenced is "id", we can still skip since
// dst.id is available from the edge's dst column.
val messageExpressions = sendMsgs.toList.map { case (_, msgExpr) => msgExpr }
Expand All @@ -411,12 +410,10 @@ class Pregel(val graph: GraphFrame)
}
val dstPrefixReferenced = allDstRefs.nonEmpty
val dstFieldsReferenced = allDstRefs.flatten.toSet
// We need the dst join if:
// 1. skipMessagesFromNonActiveVertices is enabled (needs dst._pregel_is_active), OR
// 2. dst is referenced AND fields other than just "id" are accessed
// (empty set means whole struct access like col("dst"), which also needs the join)
val needsDstState = skipMessagesFromNonActiveVertices ||
(dstPrefixReferenced && (dstFieldsReferenced.isEmpty || dstFieldsReferenced != Set(ID)))

// We need the dst join if dst is referenced AND fields other than just "id" are accessed
val needsDstState =
dstPrefixReferenced && (dstFieldsReferenced.isEmpty || dstFieldsReferenced != Set(ID))
if (!needsDstState) {
logDebug(
"Optimization: skipping second join (dst state not required by message expressions)")
Expand Down Expand Up @@ -456,8 +453,15 @@ class Pregel(val graph: GraphFrame)
val currRoundPersistent = scala.collection.mutable.Queue[DataFrame]()
currRoundPersistent.enqueue(currentVertices.persist(intermediateStorageLevel))

// Prune non-active vertices early if skipMessagesFromNonActiveVertices
// is enabled and we don't need the dst state.
val srcVertices =
if (!needsDstState && skipMessagesFromNonActiveVertices)
currentVertices.filter(col(Pregel.ACTIVE_FLAG_COL))
else currentVertices

// Build triplets: start with src vertex state joined with edges
val srcWithEdges = currentVertices
val srcWithEdges = srcVertices
.select(struct(srcCols: _*).as(SRC))
.join(edges, Pregel.src(ID) === col("edge_src"))

Expand All @@ -476,7 +480,8 @@ class Pregel(val graph: GraphFrame)
.drop(col("edge_src"), col("edge_dst"))
}

if (skipMessagesFromNonActiveVertices) {
// Only prune here if we didn't prune above.
if (needsDstState && skipMessagesFromNonActiveVertices) {
tripletsDF = tripletsDF.filter(
Pregel.src(Pregel.ACTIVE_FLAG_COL) || Pregel.dst(Pregel.ACTIVE_FLAG_COL))
}
Expand Down
11 changes: 6 additions & 5 deletions core/src/test/scala/org/graphframes/lib/PregelSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -357,9 +357,10 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {
assert(resultDF.sort("id").select("value").as[Int].collect() === Array.fill(n)(1))
}

test("automatic dst join NOT skipped when skipMessagesFromNonActiveVertices is enabled") {
// When skipMessagesFromNonActiveVertices is true, we need dst._pregel_is_active,
// so the second join must NOT be skipped even if message expressions don't use dst.
test("automatic dst join skipping with skipMessagesFromNonActiveVertices enabled") {
// When skipMessagesFromNonActiveVertices is true but message expressions don't
// reference dst columns, the dst join is still skipped. Active-vertex filtering
// is pushed before the src-edge join to reduce data volume.

val n = 5
val verDF = (1 to n).toDF("id").repartition(3)
Expand All @@ -370,8 +371,8 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {

val graph = GraphFrame(verDF, edgeDF)

// This only uses Pregel.src("value"), but skipMessagesFromNonActiveVertices
// requires dst._pregel_is_active, so dst join should NOT be skipped
// This only uses Pregel.src("value") - dst join should be skipped,
// and active-vertex filtering is applied before the src-edge join.
val resultDF = graph.pregel
.setMaxIter(n - 1)
.setSkipMessagesFromNonActiveVertices(true)
Expand Down
Loading