Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -76,3 +76,6 @@ connect/project

# Local spark distro
spark-*

# Zed
.zed
95 changes: 45 additions & 50 deletions python/graphframes/classic/graphframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,15 +108,13 @@ def detectingCycles(
use_local_checkpoints: bool = False,
intermediate_storage_level: StorageLevel = StorageLevel.MEMORY_AND_DISK_DESER,
) -> DataFrame:
jdf = (
self._jvm_graph.detectingCycles()
.setUseLocalCheckpoints(use_local_checkpoints)
.setCheckpointInterval(checkpoint_interval)
.setIntermediateStorageLevel(
storage_level_to_jvm(intermediate_storage_level, self._spark)
)
.run()
builder = self._jvm_graph.detectingCycles()
builder.setUseLocalCheckpoints(use_local_checkpoints)
builder.setCheckpointInterval(checkpoint_interval)
builder.setIntermediateStorageLevel(
storage_level_to_jvm(intermediate_storage_level, self._spark)
)
jdf = builder.run()

return DataFrame(jdf, self._spark)

Expand Down Expand Up @@ -147,7 +145,7 @@ def aggregateMessages(
intermediate_storage_level: StorageLevel,
) -> DataFrame:
builder = self._jvm_graph.aggregateMessages()
builder = builder.setIntermediateStorageLevel(
builder.setIntermediateStorageLevel(
storage_level_to_jvm(intermediate_storage_level, self._spark)
)
if len(sendToSrc) == 1:
Expand Down Expand Up @@ -212,17 +210,16 @@ def connectedComponents(
max_iter: int,
storage_level: StorageLevel,
) -> DataFrame:
jdf = (
self._jvm_graph.connectedComponents()
.setAlgorithm(algorithm)
.setCheckpointInterval(checkpointInterval)
.setBroadcastThreshold(broadcastThreshold)
.setUseLabelsAsComponents(useLabelsAsComponents)
.setUseLocalCheckpoints(use_local_checkpoints)
.maxIter(max_iter)
.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
.run()
)
java_cc = self._jvm_graph.connectedComponents()
java_cc.setAlgorithm(algorithm)
java_cc.setCheckpointInterval(checkpointInterval)
java_cc.setBroadcastThreshold(broadcastThreshold)
java_cc.setUseLabelsAsComponents(useLabelsAsComponents)
java_cc.setUseLocalCheckpoints(use_local_checkpoints)
java_cc.maxIter(max_iter)
java_cc.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
jdf = java_cc.run()

return DataFrame(jdf, self._spark)

def labelPropagation(
Expand All @@ -233,15 +230,14 @@ def labelPropagation(
checkpoint_interval: int,
storage_level: StorageLevel,
) -> DataFrame:
jdf = (
self._jvm_graph.labelPropagation()
.maxIter(maxIter)
.setAlgorithm(algorithm)
.setUseLocalCheckpoints(use_local_checkpoints)
.setCheckpointInterval(checkpoint_interval)
.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
.run()
)
java_cdlp = self._jvm_graph.labelPropagation()
java_cdlp.maxIter(maxIter)
java_cdlp.setAlgorithm(algorithm)
java_cdlp.setUseLocalCheckpoints(use_local_checkpoints)
java_cdlp.setCheckpointInterval(checkpoint_interval)
java_cdlp.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
jdf = java_cdlp.run()

return DataFrame(jdf, self._spark)

def pageRank(
Expand All @@ -253,13 +249,13 @@ def pageRank(
) -> "GraphFrame":
builder = self._jvm_graph.pageRank().resetProbability(resetProbability)
if sourceId is not None:
builder = builder.sourceId(sourceId)
builder.sourceId(sourceId)
if maxIter is not None:
builder = builder.maxIter(maxIter)
builder.maxIter(maxIter)
assert tol is None, "Exactly one of maxIter or tol should be set."
else:
assert tol is not None, "Exactly one of maxIter or tol should be set."
builder = builder.tol(tol)
builder.tol(tol)
jgf = builder.run()
return _from_java_gf(jgf, self._spark)

Expand All @@ -275,9 +271,9 @@ def parallelPersonalizedPageRank(
assert maxIter is not None, "Max number of iterations maxIter must be provided"
sourceIds = self._sc._jvm.PythonUtils.toArray(sourceIds)
builder = self._jvm_graph.parallelPersonalizedPageRank()
builder = builder.resetProbability(resetProbability)
builder = builder.sourceIds(sourceIds)
builder = builder.maxIter(maxIter)
builder.resetProbability(resetProbability)
builder.sourceIds(sourceIds)
builder.maxIter(maxIter)
jgf = builder.run()
return _from_java_gf(jgf, self._spark)

Expand All @@ -289,19 +285,20 @@ def shortestPaths(
checkpoint_interval: int,
storage_level: StorageLevel,
) -> DataFrame:
jdf = (
self._jvm_graph.shortestPaths()
.landmarks(landmarks)
.setAlgorithm(algorithm)
.setUseLocalCheckpoints(use_local_checkpoints)
.setCheckpointInterval(checkpoint_interval)
.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
.run()
)
java_sp = self._jvm_graph.shortestPaths()
java_sp.landmarks(landmarks)
java_sp.setAlgorithm(algorithm)
java_sp.setUseLocalCheckpoints(use_local_checkpoints)
java_sp.setCheckpointInterval(checkpoint_interval)
java_sp.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
jdf = java_sp.run()

return DataFrame(jdf, self._spark)

def stronglyConnectedComponents(self, maxIter: int) -> DataFrame:
jdf = self._jvm_graph.stronglyConnectedComponents().maxIter(maxIter).run()
builder = self._jvm_graph.stronglyConnectedComponents()
builder.maxIter(maxIter)
jdf = builder.run()
return DataFrame(jdf, self._spark)

def svdPlusPlus(
Expand All @@ -325,11 +322,9 @@ def svdPlusPlus(
return (v, loss)

def triangleCount(self, storage_level: StorageLevel) -> DataFrame:
jdf = (
self._jvm_graph.triangleCount()
.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
.run()
)
builder = self._jvm_graph.triangleCount()
builder.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
jdf = builder.run()
return DataFrame(jdf, self._spark)

def powerIterationClustering(
Expand Down
24 changes: 21 additions & 3 deletions python/tests/test_graphframes.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,9 +377,7 @@ def test_connected_components(
spark: SparkSession, args: PregelArguments, cc_args: tuple[int, bool]
) -> None:
v = spark.createDataFrame([(0, "a", "b")], ["id", "vattr", "gender"])
e = spark.createDataFrame([(0, 0, 1)], ["src", "dst", "test"]).filter("src > 10")
v = spark.createDataFrame([(0, "a", "b")], ["id", "vattr", "gender"])
e = spark.createDataFrame([(0, 0, 1)], ["src", "dst", "test"]).filter("src > 10")
e = spark.createDataFrame([(0, 0, 1)], ["src", "dst", "test"])
g = GraphFrame(v, e)
comps = g.connectedComponents(
algorithm=args.algorithm,
Expand Down Expand Up @@ -419,6 +417,26 @@ def test_connected_components2(
_ = comps.unpersist()


def test_connected_components_example(spark: SparkSession) -> None:
nodes = [(1, "Alice", 30), (2, "Bob", 25), (3, "Charlie", 35)]
nodes_df = spark.createDataFrame(nodes, ["id", "name", "age"])

edges = [
(1, 2, "friend"),
(2, 1, "friend"),
(2, 3, "friend"),
(3, 2, "enemy"), # eek!
]
edges_df = spark.createDataFrame(edges, ["src", "dst", "relationship"])

g = GraphFrame(nodes_df, edges_df)
cc = g.connectedComponents()
cc.write.mode("overwrite").format("noop").save()
res = cc.collect()
assert len(res) == 3
_ = cc.unpersist()


@pytest.mark.parametrize("args", PREGEL_ARGUMENTS, ids=PREGEL_IDS)
def test_shortest_paths(spark: SparkSession, args: PregelArguments) -> None:
edges = [(1, 2), (1, 5), (2, 3), (2, 5), (3, 4), (4, 5), (4, 6)]
Expand Down