Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
2 changes: 2 additions & 0 deletions connect/src/main/protobuf/graphframes.proto
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,8 @@ message SVDPlusPlus {

message TriangleCount {
optional StorageLevel storage_level = 1;
optional string algorithm = 2;
optional int32 lg_nom_entries = 3;
}

message Triplets {}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -405,7 +405,20 @@ object GraphFramesConnectUtils {
svdResult.withColumn("loss", lit(svd.loss))
}
case proto.GraphFramesAPI.MethodCase.TRIANGLE_COUNT => {
val trCounter = graphFrame.triangleCount
val message = apiMessage.getTriangleCount()

val algorithm = if (message.hasAlgorithm) {
message.getAlgorithm
} else {
"exact"
}

val lgNomEntries = if (message.hasLgNomEntries) {
message.getLgNomEntries
} else { 12 }

val trCounter =
graphFrame.triangleCount.setAlgorithm(algorithm).setLgNomEntries(lgNomEntries)
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated

if (apiMessage.getTriangleCount.hasStorageLevel) {
trCounter
Expand Down
10 changes: 10 additions & 0 deletions core/src/main/scala/org/graphframes/exceptions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -43,3 +43,13 @@ class InvalidPropertyGroupException(message: String) extends Exception(message)
* A descriptive error message providing details about why the graph operation is invalid.
*/
class InvalidGraphException(message: String) extends Exception(message)

/**
* Exception thrown when a Spark version requirement is not met.
*
* @param version
* The minimum version of Apache Spark required.
*/
class GraphFramesRequireSpark(version: String)
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated
extends Exception(
s"Called GraphFrames feature require at least $version or above version of Apache Spark")
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated
109 changes: 98 additions & 11 deletions core/src/main/scala/org/graphframes/lib/TriangleCount.scala
Original file line number Diff line number Diff line change
Expand Up @@ -21,30 +21,58 @@ import org.apache.spark.sql.DataFrame
import org.apache.spark.sql.functions.*
import org.apache.spark.storage.StorageLevel
import org.graphframes.GraphFrame
import org.graphframes.GraphFramesRequireSpark
import org.graphframes.Logging
import org.graphframes.WithIntermediateStorageLevel

/**
* Computes the number of triangles passing through each vertex.
* Triangle count implementation.
*
* This algorithm ignores edge direction; i.e., all edges are treated as undirected. In a
* multigraph, duplicate edges will be counted only once.
* This class provides two algorithms for counting triangles:
* - A direct version that computes exact triangle counts using set intersection of neighbor
* lists.
* - An approximate version based on the DataSketches library (Theta sketches), which trades off
* accuracy for performance on large-scale graphs.
*
* **WARNING** This implementation is based on intersections of neighbor sets, which requires
* collecting both SRC and DST neighbors per edge! This will blow up memory in case the graph
* contains very high-degree nodes (power-law networks). Consider sampling strategies for that
* case!
*
* The returned DataFrame contains all the original vertex information and one additional column:
* - count (`LongType`): the count of triangles
* The output DataFrame contains two columns:
* - "id": the vertex id
* - "count": the number of triangles passing through the vertex
*/
class TriangleCount private[graphframes] (private val graph: GraphFrame)
extends Arguments
with Serializable
with WithIntermediateStorageLevel {

private var algorithm: String = "exact"
private val supportedAlgorithms: Set[String] = Set("exact", "approx")
private var lgNomEntries: Int = 12

/**
* Sets the log2 of the nominal entries for the Theta sketch (only for "approx" algorithm).
* Default is 12 (4096 entries).
*/
def setLgNomEntries(value: Int): this.type = {
lgNomEntries = value
Comment thread
SemyonSinchenko marked this conversation as resolved.
this
}

/**
* Sets the triangle counting algorithm. Options are "exact" (default) or "approx".
*/
def setAlgorithm(value: String): this.type = {
require(
supportedAlgorithms.contains(value),
s"supported algorithms: ${supportedAlgorithms.mkString}")
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated
algorithm = value
this
}

def run(): DataFrame = {
TriangleCount.run(graph, intermediateStorageLevel)
if (algorithm == "exact") {
TriangleCount.run(graph, intermediateStorageLevel)
} else {
TriangleCount.approximateRun(graph, intermediateStorageLevel, lgNomEntries)
}
}
}

Expand All @@ -65,6 +93,65 @@ private object TriangleCount extends Logging {
GraphFrame(graph.vertices.select(ID), dedupedE).dropIsolatedVertices()
}

private def approximateRun(
graph: GraphFrame,
intermediateStorageLevel: StorageLevel,
lgNomEntries: Int): DataFrame = {
val spark = graph.vertices.sparkSession
val sparkVersion = spark.version

if (sparkVersion.substring(0, 3) < "4.1") {
throw new GraphFramesRequireSpark("4.1.0")
}

val thetaSketchAgg = (colName: String) => expr(s"theta_sketch_agg($colName, $lgNomEntries)")
val thetaSketchIntersect = (colLeft: String, colRight: String) =>
expr(s"theta_sketch_estimate(theta_intersection($colLeft, $colRight))")

val g2 = prepareGraph(graph)

val verticesWithNeighbors = g2.aggregateMessages
Comment thread
SemyonSinchenko marked this conversation as resolved.
.setIntermediateStorageLevel(intermediateStorageLevel)
.sendToSrc(AggregateMessages.dst(ID))
.sendToDst(AggregateMessages.src(ID))
.agg(thetaSketchAgg(AggregateMessages.MSG_COL_NAME).alias("neighbors"))

val triangles = verticesWithNeighbors
.select(col(ID), col("neighbors").alias("src_set"))
.join(g2.edges, col(ID) === col(SRC))
.drop(ID)
.join(
verticesWithNeighbors.select(col(ID), col("neighbors").alias("dst_set")),
col(ID) === col(DST))
.drop(ID)
// Count of common neighbors of SRC and DST
.withColumn("triplets", thetaSketchIntersect("src_set", "dst_set"))
.filter(col("triplets") > lit(0))
.persist(intermediateStorageLevel)

val srcTriangles = triangles.groupBy(SRC).agg(sum(col("triplets")).alias("src_triplets"))
val dstTriangles = triangles.groupBy(DST).agg(sum(col("triplets")).alias("dst_triplets"))

val result = graph.vertices
.join(srcTriangles, col(ID) === col(SRC), "left_outer")
.join(dstTriangles, col(ID) === col(DST), "left_outer")
// Each triangle counted twice, so divide by 2.
.withColumn(
COUNT_ID,
floor(
when(col("src_triplets").isNull && col("dst_triplets").isNull, lit(0))
.when(col("src_triplets").isNull, col("dst_triplets"))
.when(col("dst_triplets").isNull, col("src_triplets"))
.otherwise(col("src_triplets") + col("dst_triplets")) / lit(2)))
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated

result.persist(intermediateStorageLevel)
result.count()
verticesWithNeighbors.unpersist()
triangles.unpersist()
resultIsPersistent()
result
Comment thread
SemyonSinchenko marked this conversation as resolved.
}

private def run(graph: GraphFrame, intermediateStorageLevel: StorageLevel): DataFrame = {
val g2 = prepareGraph(graph)

Expand Down
79 changes: 79 additions & 0 deletions core/src/test/scala/org/graphframes/lib/TriangleCountSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -148,4 +148,83 @@ class TriangleCountSuite extends SparkFunSuite with GraphFrameTestSparkContext {
}
v2.unpersist()
}

test("Approximate triangle count") {
val sparkVersion = spark.version
if (sparkVersion.substring(0, 3) >= "4.1") {
Comment thread
SemyonSinchenko marked this conversation as resolved.
Outdated
val edges = spark
.createDataFrame(
Seq(0L -> 1L, 1L -> 2L, 2L -> 0L) ++
Seq(0L -> -1L, -1L -> -2L, -2L -> 0L))
.toDF("src", "dst")
val g = GraphFrame.fromEdges(edges)
val v2 = g.triangleCount.setAlgorithm("approx").run()

v2.select("id", "count").collect().foreach {
case Row(id: Long, count: Long) =>
if (id == 0L) {
// Approx might have variation but for this small graph should be exact
assert(count >= 1)
} else {
assert(count >= 0)
}
case _ => throw new GraphFramesUnreachableException()
}
v2.unpersist()
} else {
cancel(s"skip for spark $sparkVersion")
}
}

test("Approximate triangle count - no triangles") {
val sparkVersion = spark.version
if (sparkVersion.substring(0, 3) >= "4.1") {
val edges = spark.createDataFrame(Seq(0L -> 1L, 1L -> 2L, 3L -> 4L)).toDF("src", "dst")
val g = GraphFrame.fromEdges(edges)
val v2 = g.triangleCount.setAlgorithm("approx").run()
v2.select("count").collect().foreach {
case Row(count: Long) => assert(count === 0)
case _ => throw new GraphFramesUnreachableException()
}
v2.unpersist()
} else {
cancel(s"skip for spark $sparkVersion")
}
}

test("Approximate triangle count - bipartite graph") {
val sparkVersion = spark.version
if (sparkVersion.substring(0, 3) >= "4.1") {
val edges =
spark.createDataFrame(Seq(0L -> 2L, 0L -> 3L, 1L -> 2L, 1L -> 3L)).toDF("src", "dst")
val g = GraphFrame.fromEdges(edges)
val v2 = g.triangleCount.setAlgorithm("approx").run()
v2.select("count").collect().foreach {
case Row(count: Long) => assert(count === 0)
case _ => throw new GraphFramesUnreachableException()
}
v2.unpersist()
} else {
cancel(s"skip for spark $sparkVersion")
}
}

test("Approximate triangle count - large lgNomEntries") {
val sparkVersion = spark.version
if (sparkVersion.substring(0, 3) >= "4.1") {
val edges = spark.createDataFrame(Seq(0L -> 1L, 1L -> 2L, 2L -> 0L)).toDF("src", "dst")
val g = GraphFrame.fromEdges(edges)
val v2 = g.triangleCount
.setAlgorithm("approx")
.setLgNomEntries(16)
.run()
v2.select("count").collect().foreach {
case Row(count: Long) => assert(count === 1)
case _ => throw new GraphFramesUnreachableException()
}
v2.unpersist()
} else {
cancel(s"skip for spark $sparkVersion")
}
}
}
29 changes: 24 additions & 5 deletions docs/src/04-user-guide/05-traversals.md
Original file line number Diff line number Diff line change
Expand Up @@ -245,15 +245,34 @@ result.select("id", "component").orderBy("component").show()

## Triangle count

Computes the number of triangles passing through each vertex.
Triangle count computes the number of triangles passing through each vertex. A triangle is a set of three vertices where each pair is connected by an edge (A is connected to B, B to C, and C to A).

---
**WARNING!**
### Performance and Use Cases

Counting triangles is a fundamental task in network analysis:
- **Clustering Coefficient**: It is used to compute the local and global clustering coefficients, which measure the degree to which nodes in a graph tend to cluster together.
- **Community Detection**: A high density of triangles often indicates the presence of a tightly knit community or "clique."
- **Spam and Fraud Detection**: In social networks and financial transactions, unusual triangle patterns can help identify botnets or money-laundering rings.

*The current implementation is based on collecting neighbor sets for vertices and compute a pairwise intersection of them. While this works for regular graphs, it will most probably fail on any kind of power-law graphs (graphs with a few very high-degree vertices) or at least will require a lot of memory for Spark Cluster. Consider edge sampling strategies before running the algorithm to get an approximate count of triangles.*
### How it works

---
The core logic of the algorithm is based on **neighborhood intersection**. For every edge (u, v) in the graph, the algorithm finds the intersection of the neighbor sets of u and v. Every common neighbor w completes a triangle with u and v.

### Algorithms and Trade-offs

GraphFrames provides two implementations with different performance characteristics:

- **Exact**: This is the default algorithm. It computes the precise intersection of adjacency lists.
- **Pros**: 100% accuracy.
- **Cons**: Extremely memory-intensive. For high-degree nodes (hubs), collecting and intersecting large neighbor sets can lead to Out-of-Memory (OOM) errors or severe skew.
- **Approximate** (Starting from Spark 4.1): This version uses **DataSketches (Theta sketches)** to estimate the size of the intersection.
- **Pros**: Highly scalable. It uses a fixed-size probabilistic structure to represent neighborhoods, dramatically reducing memory overhead and execution time.
- **Cons**: Provides an estimate rather than an exact count.

### Selection Guide

- **Sparse/Regular Graphs**: For graphs like grids, spatial meshes, or infrastructure networks where the maximum degree is relatively low and there is no power-law distribution, the **Exact** algorithm is recommended.
- **Power-Law/Scale-Free Networks**: For social networks, web graphs, or biological networks containing "hubs" (vertices with thousands or millions of connections), the **Approximate** algorithm is often the only viable choice to avoid job failures.

### Python API

Expand Down
6 changes: 5 additions & 1 deletion python/graphframes/classic/graphframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,9 +323,13 @@ def svdPlusPlus(
v = DataFrame(jdf, self._spark)
return (v, loss)

def triangleCount(self, storage_level: StorageLevel) -> DataFrame:
def triangleCount(
self, storage_level: StorageLevel, algorithm: str, log_nom_entries: int
Comment thread
SemyonSinchenko marked this conversation as resolved.
) -> DataFrame:
builder = self._jvm_graph.triangleCount()
builder.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))
builder.setAlgorithm(algorithm)
builder.setLgNomEntries(log_nom_entries)
jdf = builder.run()
return DataFrame(jdf, self._spark)

Expand Down
10 changes: 8 additions & 2 deletions python/graphframes/connect/graphframes_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -1085,7 +1085,9 @@ def plan(self, session: SparkConnectClient) -> proto.Relation:
else:
return (output.drop("loss"), -1.0)

def triangleCount(self, storage_level: StorageLevel) -> DataFrame:
def triangleCount(
Comment thread
SemyonSinchenko marked this conversation as resolved.
self, storage_level: StorageLevel, algorithm: str, log_nom_entries: int
) -> DataFrame:
@final
class TriangleCount(LogicalPlan):
def __init__(self, v: DataFrame, e: DataFrame, storage_level: StorageLevel) -> None:
Expand All @@ -1100,7 +1102,11 @@ def plan(self, session: SparkConnectClient) -> proto.Relation:
self.v, self.e, session
)
graphframes_api_call.triangle_count.CopyFrom(
pb.TriangleCount(storage_level=storage_level_to_proto(self.storage_level))
pb.TriangleCount(
storage_level=storage_level_to_proto(self.storage_level),
algorithm=algorithm,
lg_nom_entries=log_nom_entries,
)
)
plan = self._create_proto_relation()
plan.extension.Pack(graphframes_api_call)
Expand Down
Loading