Skip to content

Commit d36a5ac

Browse files
feat: Randomized Contraction CC (#776)
* implementation of new CC algorithm + tests * from comments
1 parent 8292846 commit d36a5ac

4 files changed

Lines changed: 593 additions & 3 deletions

File tree

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
1+
package org.apache.spark.sql.graphframes.expressions
2+
3+
import org.apache.spark.sql.catalyst.expressions.Expression
4+
import org.apache.spark.sql.catalyst.expressions.TernaryExpression
5+
import org.apache.spark.sql.catalyst.expressions.codegen.Block.*
6+
import org.apache.spark.sql.catalyst.expressions.codegen.CodegenContext
7+
import org.apache.spark.sql.catalyst.expressions.codegen.CodegenFallback
8+
import org.apache.spark.sql.catalyst.expressions.codegen.ExprCode
9+
import org.apache.spark.sql.types.DataType
10+
import org.apache.spark.sql.types.LongType
11+
12+
case class FiniteAXPlusB(first: Expression, second: Expression, third: Expression)
13+
extends TernaryExpression
14+
with CodegenFallback {
15+
override def dataType: DataType = LongType
16+
17+
override protected def withNewChildrenInternal(
18+
newFirst: Expression,
19+
newSecond: Expression,
20+
newThird: Expression): Expression = copy(newFirst, newSecond, newThird)
21+
22+
override protected def nullSafeEval(input1: Any, input2: Any, input3: Any): Any = {
23+
val a = input1.asInstanceOf[Long]
24+
val x = input2.asInstanceOf[Long]
25+
val b = input3.asInstanceOf[Long]
26+
27+
FiniteAXPlusB.axpb(a, x, b)
28+
}
29+
30+
override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = {
31+
val a = ctx.freshName("a")
32+
val x = ctx.freshName("x")
33+
val b = ctx.freshName("b")
34+
val r = ctx.freshName("r")
35+
36+
val aGenCode = first.genCode(ctx)
37+
val xGenCode = second.genCode(ctx)
38+
val bGenCode = third.genCode(ctx)
39+
40+
ev.copy(code = code"""
41+
${aGenCode.code}
42+
${xGenCode.code}
43+
${bGenCode.code}
44+
long $a = ${aGenCode.value};
45+
long $x = ${xGenCode.value};
46+
long $b = ${bGenCode.value};
47+
long $r = 0L;
48+
long irrpoly = 0x1bL;
49+
while ($x != 0L) {
50+
if (($x & 1L) != 0L) {
51+
$r ^= $a;
52+
}
53+
$x = ($x >>> 1) & 0x7fffffffffffffffL;
54+
if (($a & (1L << 63)) != 0L) {
55+
$a = ($a << 1) ^ irrpoly;
56+
} else {
57+
$a <<= 1;
58+
}
59+
}
60+
boolean ${ev.isNull} = false;
61+
long ${ev.value} = $r ^ $b;
62+
""")
63+
}
64+
}
65+
66+
object FiniteAXPlusB extends Serializable {
67+
def axpb(a: Long, x: Long, b: Long): Long = {
68+
var r = 0L
69+
val irrpoly = 0x1bL
70+
var currentA = a
71+
var currentX = x
72+
while (currentX != 0L) {
73+
if ((currentX & 1L) != 0L) {
74+
r ^= currentA
75+
}
76+
currentX = (currentX >>> 1) & 0x7fffffffffffffffL
77+
if ((currentA & (1L << 63)) != 0L) {
78+
currentA = (currentA << 1) ^ irrpoly
79+
} else {
80+
currentA <<= 1
81+
}
82+
}
83+
r ^ b
84+
}
85+
}

‎core/src/main/scala/org/graphframes/lib/ConnectedComponents.scala‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -90,8 +90,8 @@ object ConnectedComponents extends Logging {
9090

9191
import org.graphframes.GraphFrame.*
9292

93-
private val COMPONENT = "component"
94-
private val ORIG_ID = "orig_id"
93+
private[graphframes] val COMPONENT = "component"
94+
private[graphframes] val ORIG_ID = "orig_id"
9595
private val MIN_NBR = "min_nbr"
9696
private val CNT = "cnt"
9797
private val CHECKPOINT_NAME_PREFIX = "connected-components"
@@ -101,7 +101,7 @@ object ConnectedComponents extends Logging {
101101
* @param ee
102102
* non-bidirectional edges
103103
*/
104-
private def symmetrize(ee: DataFrame): DataFrame = {
104+
private[graphframes] def symmetrize(ee: DataFrame): DataFrame = {
105105
val EDGE = "_edge"
106106
ee.select(explode(
107107
array(struct(col(SRC), col(DST)), struct(col(DST).as(SRC), col(SRC).as(DST)))).as(EDGE))
Lines changed: 255 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,255 @@
1+
package org.graphframes.lib
2+
3+
import org.apache.hadoop.fs.Path
4+
import org.apache.spark.sql.Column
5+
import org.apache.spark.sql.DataFrame
6+
import org.apache.spark.sql.catalyst.FunctionIdentifier
7+
import org.apache.spark.sql.catalyst.expressions.Expression
8+
import org.apache.spark.sql.functions.*
9+
import org.apache.spark.sql.graphframes.expressions.FiniteAXPlusB
10+
import org.apache.spark.storage.StorageLevel
11+
import org.graphframes.GraphFrame
12+
import org.graphframes.GraphFrame.DST
13+
import org.graphframes.GraphFrame.ID
14+
import org.graphframes.GraphFrame.LONG_DST
15+
import org.graphframes.GraphFrame.LONG_ID
16+
import org.graphframes.GraphFrame.LONG_SRC
17+
import org.graphframes.GraphFrame.SRC
18+
import org.graphframes.Logging
19+
20+
import java.io.IOException
21+
import java.util.UUID
22+
import scala.collection.mutable.Stack
23+
import scala.util.Random
24+
25+
/**
26+
* Implementation of parallel connected components algorithm using randomized contraction, based
27+
* on Bögeholz, Harald, Michael Brand, and Radu-Alexandru Todor. "In-database connected component
28+
* analysis." 2020 IEEE 36th International Conference on Data Engineering (ICDE). IEEE, 2020.
29+
*
30+
* The algorithm contracts the graph iteratively using random linear functions, until no edges
31+
* remain, then reconstructs the component identifiers.
32+
*/
33+
private[graphframes] object RandomizedContraction extends Logging with Serializable {
34+
private val CHECKPOINT_NAME_PREFIX = "randomized-contraction"
35+
36+
private def prepare(graph: GraphFrame): GraphFrame = {
37+
val vertices = graph.indexedVertices
38+
.select(col(LONG_ID).as(ID))
39+
40+
val edges = graph.indexedEdges
41+
.select(col(LONG_SRC).as(SRC), col(LONG_DST).as(DST))
42+
val symmetricEdges = edges
43+
.union(edges.select(col(DST).alias(SRC), col(SRC).alias(DST)))
44+
.distinct()
45+
GraphFrame(vertices, symmetricEdges)
46+
47+
}
48+
49+
def run(
50+
inputGraph: GraphFrame,
51+
useLabelsAsComponents: Boolean,
52+
intermediateStorageLevel: StorageLevel,
53+
isGraphPrepared: Boolean): DataFrame = {
54+
val spark = inputGraph.vertices.sparkSession
55+
val sc = spark.sparkContext
56+
val runId = UUID.randomUUID().toString.takeRight(8)
57+
val logPrefix = s"[CC $runId]"
58+
59+
val checkpointDir = sc.getCheckpointDir
60+
.map { d =>
61+
new Path(d, s"$CHECKPOINT_NAME_PREFIX-$runId").toString
62+
}
63+
.getOrElse {
64+
// Spark-Connect workaround
65+
spark.conf.getOption("spark.checkpoint.dir") match {
66+
case Some(d) => new Path(d, s"$CHECKPOINT_NAME_PREFIX-$runId").toString
67+
case None =>
68+
throw new IOException(
69+
"Checkpoint directory is not set. Please set it first using sc.setCheckpointDir()" +
70+
"or by specifying the conf 'spark.checkpoint.dir'.")
71+
}
72+
}
73+
logInfo(s"$logPrefix Using $checkpointDir for storing intermediate tables.")
74+
75+
val functionRegistry = spark.sessionState.functionRegistry
76+
functionRegistry.registerFunction(
77+
FunctionIdentifier("_axpb"),
78+
(children: Seq[Expression]) => FiniteAXPlusB(children(0), children(1), children(2)),
79+
"scala_udf")
80+
81+
val random = new Random()
82+
random.setSeed(42L)
83+
val stackA = Stack.empty[Long]
84+
val stackB = Stack.empty[Long]
85+
var iter = 0
86+
87+
def tableName(iter: Int): String = s"${checkpointDir}/ccreps-${iter}"
88+
89+
val graph = if (isGraphPrepared) {
90+
inputGraph
91+
} else {
92+
prepare(inputGraph)
93+
}
94+
95+
var edges =
96+
graph.edges.select(SRC, DST).persist(intermediateStorageLevel)
97+
98+
def axpb(a: Long, x: Column, b: Long): Column = call_function("_axpb", lit(a), x, lit(b))
99+
100+
try {
101+
var rA = 0L
102+
var graphSize = edges.count()
103+
var ccRepresentatives: DataFrame = null
104+
105+
// "no edges graph"
106+
if (graphSize == 0L) {
107+
val result = inputGraph.vertices
108+
.select(col(ID), col(ID).alias(ConnectedComponents.COMPONENT))
109+
.persist(intermediateStorageLevel)
110+
result.count()
111+
edges.unpersist()
112+
113+
return result
114+
}
115+
116+
while (graphSize > 0) {
117+
logInfo(s"iteration ${iter}, edges left ${graphSize}")
118+
iter += 1
119+
rA = 0L
120+
while (rA == 0L) {
121+
rA = random.nextLong()
122+
}
123+
val rB = random.nextLong()
124+
stackA.push(rA)
125+
stackB.push(rB)
126+
127+
ccRepresentatives = edges
128+
.groupBy(SRC)
129+
.agg(min(axpb(rA, col(DST), rB)).alias("rep"))
130+
.select(col(SRC).alias("v"), least(axpb(rA, col(SRC), rB), col("rep")).alias("rep"))
131+
132+
// "free" checkpointing
133+
ccRepresentatives.write.parquet(tableName(iter))
134+
ccRepresentatives = spark.read.parquet(tableName(iter))
135+
136+
val edges2 = edges
137+
.join(ccRepresentatives, col(SRC) === col("v"))
138+
.select(col("rep").alias(SRC), col(DST))
139+
140+
// save ref to unpersist
141+
val oldEdges = edges
142+
143+
edges = edges2
144+
.alias("e")
145+
.join(
146+
ccRepresentatives.alias("r2"),
147+
col(s"e.$DST") === col("r2.v") &&
148+
col(s"e.$SRC") =!= col("r2.rep"))
149+
.select(col(s"e.$SRC").alias(SRC), col("r2.rep").alias(DST))
150+
.distinct()
151+
.persist(intermediateStorageLevel)
152+
153+
graphSize = edges.count()
154+
oldEdges.unpersist()
155+
}
156+
157+
logInfo(s"graph was successfully contracted for $iter iterations")
158+
logInfo("start reverse tranformation")
159+
160+
var accA = 1L
161+
var accB = 0L
162+
163+
while (iter > 1) {
164+
iter -= 1
165+
val poppedA = stackA.pop()
166+
val poppedB = stackB.pop()
167+
168+
val oldAccA = accA
169+
accA = FiniteAXPlusB.axpb(oldAccA, poppedA, 0L)
170+
accB = FiniteAXPlusB.axpb(oldAccA, poppedB, accB)
171+
172+
val ccRepsR = tableName(iter)
173+
val ccRepsR1 = tableName(iter + 1)
174+
175+
val result = spark.read
176+
.parquet(ccRepsR)
177+
.alias("r1")
178+
.join(
179+
spark.read.parquet(ccRepsR1).alias("r2"),
180+
col("r1.rep") === col("r2.v"),
181+
"left_outer")
182+
.select(
183+
col("r1.v"),
184+
coalesce(col("r2.rep"), axpb(accA, col("r1.rep"), accB)).alias("rep"))
185+
.persist(intermediateStorageLevel)
186+
187+
result.write.mode("overwrite").parquet(ccRepsR)
188+
val oldPath = new Path(ccRepsR1)
189+
val fs = oldPath.getFileSystem(sc.hadoopConfiguration)
190+
191+
if (fs.exists(oldPath)) {
192+
fs.delete(oldPath, true)
193+
}
194+
}
195+
196+
val finalReps = spark.read
197+
.parquet(tableName(1))
198+
.select(col("v").alias(ID), col("rep").alias(ConnectedComponents.COMPONENT))
199+
200+
val outputComponents = if (useLabelsAsComponents && (!inputGraph.hasIntegralIdType)) {
201+
val labels = inputGraph.indexedVertices
202+
.withColumnRenamed(ID, ConnectedComponents.ORIG_ID)
203+
.join(finalReps, col(ID) === col(LONG_ID))
204+
.groupBy(ConnectedComponents.COMPONENT)
205+
.agg(min(ConnectedComponents.ORIG_ID).alias("new_component"))
206+
207+
inputGraph.indexedVertices
208+
.withColumnRenamed(ID, ConnectedComponents.ORIG_ID)
209+
.join(finalReps, col(ID) === col(LONG_ID), "left")
210+
.join(labels, ConnectedComponents.COMPONENT, "left")
211+
.select(
212+
col(ConnectedComponents.ORIG_ID).alias(ID),
213+
coalesce(col("new_component"), col(ConnectedComponents.ORIG_ID))
214+
.alias(ConnectedComponents.COMPONENT))
215+
} else if (useLabelsAsComponents) {
216+
val labels =
217+
finalReps.groupBy(ConnectedComponents.COMPONENT).agg(min(ID).alias("new_component"))
218+
inputGraph.vertices
219+
.join(finalReps, ID, "left")
220+
.join(labels, ConnectedComponents.COMPONENT, "left")
221+
.select(
222+
col(ID),
223+
coalesce(col("new_component"), col(ID))
224+
.alias(ConnectedComponents.COMPONENT))
225+
} else {
226+
inputGraph.vertices
227+
.join(finalReps, ID, "left")
228+
.select(
229+
col(ID),
230+
coalesce(col(ConnectedComponents.COMPONENT), col(ID))
231+
.alias(ConnectedComponents.COMPONENT))
232+
}
233+
234+
outputComponents.persist(intermediateStorageLevel)
235+
// materialize to be able to clean up everything
236+
outputComponents.count()
237+
238+
// clean-up
239+
edges.unpersist()
240+
val chDirPath = new Path(checkpointDir)
241+
val fs = chDirPath.getFileSystem(sc.hadoopConfiguration)
242+
if (fs.exists(chDirPath)) {
243+
fs.delete(chDirPath, true)
244+
}
245+
246+
outputComponents
247+
} finally {
248+
val dereg = functionRegistry.dropFunction(FunctionIdentifier("_axpb"))
249+
if (!dereg) {
250+
logWarn(
251+
"graphframes faced an internal error and was not able to de-register function _axpb; Spark' functionRegistry is in a bad state")
252+
}
253+
}
254+
}
255+
}

0 commit comments

Comments
 (0)