Skip to content
Merged
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
Next Next commit
Add early stopping to Pregel
  • Loading branch information
SemyonSinchenko committed Mar 21, 2025
commit 6265d19a577f5fb7d0b87ff22a4e35ad994bc670
97 changes: 54 additions & 43 deletions src/main/scala/org/graphframes/lib/Pregel.scala
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,8 @@ import org.graphframes.GraphFrame._
import org.apache.spark.sql.{Column, DataFrame}
import org.apache.spark.sql.functions.{array, col, explode, struct}

import scala.util.control.Breaks.{breakable, break}

/**
* Implements a Pregel-like bulk-synchronous message-passing API based on DataFrame operations.
*
Expand Down Expand Up @@ -231,54 +233,63 @@ class Pregel(val graph: GraphFrame) {

val shouldCheckpoint = checkpointInterval > 0

while (iteration <= maxIter) {
val tripletsDF = currentVertices
.select(struct(col("*")).as(SRC))
.join(edges.select(struct(col("*")).as(EDGE)), Pregel.src(ID) === Pregel.edge(SRC))
.join(
currentVertices.select(struct(col("*")).as(DST)),
Pregel.edge(DST) === Pregel.dst(ID))

var msgDF: DataFrame = tripletsDF
.select(explode(array(sendMsgsColList: _*)).as("msg"))
.select(col("msg.id"), col("msg.msg").as(Pregel.MSG_COL_NAME))

val newAggMsgDF = msgDF
.filter(Pregel.msg.isNotNull)
.groupBy(ID)
.agg(aggMsgsCol.as(Pregel.MSG_COL_NAME))

val verticesWithMsg = currentVertices.join(newAggMsgDF, Seq(ID), "left_outer")

var newVertexUpdateColDF = verticesWithMsg.select((col(ID) :: updateVertexCols): _*)

if (shouldCheckpoint && graph.spark.sparkContext.getCheckpointDir.isEmpty) {
// Spark Connect workaround
graph.spark.conf.getOption("spark.checkpoint.dir") match {
case Some(d) => graph.spark.sparkContext.setCheckpointDir(d)
case None =>
throw new IOException(
"Checkpoint directory is not set. Please set it first using sc.setCheckpointDir()" +
"or by specifying the conf 'spark.checkpoint.dir'.")
breakable {
while (iteration <= maxIter) {
val tripletsDF = currentVertices
.select(struct(col("*")).as(SRC))
.join(edges.select(struct(col("*")).as(EDGE)), Pregel.src(ID) === Pregel.edge(SRC))
.join(
currentVertices.select(struct(col("*")).as(DST)),
Pregel.edge(DST) === Pregel.dst(ID))

var msgDF: DataFrame = tripletsDF
.select(explode(array(sendMsgsColList: _*)).as("msg"))
.select(col("msg.id"), col("msg.msg").as(Pregel.MSG_COL_NAME))

if (msgDF.filter(Pregel.msg.isNotNull).isEmpty) {
if (vertexUpdateColDF != null) {
vertexUpdateColDF.unpersist()
}
break
}
}

if (shouldCheckpoint && iteration % checkpointInterval == 0) {
// do checkpoint, use lazy checkpoint because later we will materialize this DF.
newVertexUpdateColDF = newVertexUpdateColDF.checkpoint(eager = false)
// TODO: remove last checkpoint file.
}
newVertexUpdateColDF.cache()
newVertexUpdateColDF.count() // materialize it
val newAggMsgDF = msgDF
.filter(Pregel.msg.isNotNull)
.groupBy(ID)
.agg(aggMsgsCol.as(Pregel.MSG_COL_NAME))

if (vertexUpdateColDF != null) {
vertexUpdateColDF.unpersist()
}
vertexUpdateColDF = newVertexUpdateColDF
val verticesWithMsg = currentVertices.join(newAggMsgDF, Seq(ID), "left_outer")

var newVertexUpdateColDF = verticesWithMsg.select((col(ID) :: updateVertexCols): _*)

if (shouldCheckpoint && graph.spark.sparkContext.getCheckpointDir.isEmpty) {
// Spark Connect workaround
graph.spark.conf.getOption("spark.checkpoint.dir") match {
case Some(d) => graph.spark.sparkContext.setCheckpointDir(d)
case None =>
throw new IOException(
"Checkpoint directory is not set. Please set it first using sc.setCheckpointDir()" +
"or by specifying the conf 'spark.checkpoint.dir'.")
}
}

if (shouldCheckpoint && iteration % checkpointInterval == 0) {
// do checkpoint, use lazy checkpoint because later we will materialize this DF.
newVertexUpdateColDF = newVertexUpdateColDF.checkpoint(eager = false)
// TODO: remove last checkpoint file.
}
newVertexUpdateColDF.cache()
newVertexUpdateColDF.count() // materialize it

currentVertices = graph.vertices.join(vertexUpdateColDF, ID)
if (vertexUpdateColDF != null) {
vertexUpdateColDF.unpersist()
}
vertexUpdateColDF = newVertexUpdateColDF

currentVertices = graph.vertices.join(vertexUpdateColDF, ID)

iteration += 1
iteration += 1
}
}

currentVertices
Expand Down