Skip to content

Commit 6a0f34c

Browse files
feat: random walks and embeddings (#752)
* edges sampling API (scala) * add seed to z-estimation * wip * WIP * scalfix * docstrings to RandomWalkBase and RandomWalkWithRestart - Added Scala-style docstrings to all classes, traits, methods, and fields - Improved documentation for random walk algorithms and configurations * Fix RandomWalk implementation bugs and add example - Correct element_at index from 0 to 1 for 1-based Spark SQL arrays - Fix walk array construction by appending nextNode instead of currVisitingVertex - Add null handling for nodes with no outgoing neighbors in restart logic - Add comprehensive Scala docstrings to RandomWalkBase and RandomWalkWithRestart - Create RWExample.scala demonstrating RandomWalkWithRestart on LDBC datasets ... * Add Word2VecHashingTrick implementation for graph embeddings This commit introduces a new Word2Vec-based embedding method using the hashing trick to handle large vocabularies efficiently in graph frames, particularly for random walk sequences. It includes configurable parameters like number of hashing functions, max features, and standard W2V settings, with comprehensive Scaladoc for public APIs. - Added core/src/main/scala/org/graphframes/embeddings/Word2VecHashingTrick.scala: New class implementing hashing trick by applying multiple Murmur3 hash functions and modulo to map features to a fixed-size space, reducing collisions and memory usage. It trains a W2V model on expanded sequences and provides a companion model class for vector retrieval via averaging hashed embeddings. Setters include docstrings explaining trade-offs (e.g., more hashes improve quality but multiply dataset size). - Modified core/src/main/scala/org/graphframes/examples/RWExample.scala: Updated main method to accept a single file path argument for edge loading instead of downloading LDBC datasets, simplifying usage for local files. Replaced vertex loading with direct derivation from edges for consistency and reduced I/O. - Modified core/src/main/scala/org/graphframes/exceptions.scala: Added GraphFramesW2VException class to handle W2V-specific errors, such as unsupported input types in hashing. * Implement reservoir sampling for neighbor selection in random walks Replace collect_set + shuffle + slice with ReservoirSamplingAgg UDAF for efficient sampling of up to maxNbrs neighbors per vertex. This improves performance by avoiding full neighbor list aggregation and shuffling, especially beneficial for high-degree vertices. - Add ReservoirSamplingAgg trait: generic aggregator using reservoir sampling algorithm, supporting merge operations for distributed computation. - Handle various vertex ID types (String, Short, Byte, Int, Long) with appropriate encoders. - Raise GraphFramesUnsupportedVertexTypeException for unsupported types. - Add comprehensive test suite covering reduce, merge, and finish operations with edge cases and fixed seeds for determinism. Modified files: - .gitignore: Ignore Emacs temp files for cleaner diffs. - core/src/main/scala/org/graphframes/exceptions.scala: New exception class. - core/src/main/scala/org/graphframes/rw/RandomWalkBase.scala: Integrate ReservoirSamplingAgg in prepareGraph method. New files: - core/src/main/scala/org/apache/spark/sql/graphframes/expressions/ReservoirSamplingAgg.scala - core/src/test/scala/org/apache/spark/sql/graphframes/expressions/ReservoirSamplingAggSuite.scala * fix scalastyle? * fix reservoir * add hash2vec delete wrong implementation of w2v + hashing * fixes * docstrings + scalfix * remove sampling as not needed * fixes in build and code * Big update - replace Reservoir sampling by KMinSampling - add L2norm to Hash2vec - add an optional convolution step to RW embeddings - small updates and performance fixes * Tests and updates * workaround scala 2.13 deprecation of Searching.search * Fixes Fix the problem `sun.security.action` access * Fix access * fallback to java Serialization * Fix some bugs Tested on an end2end case * Sampling Convolution tests and docstrings * docstrings for RW Emebddings * Python API * hash2vec tests * Explicit types * hash2vec and random walks with restart tests * performance * ignore unused nowarn spark 4 vs spark 3 * performance + cached walks support + continous mode * fix rw and update the branch * fix * protobuf & connect * public API for embeddings and small refactoring of methods * initial Py API for embeddings * fix * tests + docs - python tests - docs (I use AI to generate but checked by myself and fixed) - fix in python APIs - small changes * fix connect tests * decrease GC pressure * further optimizations - reduce GC pressure from UTF8String - avoid hash recomputations * refactor: optimize Hash2Vec string hashing and partitioning logic * refactor: replace generic hash function with type-specific implementations and add hash caching * refactor: simplify hash function logic and improve performance in Hash2Vec * refactor: inline generic processPartitionGeneric into specialized String and Long methods * test: add tests for PagedMatrixDouble helper covering page extension, add, and getVector * refactor: replace unsafe hash functions with MurmurHash3 and optimize memory usage with paged matrix * refactor: update processStringPartition to use PagedMatrixDouble for memory efficiency * docs: add internal documentation for PagedMatrixDouble explaining memory layout and GC benefits * chore: remove commented helper section from Hash2Vec * refactor: replace case class with class for PagedMatrixDouble and remove Scala version-specific compat files * fix: correct LongMap type parameter in Hash2Vec vocabIndex initialization * refactor: reduce PAGE_BITS from 16 to 12 and update related constants and tests * feat: add max vectors per partition limit and batched processing for memory control * refactor: process long partitions in batches respecting max vectors limit * test: add Hash2Vec tests for co-occurrence patterns and cosine similarity validation * chore: clean up code formatting and improve test readability in Hash2Vec * fix: correct typo in error message from 'gor' to 'got' in Hash2Vec exception * docs: add docstrings for Hash2Vec setters setDoNormalization and setMaxVectorsPerPartition * fix: correct typo in NOTICE and KMinSampling, update Hash2Vec defaults, add null check and memory management in RandomWalkEmbeddings, simplify RandomWalkBase seed setting, and fix randomness in RandomWalkWithRestart * fix: skip seeds for previous batches to maintain consistency when starting from non-first batch * fix: add overwrite mode when writing batch results to allow re-running with same walkID * feat: add cleanUp method to remove temporary files for a walk ID using Hadoop FS * docs: improve documentation for cleanUp method in RandomWalkBase trait * test: add cleanUp call in RandomWalkWithRestart test * test: move walks execution inside try block in RandomWalkWithRestartSuite * refactor: set walkID default to UUID and remove redundant runID variable * docs: update comment to clarify walkID retrieval method behavior * refactor: rename walkID to runID for clarity in random walk operations * refactor: remove runId parameter from cleanUp method in RandomWalkBase * refactor: move cleanUp method to companion object with parameters and update instance method * refactor: improve documentation formatting and remove redundant log setup in cleanUp * test: verify temporary files are deleted after RandomWalkWithRestart test * refactor: use numBatches variable and fix run path in RandomWalkWithRestartSuite test * test: add test for RandomWalkWithRestart resuming from middle iteration * style: Format RandomWalkWithRestartSuite with consistent spacing and indentation * refactor: use getSeq instead of getAs[Seq[String]] in RandomWalkWithRestartSuite * fix: correct typo in error message and parameter name for gaussian sigma * feat: add clean-up option for temporary random walk files in embeddings pipeline * docs: improve formatting and consistency in graph-ml documentation tables and sections * feat: add clean_up_after_run field to RandomWalkEmbeddings proto definition * chore: add clean_up_after_run parameter to _RandomWalksEmbeddingsParameters
1 parent 12ad0c8 commit 6a0f34c

32 files changed

Lines changed: 3792 additions & 163 deletions

File tree

‎.gitignore‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -79,3 +79,8 @@ spark-*
7979

8080
# Zed
8181
.zed
82+
83+
# Emacs
84+
.dir-locals.el
85+
*~
86+
.aider*

‎.pre-commit-config.yaml‎

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,15 +21,14 @@ repos:
2121

2222
- id: scalafmt
2323
name: scalafmt
24-
entry: build/sbt scalafmtCheckAll
24+
entry: build/sbt scalafmtAll
2525
language: system
2626
types: [scala]
2727
pass_filenames: false
2828

2929
- id: scalafix
3030
name: scalafix
31-
entry: build/sbt "scalafixAll --check"
31+
entry: build/sbt scalafixAll
3232
language: system
3333
types: [scala]
3434
pass_filenames: false
35-

‎NOTICE‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,3 +8,10 @@ Copyright 2014-2025 The Apache Software Foundation.
88

99
This product includes software developed at
1010
The Apache Software Foundation (http://www.apache.org/).
11+
12+
Part of the code of the project is heavily inspired or copied from the Apache Spark ML project, which are licensed under the Apache Software License, Version 2.0. The Apache Spark project has the following NOTICE:
13+
Apache Spark
14+
Copyright 2014 and onwards The Apache Software Foundation.
15+
16+
This product includes software developed at
17+
The Apache Software Foundation (http://www.apache.org/).

‎build.sbt‎

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -96,10 +96,12 @@ lazy val commonSetting = Seq(
9696
"--add-opens=java.base/java.lang=ALL-UNNAMED",
9797
"--add-opens=java.base/java.nio=ALL-UNNAMED",
9898
"--add-opens=java.base/java.lang.invoke=ALL-UNNAMED",
99-
"--add-opens=java.base/java.util=ALL-UNNAMED"),
99+
"--add-opens=java.base/java.util=ALL-UNNAMED",
100+
"--add-opens=java.base/sun.security.action=ALL-UNNAMED",
101+
"--add-opens=java.base/java.io=ALL-UNNAMED"),
100102

101103
// Scalac options
102-
tpolecatScalacOptions ++= Set(
104+
Compile / tpolecatScalacOptions ++= Set(
103105
ScalacOptions.lint,
104106
ScalacOptions.deprecation,
105107
ScalacOptions.warnDeadCode,
@@ -111,7 +113,10 @@ lazy val commonSetting = Seq(
111113
ScalacOptions.warnUnusedNoWarn,
112114
ScalacOptions.source3,
113115
ScalacOptions.fatalWarnings),
114-
tpolecatExcludeOptions ++= Set(ScalacOptions.warnNonUnitStatement),
116+
Compile / tpolecatExcludeOptions ++= Set(
117+
ScalacOptions.warnNonUnitStatement,
118+
ScalacOptions.privateWarnUnusedNoWarn,
119+
ScalacOptions.warnUnusedNoWarn),
115120
Test / tpolecatExcludeOptions ++= Set(
116121
ScalacOptions.warnValueDiscard,
117122
ScalacOptions.warnUnusedLocals,
@@ -122,8 +127,7 @@ lazy val commonSetting = Seq(
122127
ScalacOptions.warnNumericWiden,
123128
ScalacOptions.privateWarnNumericWiden,
124129
ScalacOptions.warnUnusedNoWarn,
125-
ScalacOptions.privateWarnUnusedNoWarn,
126-
))
130+
ScalacOptions.privateWarnUnusedNoWarn))
127131

128132
lazy val graphx = (project in file("graphx"))
129133
.settings(
@@ -136,7 +140,9 @@ lazy val graphx = (project in file("graphx"))
136140
// for scala 2.13 we should mark "unused" class tags by @nowarn,
137141
// for scala 2.12 we shouldn't
138142
// the only way at the moment is to not check unused @nowarn for GraphX
139-
tpolecatExcludeOptions ++= Set(ScalacOptions.warnUnusedNoWarn, ScalacOptions.privateWarnUnusedNoWarn),
143+
tpolecatExcludeOptions ++= Set(
144+
ScalacOptions.warnUnusedNoWarn,
145+
ScalacOptions.privateWarnUnusedNoWarn),
140146

141147
// Global settings
142148
Global / concurrentRestrictions := Seq(Tags.limitAll(1)),

‎connect/src/main/protobuf/graphframes.proto‎

Lines changed: 37 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,8 +35,9 @@ message GraphFramesAPI {
3535
SVDPlusPlus svd_plus_plus = 18;
3636
TriangleCount triangle_count = 19;
3737
Triplets triplets = 20;
38-
MaximalIndependentSet mis = 22;
3938
KCore kcore = 21;
39+
MaximalIndependentSet mis = 22;
40+
RandomWalkEmbeddings rw_embeddings = 23;
4041
}
4142
}
4243

@@ -208,3 +209,38 @@ message KCore {
208209
int32 checkpoint_interval = 2;
209210
optional StorageLevel storage_level = 3;
210211
}
212+
213+
message RandomWalkEmbeddings {
214+
bool use_edge_direction = 1;
215+
string rw_model = 2;
216+
int32 rw_max_nbrs = 3;
217+
int32 rw_num_walks_per_node = 4;
218+
int32 rw_batch_size = 5;
219+
int32 rw_num_batches = 6;
220+
int64 rw_seed = 7;
221+
double rw_restart_probability = 8;
222+
string rw_temporary_prefix = 9;
223+
string rw_cached_walks = 10;
224+
string sequence_model = 11;
225+
int32 hash2vec_context_size = 12;
226+
int32 hash2vec_num_partitions = 13;
227+
int32 hash2vec_embeddings_dim = 14;
228+
string hash2vec_decay_function = 15;
229+
double hash2vec_gaussian_sigma = 16;
230+
int32 hash2vec_hashing_seed = 17;
231+
int32 hash2vec_sign_seed = 18;
232+
bool hash2vec_do_l2_norm = 19;
233+
bool hash2vec_safe_l2 = 20;
234+
int32 word2vec_max_iter = 21;
235+
int32 word2vec_embeddings_dim = 22;
236+
int32 word2vec_window_size = 23;
237+
int32 word2vec_num_partitions = 24;
238+
int32 word2vec_min_count = 25;
239+
int32 word2vec_max_sentence_length = 26;
240+
int64 word2vec_seed = 27;
241+
double word2vec_step_size = 28;
242+
bool aggregate_neighbors = 29;
243+
int32 aggregate_neighbors_max_nbrs = 30;
244+
int64 aggregate_neighbors_seed = 31;
245+
bool clean_up_after_run = 32;
246+
}

‎connect/src/main/scala/org/apache/spark/sql/graphframes/GraphFramesConnectUtils.scala‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ import org.apache.spark.storage.StorageLevel
1212
import org.graphframes.GraphFrame
1313
import org.graphframes.GraphFramesUnreachableException
1414
import org.graphframes.connect.proto
15+
import org.graphframes.embeddings.RandomWalkEmbeddings
1516

1617
import scala.jdk.CollectionConverters.*
1718

@@ -454,6 +455,44 @@ object GraphFramesConnectUtils {
454455

455456
kCoreBuilder.run()
456457
}
458+
case proto.GraphFramesAPI.MethodCase.RW_EMBEDDINGS => {
459+
val message = apiMessage.getRwEmbeddings()
460+
461+
RandomWalkEmbeddings.pythonAPI(
462+
graph = graphFrame,
463+
useEdgeDirection = message.getUseEdgeDirection(),
464+
rwModel = message.getRwModel(),
465+
rwMaxNbrs = message.getRwMaxNbrs(),
466+
rwNumWalksPerNode = message.getRwNumWalksPerNode(),
467+
rwBatchSize = message.getRwBatchSize(),
468+
rwNumBatches = message.getRwNumBatches(),
469+
rwSeed = message.getRwSeed(),
470+
rwRestartProbability = message.getRwRestartProbability(),
471+
rwTemporaryPrefix = message.getRwTemporaryPrefix(),
472+
rwCachedWalks = message.getRwCachedWalks(),
473+
sequenceModel = message.getSequenceModel(),
474+
hash2vecContextSize = message.getHash2VecContextSize(),
475+
hash2vecNumPartitions = message.getHash2VecNumPartitions(),
476+
hash2vecEmbeddingsDim = message.getHash2VecEmbeddingsDim(),
477+
hash2vecDecayFunction = message.getHash2VecDecayFunction(),
478+
hash2vecGaussianSigma = message.getHash2VecGaussianSigma(),
479+
hash2vecHashingSeed = message.getHash2VecHashingSeed(),
480+
hash2vecSignSeed = message.getHash2VecSignSeed(),
481+
hash2vecDoL2Norm = message.getHash2VecDoL2Norm(),
482+
hash2vecSafeL2 = message.getHash2VecSafeL2(),
483+
word2vecMaxIter = message.getWord2VecMaxIter(),
484+
word2vecEmbeddingsDim = message.getWord2VecEmbeddingsDim(),
485+
word2vecWindowSize = message.getWord2VecWindowSize(),
486+
word2vecNumPartitions = message.getWord2VecNumPartitions(),
487+
word2vecMinCount = message.getWord2VecMinCount(),
488+
word2vecMaxSentenceLength = message.getWord2VecMaxSentenceLength(),
489+
word2vecSeed = message.getWord2VecSeed(),
490+
word2vecStepSize = message.getWord2VecStepSize(),
491+
aggregateNeighbors = message.getAggregateNeighbors(),
492+
aggregateNeighborsMaxNbrs = message.getAggregateNeighborsMaxNbrs(),
493+
aggregateNeighborsSeed = message.getAggregateNeighborsSeed(),
494+
cleanUpAfterRun = message.getCleanUpAfterRun())
495+
}
457496
case _ => throw new GraphFramesUnreachableException() // Unreachable
458497
}
459498
}
Lines changed: 165 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,165 @@
1+
package org.apache.spark.sql.graphframes.expressions
2+
3+
import org.apache.spark.sql.Encoder
4+
import org.apache.spark.sql.Encoders
5+
import org.apache.spark.sql.Row
6+
import org.apache.spark.sql.SparkSession
7+
import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder
8+
import org.apache.spark.sql.expressions.Aggregator
9+
import org.apache.spark.sql.expressions.UserDefinedFunction
10+
import org.apache.spark.sql.functions.udaf
11+
import org.apache.spark.sql.types.*
12+
import org.apache.spark.sql.types.DataType
13+
import org.graphframes.GraphFramesUnsupportedVertexTypeException
14+
15+
import scala.annotation.nowarn
16+
import scala.reflect.ClassTag
17+
import scala.reflect.runtime.universe.TypeTag
18+
19+
case class KMinAccum[T](values: Array[T], weights: Array[Long], var cnt: Int) extends Serializable
20+
21+
case class KMinSampling[T: ClassTag](size: Int)(implicit
22+
@nowarn tag: TypeTag[T],
23+
ord: Ordering[T])
24+
extends Aggregator[Row, KMinAccum[T], Seq[T]]
25+
with Serializable {
26+
27+
override def zero: KMinAccum[T] = KMinAccum(Array.ofDim[T](size), Array.ofDim[Long](size), 0)
28+
29+
override def reduce(b: KMinAccum[T], a: Row): KMinAccum[T] = {
30+
val newWeight = a.getLong(1)
31+
val newValue = a.getAs[T](0)
32+
// fast-path: buffer is already full of "strong" elements
33+
// the case of "influencer" vertex
34+
if (b.cnt == size) {
35+
val lastWeight = b.weights.last
36+
if ((lastWeight < newWeight) || ((lastWeight == newWeight) && (ord.compare(
37+
newValue,
38+
b.values.last) >= 0))) {
39+
return b
40+
}
41+
}
42+
43+
// slow-path: custom binary search for (Weight, Value)
44+
// We want to find the first index where (b.w, b.v) > (newWeight, newValue)
45+
var low = 0
46+
var high = b.cnt - 1
47+
var idx = b.cnt // Default insertion point is at the end
48+
49+
while (low <= high) {
50+
val mid = (low + high) / 2
51+
val midWeight = b.weights(mid)
52+
53+
// Compare (midWeight, midValue) vs (newWeight, newValue)
54+
val res =
55+
if (midWeight < newWeight) -1
56+
else if (midWeight > newWeight) 1
57+
else ord.compare(b.values(mid), newValue)
58+
59+
if (res <= 0) {
60+
// mid is smaller or equal: we must insert after mid
61+
low = mid + 1
62+
} else {
63+
// mid is larger: potential insertion point here
64+
idx = mid
65+
high = mid - 1
66+
}
67+
}
68+
69+
if (idx < size) {
70+
val newCount = math.min(b.cnt + 1, size)
71+
if (idx < newCount - 1) {
72+
// shift to the right if needed
73+
System.arraycopy(b.weights, idx, b.weights, idx + 1, newCount - idx - 1)
74+
System.arraycopy(b.values, idx, b.values, idx + 1, newCount - idx - 1)
75+
}
76+
77+
b.weights(idx) = newWeight
78+
b.values(idx) = newValue
79+
b.cnt = newCount
80+
}
81+
82+
b
83+
}
84+
85+
override def merge(b1: KMinAccum[T], b2: KMinAccum[T]): KMinAccum[T] = {
86+
87+
if (b1.cnt == 0) {
88+
return b2
89+
}
90+
91+
if (b2.cnt == 0) {
92+
return b1
93+
}
94+
95+
val resultSize = math.min(b1.cnt + b2.cnt, size)
96+
val newValues = Array.ofDim[T](resultSize)
97+
val newWeights = Array.ofDim[Long](resultSize)
98+
99+
var i = 0
100+
var j = 0
101+
var r = 0
102+
103+
while (r < resultSize) {
104+
val useLeft = if (i >= b1.cnt) {
105+
false
106+
} else if (j >= b2.cnt) {
107+
true
108+
} else {
109+
val wLeft = b1.weights(i)
110+
val wRight = b2.weights(j)
111+
112+
if (wLeft < wRight) {
113+
true
114+
} else if (wLeft > wRight) {
115+
false
116+
} else {
117+
ord.compare(b1.values(i), b2.values(j)) <= 0
118+
}
119+
}
120+
121+
if (useLeft) {
122+
newWeights(r) = b1.weights(i)
123+
newValues(r) = b1.values(i)
124+
i += 1
125+
} else {
126+
newWeights(r) = b2.weights(j)
127+
newValues(r) = b2.values(j)
128+
j += 1
129+
}
130+
131+
r += 1
132+
}
133+
134+
KMinAccum(newValues, newWeights, resultSize)
135+
}
136+
137+
override def finish(reduction: KMinAccum[T]): Seq[T] =
138+
reduction.values.slice(0, reduction.cnt).toSeq
139+
// TODO: replace by Kryo after 4.0.2 is released, see SPARK-52819
140+
override def bufferEncoder: Encoder[KMinAccum[T]] = Encoders.product
141+
override def outputEncoder: Encoder[Seq[T]] = ExpressionEncoder[Seq[T]]()
142+
}
143+
144+
object KMinSampling extends Serializable {
145+
def getEncoder(spark: SparkSession, dataType: DataType, colNames: Seq[String]): Encoder[Row] = {
146+
// That is very stupid way actually. But it is the only way with public API
147+
spark
148+
.createDataFrame(
149+
java.util.List.of[Row](),
150+
StructType(
151+
StructField(colNames(0), dataType) :: StructField(colNames(1), LongType) :: Nil))
152+
.encoder
153+
}
154+
155+
def fromSparkType(dataType: DataType, size: Int, encoder: Encoder[Row]): UserDefinedFunction = {
156+
dataType match {
157+
case StringType => udaf(KMinSampling[java.lang.String](size), encoder)
158+
case ShortType => udaf(KMinSampling[java.lang.Short](size), encoder)
159+
case ByteType => udaf(KMinSampling[java.lang.Byte](size), encoder)
160+
case IntegerType => udaf(KMinSampling[java.lang.Integer](size), encoder)
161+
case LongType => udaf(KMinSampling[java.lang.Long](size), encoder)
162+
case _ => throw new GraphFramesUnsupportedVertexTypeException("unsupported vertex type")
163+
}
164+
}
165+
}

0 commit comments

Comments
 (0)