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
Add an experimental implementation of SP in GF
  • Loading branch information
SemyonSinchenko committed Mar 22, 2025
commit 2bca985678fba457b7cee31da4d22fddb5b33126
5 changes: 5 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -18,13 +18,18 @@ project/boot/
project/plugins/project/
.bsp
.metals
metals.sbt
Comment thread
SemyonSinchenko marked this conversation as resolved.
.bloop

# intellij
.idea/

# VSCode
.vscode

# Helix
.helix

# Mac
*.DS_Store

Expand Down
121 changes: 117 additions & 4 deletions src/main/scala/org/graphframes/lib/ShortestPaths.scala
Original file line number Diff line number Diff line change
Expand Up @@ -24,10 +24,12 @@ import scala.jdk.CollectionConverters._
import org.apache.spark.graphx.{lib => graphxlib}
import org.apache.spark.sql.{Column, DataFrame, Row}
import org.apache.spark.sql.api.java.UDF1
import org.apache.spark.sql.functions.{col, udf}
import org.apache.spark.sql.functions.{col, udf, map, lit, when, map_zip_with, reduce, map_values, transform_values, collect_list}
import org.apache.spark.sql.types.{IntegerType, MapType}

import org.graphframes.GraphFrame
import org.apache.spark.annotation.Experimental
import org.graphframes.Logging

/**
* Computes shortest paths from every vertex to the given set of landmark vertices. Note that this
Expand All @@ -39,7 +41,10 @@ import org.graphframes.GraphFrame
* shortest-path distance to each reachable landmark vertex.
*/
class ShortestPaths private[graphframes] (private val graph: GraphFrame) extends Arguments {
import org.graphframes.lib.ShortestPaths._

private var lmarks: Option[Seq[Any]] = None
private var algorithm: String = ALGO_GRAPHX

/**
* The list of landmark vertex ids. Shortest paths will be computed to each landmark.
Expand All @@ -57,14 +62,25 @@ class ShortestPaths private[graphframes] (private val graph: GraphFrame) extends
landmarks(value.asScala.toSeq)
}

def setAlgorithm(value: String): this.type = {
require(
supportedAlgorithms.contains(value),
s"Supported algorithms are {${supportedAlgorithms.mkString(", ")}}, but got $value.)")
algorithm = value
this
}

def run(): DataFrame = {
ShortestPaths.run(graph, check(lmarks, "landmarks"))
algorithm match {
case ALGO_GRAPHX => runInGraphX(graph, check(lmarks, "landmarks"))
case ALGO_GRAPHFRAMES => runInGraphFrames(graph, check(lmarks, "landmarks"))
}
}
}

private object ShortestPaths {
private object ShortestPaths extends Logging {

private def run(graph: GraphFrame, landmarks: Seq[Any]): DataFrame = {
private def runInGraphX(graph: GraphFrame, landmarks: Seq[Any]): DataFrame = {
val idType = graph.vertices.schema(GraphFrame.ID).dataType
val longIdToLandmark = landmarks.map(l => GraphXConversions.integralId(graph, l) -> l).toMap
val gx = graphxlib.ShortestPaths
Expand Down Expand Up @@ -94,6 +110,103 @@ private object ShortestPaths {
g.vertices.select(cols.toSeq: _*)
}

@Experimental
private def runInGraphFrames(
graph: GraphFrame,
landmarks: Seq[Any],
isDirected: Boolean = true): DataFrame = {
logWarn("The GraphFrames based implementation is slow and considered experimental!")
val vertexType = graph.vertices.schema(GraphFrame.ID).dataType

// For landmark vertices the initial distance to itself is set to 0
Comment thread
SemyonSinchenko marked this conversation as resolved.
def initDistancesMap(vertexId: Column): Column = {
var initCol = when(vertexId === lit(landmarks.head), map(lit(landmarks.head), lit(0)))
for (lmark <- landmarks.tail) {
initCol = initCol.when(vertexId === lit(lmark), map(lit(lmark), lit(0)))
}
initCol
}

// Concatenations of two distance maps:
// If one map is null just take another.
// In case both maps are not null:
// - iterate over keys
// - if value in the left map is null or greater than value from the right map take right one
// else take left one
def concatMaps(distancesLeft: Column, distancesRight: Column): Column =
when(distancesLeft.isNull, distancesRight)
.when(distancesRight.isNull, distancesLeft)
.otherwise(map_zip_with(
distancesLeft,
distancesRight,
(_, leftDistance, rightDistance) => {
when(leftDistance.isNull || (leftDistance > rightDistance), rightDistance)
.otherwise(leftDistance)
}))

// If distance is null, result of d + 1 will be null too
def incrementDistances(distancesMap: Column): Column =
transform_values(distancesMap, (_, distance) => distance + lit(1))

// Takes an array of distance maps and reduce them with concatMaps
def aggregateArrayOfDistanceMaps(arrayCol: Column): Column =
reduce(arrayCol, lit(null).cast(MapType(vertexType, IntegerType)), concatMaps)

// Checks that a sended distances map can change the destination distances.
// Evaluation would be "true" in case in the new distances map
// for one of keys present a non null value but in the old distances map it is null
// or new distance is less than old one.
def isDistanceImprovedWithMessage(newMap: Column, oldMap: Column): Column = reduce(
map_values(
map_zip_with(
newMap,
oldMap,
(_, newDistance, rightDistance) =>
(newDistance.isNotNull && rightDistance.isNull) || (newDistance < rightDistance))),
lit(false),
(left, right) => left || right)

// Overall:
// 1. Initialize distances
// 2. If new message can improve distances send it
// 3. Collect and aggregate messages
val pregel = graph.pregel
// TODO: set maxIter to Int.MaxValue and earlyStopping = true after merging #550
.setMaxIter(15)
.withVertexColumn(
DISTANCE_ID,
when(col(GraphFrame.ID).isInCollection(landmarks), initDistancesMap(col(GraphFrame.ID)))
.otherwise(map().cast(MapType(vertexType, IntegerType))),
concatMaps(col(DISTANCE_ID), Pregel.msg))
.sendMsgToSrc(
when(
isDistanceImprovedWithMessage(
incrementDistances(Pregel.dst(DISTANCE_ID)),
Pregel.src(DISTANCE_ID)),
incrementDistances(Pregel.dst(DISTANCE_ID))))
.aggMsgs(aggregateArrayOfDistanceMaps(collect_list(Pregel.msg)))

// Experimental feature
if (isDirected) {
pregel.run()
} else {
// For consider edges as undireceted,
// it is enough to send messages in both directions
pregel
.sendMsgToDst(
when(
isDistanceImprovedWithMessage(
incrementDistances(Pregel.src(DISTANCE_ID)),
Pregel.dst(DISTANCE_ID)),
incrementDistances(Pregel.src(DISTANCE_ID))))
.run()
}

}

private val DISTANCE_ID = "distances"
private val ALGO_GRAPHX = "graphx"
private val ALGO_GRAPHFRAMES = "graphframes"

val supportedAlgorithms: Array[String] = Array(ALGO_GRAPHX, ALGO_GRAPHFRAMES)
}
53 changes: 53 additions & 0 deletions src/test/scala/org/graphframes/lib/ShortestPathsSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,39 @@ class ShortestPathsSuite extends SparkFunSuite with GraphFrameTestSparkContext {
assert(results.toSet === shortestPaths)
}

test("Simple test with GraphFrames") {
val edgeSeq = Seq((1, 2), (1, 5), (2, 3), (2, 5), (3, 4), (4, 5), (4, 6))
.flatMap { case e =>
Seq(e, e.swap)
}
.map { case (src, dst) => (src.toLong, dst.toLong) }
val edges = spark.createDataFrame(edgeSeq).toDF("src", "dst")
val graph = GraphFrame.fromEdges(edges)

// Ground truth
val shortestPaths = Set(
(1, Map(1 -> 0, 4 -> 2)),
(2, Map(1 -> 1, 4 -> 2)),
(3, Map(1 -> 2, 4 -> 1)),
(4, Map(1 -> 2, 4 -> 0)),
(5, Map(1 -> 1, 4 -> 1)),
(6, Map(1 -> 3, 4 -> 1)))

val landmarks = Seq(1, 4).map(_.toLong)
val v2 = graph.shortestPaths.landmarks(landmarks).setAlgorithm("graphframes").run()

TestUtils.testSchemaInvariants(graph, v2)
TestUtils.checkColumnType(
v2.schema,
"distances",
DataTypes.createMapType(v2.schema("id").dataType, DataTypes.IntegerType, true))
val newVs = v2.select("id", "distances").collect().toSeq
val results = newVs.map { case Row(id: Long, spMap: Map[Long, Int] @unchecked) =>
(id, spMap)
}
assert(results.toSet === shortestPaths)
}

test("friends graph") {
val friends = examples.Graphs.friends
val v = friends.shortestPaths.landmarks(Seq("a", "d")).run()
Expand All @@ -78,4 +111,24 @@ class ShortestPathsSuite extends SparkFunSuite with GraphFrameTestSparkContext {
assert(results === expected)
}

test("friends graph with GraphFrames") {
val friends = examples.Graphs.friends
val v = friends.shortestPaths.landmarks(Seq("a", "d")).setAlgorithm("graphframes").run()
val expected = Set[(String, Map[String, Int])](
("a", Map("a" -> 0, "d" -> 2)),
("b", Map.empty),
("c", Map.empty),
("d", Map("a" -> 1, "d" -> 0)),
("e", Map("a" -> 2, "d" -> 1)),
("f", Map.empty),
("g", Map.empty))
val results = v
.select("id", "distances")
.collect()
.map { case Row(id: String, spMap: Map[String, Int] @unchecked) =>
(id, spMap)
}
.toSet
assert(results === expected)
}
}