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
PowerIterationClustering wrapper
  • Loading branch information
SemyonSinchenko committed Feb 27, 2025
commit 29d47419c4f93996b8892f12a639202f7f08f581
164 changes: 110 additions & 54 deletions python/graphframes/graphframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
import sys
from typing import Any, Union, Optional

if sys.version > '3':
if sys.version > "3":
basestring = str

from graphframes.lib import Pregel
Expand All @@ -27,7 +27,7 @@
from pyspark.storagelevel import StorageLevel


def _from_java_gf(jgf: Any, spark: SparkSession) -> 'GraphFrame':
def _from_java_gf(jgf: Any, spark: SparkSession) -> "GraphFrame":
"""
(internal) creates a python GraphFrame wrapper from a java GraphFrame.

Expand All @@ -37,10 +37,15 @@ def _from_java_gf(jgf: Any, spark: SparkSession) -> 'GraphFrame':
pe = DataFrame(jgf.edges(), spark)
return GraphFrame(pv, pe)


def _java_api(jsc: SparkContext) -> Any:
javaClassName = "org.graphframes.GraphFramePythonAPI"
return jsc._jvm.Thread.currentThread().getContextClassLoader().loadClass(javaClassName) \
.newInstance()
return (
jsc._jvm.Thread.currentThread()
.getContextClassLoader()
.loadClass(javaClassName)
.newInstance()
)


class GraphFrame:
Expand Down Expand Up @@ -76,16 +81,22 @@ def __init__(self, v: DataFrame, e: DataFrame) -> None:
# Check that provided DataFrames contain required columns
if self.ID not in v.columns:
raise ValueError(
"Vertex ID column {} missing from vertex DataFrame, which has columns: {}"
.format(self.ID, ",".join(v.columns)))
"Vertex ID column {} missing from vertex DataFrame, which has columns: {}".format(
self.ID, ",".join(v.columns)
)
)
if self.SRC not in e.columns:
raise ValueError(
"Source vertex ID column {} missing from edge DataFrame, which has columns: {}"
.format(self.SRC, ",".join(e.columns)))
"Source vertex ID column {} missing from edge DataFrame, which has columns: {}".format(
self.SRC, ",".join(e.columns)
)
)
if self.DST not in e.columns:
raise ValueError(
"Destination vertex ID column {} missing from edge DataFrame, which has columns: {}"
.format(self.DST, ",".join(e.columns)))
"Destination vertex ID column {} missing from edge DataFrame, which has columns: {}".format(
self.DST, ",".join(e.columns)
)
)

self._jvm_graph = self._jvm_gf_api.createGraph(v._jdf, e._jdf)

Expand All @@ -109,8 +120,8 @@ def edges(self) -> DataFrame:
def __repr__(self):
return self._jvm_graph.toString()

def cache(self) -> 'GraphFrame':
""" Persist the dataframe representation of vertices and edges of the graph with the default
def cache(self) -> "GraphFrame":
"""Persist the dataframe representation of vertices and edges of the graph with the default
storage level.
"""
self._jvm_graph.cache()
Expand All @@ -124,7 +135,7 @@ def persist(self, storageLevel: StorageLevel = StorageLevel.MEMORY_ONLY) -> "Gra
self._jvm_graph.persist(javaStorageLevel)
return self

def unpersist(self, blocking: bool = False) -> 'GraphFrame':
def unpersist(self, blocking: bool = False) -> "GraphFrame":
"""Mark the dataframe representation of vertices and edges of the graph as non-persistent,
and remove all blocks for it from memory and disk.
"""
Expand Down Expand Up @@ -209,12 +220,12 @@ def find(self, pattern: str) -> DataFrame:
jdf = self._jvm_graph.find(pattern)
return DataFrame(jdf, self._spark)

def filterVertices(self, condition: Union[str, Column]) -> 'GraphFrame':
def filterVertices(self, condition: Union[str, Column]) -> "GraphFrame":
"""
Filters the vertices based on expression, remove edges containing any dropped vertices.

:param condition: String or Column describing the condition expression for filtering.
:return: GraphFrame with filtered vertices and edges.
:return: GraphFrame with filtered vertices and edges.
"""

if isinstance(condition, basestring):
Expand All @@ -225,12 +236,12 @@ def filterVertices(self, condition: Union[str, Column]) -> 'GraphFrame':
raise TypeError("condition should be string or Column")
return _from_java_gf(jdf, self._spark)

def filterEdges(self, condition: Union[str, Column]) -> 'GraphFrame':
def filterEdges(self, condition: Union[str, Column]) -> "GraphFrame":
"""
Filters the edges based on expression, keep all vertices.

:param condition: String or Column describing the condition expression for filtering.
:return: GraphFrame with filtered edges.
:return: GraphFrame with filtered edges.
"""
if isinstance(condition, basestring):
jdf = self._jvm_graph.filterEdges(condition)
Expand All @@ -240,37 +251,39 @@ def filterEdges(self, condition: Union[str, Column]) -> 'GraphFrame':
raise TypeError("condition should be string or Column")
return _from_java_gf(jdf, self._spark)

def dropIsolatedVertices(self) -> 'GraphFrame':
def dropIsolatedVertices(self) -> "GraphFrame":
"""
Drops isolated vertices, vertices are not contained in any edges.

:return: GraphFrame with filtered vertices.
:return: GraphFrame with filtered vertices.
"""
jdf = self._jvm_graph.dropIsolatedVertices()
return _from_java_gf(jdf, self._spark)

def bfs(self, fromExpr: str, toExpr: str,
edgeFilter: Optional[str] = None,
maxPathLength: int = 10) -> DataFrame:
def bfs(
self, fromExpr: str, toExpr: str, edgeFilter: Optional[str] = None, maxPathLength: int = 10
) -> DataFrame:
"""
Breadth-first search (BFS).

See Scala documentation for more details.

:return: DataFrame with one Row for each shortest path between matching vertices.
"""
builder = self._jvm_graph.bfs()\
.fromExpr(fromExpr)\
.toExpr(toExpr)\
.maxPathLength(maxPathLength)
builder = (
self._jvm_graph.bfs().fromExpr(fromExpr).toExpr(toExpr).maxPathLength(maxPathLength)
)
if edgeFilter is not None:
builder.edgeFilter(edgeFilter)
jdf = builder.run()
return DataFrame(jdf, self._spark)

def aggregateMessages(self, aggCol: Union[Column, str],
sendToSrc: Union[Column, str, None] = None,
sendToDst: Union[Column, str, None] = None) -> DataFrame:
def aggregateMessages(
self,
aggCol: Union[Column, str],
sendToSrc: Union[Column, str, None] = None,
sendToDst: Union[Column, str, None] = None,
) -> DataFrame:
"""
Aggregates messages from the neighbours.

Expand Down Expand Up @@ -314,9 +327,12 @@ def aggregateMessages(self, aggCol: Union[Column, str],

# Standard algorithms

def connectedComponents(self, algorithm: str = 'graphframes',
checkpointInterval: int = 2,
broadcastThreshold: int = 1000000) -> DataFrame:
def connectedComponents(
self,
algorithm: str = "graphframes",
checkpointInterval: int = 2,
broadcastThreshold: int = 1000000,
) -> DataFrame:
"""
Computes the connected components of the graph.

Expand All @@ -330,11 +346,13 @@ def connectedComponents(self, algorithm: str = 'graphframes',

:return: DataFrame with new vertices column "component"
"""
jdf = self._jvm_graph.connectedComponents() \
.setAlgorithm(algorithm) \
.setCheckpointInterval(checkpointInterval) \
.setBroadcastThreshold(broadcastThreshold) \
jdf = (
self._jvm_graph.connectedComponents()
.setAlgorithm(algorithm)
.setCheckpointInterval(checkpointInterval)
.setBroadcastThreshold(broadcastThreshold)
.run()
)
return DataFrame(jdf, self._spark)

def labelPropagation(self, maxIter: int) -> DataFrame:
Expand All @@ -349,10 +367,13 @@ def labelPropagation(self, maxIter: int) -> DataFrame:
jdf = self._jvm_graph.labelPropagation().maxIter(maxIter).run()
return DataFrame(jdf, self._spark)

def pageRank(self, resetProbability: float = 0.15,
sourceId: Optional[Any] = None,
maxIter: Optional[int] = None,
tol: Optional[float] = None) -> 'GraphFrame':
def pageRank(
self,
resetProbability: float = 0.15,
sourceId: Optional[Any] = None,
maxIter: Optional[int] = None,
tol: Optional[float] = None,
) -> "GraphFrame":
"""
Runs the PageRank algorithm on the graph.
Note: Exactly one of fixed_num_iter or tolerance must be set.
Expand All @@ -379,9 +400,12 @@ def pageRank(self, resetProbability: float = 0.15,
jgf = builder.run()
return _from_java_gf(jgf, self._spark)

def parallelPersonalizedPageRank(self, resetProbability: float = 0.15,
sourceIds: Optional[list[Any]] = None,
maxIter: Optional[int] = None) -> 'GraphFrame':
def parallelPersonalizedPageRank(
self,
resetProbability: float = 0.15,
sourceIds: Optional[list[Any]] = None,
maxIter: Optional[int] = None,
) -> "GraphFrame":
"""
Run the personalized PageRank algorithm on the graph,
from the provided list of sources in parallel for a fixed number of iterations.
Expand All @@ -393,7 +417,9 @@ def parallelPersonalizedPageRank(self, resetProbability: float = 0.15,
:param maxIter: the fixed number of iterations this algorithm runs
:return: GraphFrame with new vertices column "pageranks" and new edges column "weight"
"""
assert sourceIds is not None and len(sourceIds) > 0, "Source vertices Ids sourceIds must be provided"
assert (
sourceIds is not None and len(sourceIds) > 0
), "Source vertices Ids sourceIds must be provided"
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()
Expand Down Expand Up @@ -427,10 +453,17 @@ def stronglyConnectedComponents(self, maxIter: int) -> DataFrame:
jdf = self._jvm_graph.stronglyConnectedComponents().maxIter(maxIter).run()
return DataFrame(jdf, self._spark)

def svdPlusPlus(self, rank: int = 10, maxIter: int = 2,
minValue: float = 0.0, maxValue: float = 5.0,
gamma1: float = 0.007, gamma2: float = 0.007,
gamma6: float = 0.005, gamma7: float = 0.015) -> tuple[DataFrame, float]:
def svdPlusPlus(
self,
rank: int = 10,
maxIter: int = 2,
minValue: float = 0.0,
maxValue: float = 5.0,
gamma1: float = 0.007,
gamma2: float = 0.007,
gamma6: float = 0.005,
gamma7: float = 0.015,
) -> tuple[DataFrame, float]:
"""
Runs the SVD++ algorithm.

Expand Down Expand Up @@ -458,16 +491,39 @@ def triangleCount(self) -> DataFrame:
jdf = self._jvm_graph.triangleCount().run()
return DataFrame(jdf, self._spark)

def powerIterationClustering(
self, k: int, maxIter: int, weightCol: Optional[str] = None
) -> DataFrame:
"""
Power Iteration Clustering (PIC), a scalable graph clustering algorithm developed by Lin and Cohen.
From the abstract: PIC finds a very low-dimensional embedding of a dataset using truncated power iteration
on a normalized pair-wise similarity matrix of the data.

:param k: the numbers of clusters to create
:param maxIter: param for maximum number of iterations (>= 0)
:param weightCol: optional name of weight column, 1.0 is used if not provided

:return: DataFrame with new column "cluster"
"""
if weightCol:
weightCol = self._spark._jvm.scala.Option.apply(weightCol)
else:
weightCol = self._spark._jvm.scala.Option.empty()
jdf = self._jvm_graph.powerIterationClustering(k, maxIter, weightCol)
return DataFrame(jdf, self._spark)


def _test():
import doctest
import graphframe

globs = graphframe.__dict__.copy()
globs['sc'] = SparkContext('local[4]', 'PythonTest', batchSize=2)
globs['spark'] = SparkSession(globs['sc']).builder.getOrCreate()
globs["sc"] = SparkContext("local[4]", "PythonTest", batchSize=2)
globs["spark"] = SparkSession(globs["sc"]).builder.getOrCreate()
(failure_count, test_count) = doctest.testmod(
globs=globs, optionflags=doctest.ELLIPSIS | doctest.NORMALIZE_WHITESPACE)
globs['sc'].stop()
globs=globs, optionflags=doctest.ELLIPSIS | doctest.NORMALIZE_WHITESPACE
)
globs["sc"].stop()
if failure_count:
exit(-1)

Expand Down
Loading