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
Next Next commit
feat: pregel memory optimizations
  • Loading branch information
SemyonSinchenko committed Jun 9, 2026
commit b0d9f99614e1a5ac041cf5d391924d91fcc1a03e
34 changes: 30 additions & 4 deletions core/src/main/scala/org/graphframes/lib/Pregel.scala
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,13 @@ import scala.util.control.Breaks.breakable
* .run()
* }}}
*
* Migration note: pre 0.12 users that used edge columns in Pregel expressions should explicitly
* specify these columns using [[org.graphframes.lib.Pregel#requiredDstColumns]]. In 0.11 and
* earlier there was an unspecified bug that leads all the edge columns are always kept and
* persisted that created a bug memory pressure (2 columns in O(|E|) rows in the form of
* `StructType`). That behavior is considered as bug and starting from 0.12 edge columns are not
* kept by default.
*
* @param graph
* The graph that Pregel will run on.
* @see
Expand Down Expand Up @@ -106,6 +113,10 @@ class Pregel(val graph: GraphFrame)
private val requiredSrcColumnsList = collection.mutable.ListBuffer.empty[String]
private val requiredDstColumnsList = collection.mutable.ListBuffer.empty[String]

// Required columns for edges
// When empty, only src and dst are selected
private val requiredEdgeColumnsList = collection.mutable.ListBuffer.empty[String]

/** Sets the max number of iterations (default: 10). */
def setMaxIter(value: Int): this.type = {
maxIter = value
Expand Down Expand Up @@ -345,6 +356,13 @@ class Pregel(val graph: GraphFrame)
this
}

def requiredEdgeColumns(colName: String, colNames: String*): this.type = {
requiredEdgeColumnsList.clear()
requiredEdgeColumnsList += colName
requiredEdgeColumnsList ++= colNames
this
}

/**
* Defines how messages are aggregated after grouped by target vertex IDs.
*
Expand Down Expand Up @@ -419,10 +437,18 @@ class Pregel(val graph: GraphFrame)
"Optimization: skipping second join (dst state not required by message expressions)")
}

val edges = graph.edges
.select(col(SRC).alias("edge_src"), col(DST).alias("edge_dst"), struct(col("*")).as(EDGE))
.repartition(col("edge_src"))
.persist(intermediateStorageLevel)
val edges = (if (requiredEdgeColumnsList.isEmpty) {
graph.edges
.select(col(SRC).alias("edge_src"), col(DST).alias("edge_dst"))
} else {
graph.edges
.select(
col(SRC).alias("edge_src"),
col(DST).alias("edge_dst"),
struct(
requiredEdgeColumnsList.head,
requiredEdgeColumnsList.tail.toSeq: _*).as(EDGE))
}).repartition(col("edge_src")).persist(intermediateStorageLevel)

var iteration = 1

Expand Down
4 changes: 4 additions & 0 deletions core/src/test/scala/org/graphframes/lib/PregelSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -400,6 +400,7 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {

// Only uses Pregel.edge("weight") - dst join should be skipped
val resultDF = graph.pregel
.requiredEdgeColumns("weight")
.setMaxIter(1)
.withVertexColumn("received", lit(0L), coalesce(Pregel.msg, col("received")))
.sendMsgToSrc(Pregel.edge("weight"))
Expand All @@ -425,6 +426,7 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {

// Only uses Pregel.edge("weight") - dst join should be skipped
val result = graph.pregel
.requiredEdgeColumns("weight")
.setMaxIter(1) // Single iteration to simplify testing
.withVertexColumn("total", lit(0.0), coalesce(Pregel.msg, col("total")))
.sendMsgToDst(Pregel.edge("weight"))
Expand Down Expand Up @@ -523,6 +525,7 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {
val graph = GraphFrame(vertices, edges)

val result = graph.pregel
.requiredEdgeColumns("weights")
.setMaxIter(1)
.withVertexColumn("received", lit(0L), coalesce(Pregel.msg, col("received")))
// Use dst.key to look up value in edge.weights map
Expand All @@ -545,6 +548,7 @@ class PregelSuite extends SparkFunSuite with GraphFrameTestSparkContext {
val graph = GraphFrame(vertices, edges)

val result = graph.pregel
.requiredEdgeColumns("values")
.setMaxIter(1)
.withVertexColumn("received", lit(0L), coalesce(Pregel.msg, col("received")))
// Use dst.idx to index into edge.values array (element_at is 1-based)
Expand Down
Loading