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
Do not orphan out of scope persisted dataframes in ConnectedComponent…
…s.run function
  • Loading branch information
james-willis committed Feb 25, 2025
commit 3d0dcf4b6c5fc0585544e1e5199da633c44910b1
18 changes: 14 additions & 4 deletions src/main/scala/org/graphframes/lib/ConnectedComponents.scala
Original file line number Diff line number Diff line change
Expand Up @@ -435,10 +435,20 @@ object ConnectedComponents extends Logging {

logInfo(s"$logPrefix Connected components converged in ${iteration - 1} iterations.")

logInfo(s"$logPrefix Join and return component assignments with original vertex IDs.")
vv.join(ee, vv(ID) === ee(DST), "left_outer")
.select(vv(ATTR), when(ee(SRC).isNull, vv(ID)).otherwise(ee(SRC)).as(COMPONENT))
.select(col(s"$ATTR.*"), col(COMPONENT))
logInfo(s"$logPrefix Join and return component assignments with original vertex IDs.")
val output = vv.join(ee, vv(ID) === ee(DST), "left_outer")
.select(vv(ATTR), when(ee(SRC).isNull, vv(ID)).otherwise(ee(SRC)).as(COMPONENT))
.select(col(s"$ATTR.*"), col(COMPONENT)).persist(intermediateStorageLevel)

// materialize the output DataFrame
output.count()

// clean up persisted DFs
for (persisted_df <- lastRoundPersistedDFs) {
persisted_df.unpersist()
}

output
} finally {
// Restore original AQE setting
spark.conf.set("spark.sql.adaptive.enabled", originalAQE)
Expand Down
11 changes: 11 additions & 0 deletions src/test/scala/org/graphframes/lib/ConnectedComponentsSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -253,6 +253,17 @@ class ConnectedComponentsSuite extends SparkFunSuite with GraphFrameTestSparkCon
}
}

test("not leaking cached data") {
val priorCachedDFsSize = spark.sparkContext.getPersistentRDDs.size

val cc = Graphs.friends.connectedComponents
val components = cc.run()

components.unpersist(blocking = true)

assert(spark.sparkContext.getPersistentRDDs.size === priorCachedDFsSize)
}

private def assertComponents[T: ClassTag: TypeTag](
actual: DataFrame,
expected: Set[Set[T]]): Unit = {
Expand Down