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
Prev Previous commit
Next Next commit
Merge main into 785-aggregate-nbrs: integrate sampling, embeddings, r…
…andom walks, and PG changes

Resolved merge conflicts and integrated updates from graphframes/main into branch 785-aggregate-nbrs.

Key changes:
- Added sampling and convolution primitives (KMinSampling, SamplingConvolution) and related tests.
- Introduced embeddings and random-walk features: Hash2Vec, RandomWalkEmbeddings, RandomWalkBase, RandomWalkWithRestart, and example runners.
- Added new library components and algorithms: TwoPhase, Updated ConnectedComponents/KCore/Pregel/RandomizedContraction logic.
- Added benchmarks and benchmark infrastructure for various algorithms.
- Python API improvements: new property-graph package (python/graphframes/pg), internal utilities, updated client and protobuf bindings.
- Build, docs, and packaging updates: build.sbt, docs (including graph-ml page), NOTICE, AGENTS.md, and pre-commit config.
- Updated Spark shims and connect utilities to support the new features.
- Added and updated tests across Scala and Python to cover the new functionality.

All conflicts were fixed and changes staged for commit.
  • Loading branch information
SemyonSinchenko committed Mar 26, 2026
commit 56736bbf43dadec410bb2ed2fb2aaf5dce33454e
11 changes: 11 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -79,4 +79,15 @@ spark-*

# Zed
.zed

# Emacs
.dir-locals.el
*~

# AI
.claude
.opencode
.qwen
.cursor
openspec
.aider*
39 changes: 38 additions & 1 deletion connect/src/main/protobuf/graphframes.proto
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,9 @@ message GraphFramesAPI {
TriangleCount triangle_count = 19;
Triplets triplets = 20;
KCore kcore = 21;
AggregateNeighbors aggregate_neighbors = 23;
MaximalIndependentSet mis = 22;
RandomWalkEmbeddings rw_embeddings = 23;
AggregateNeighbors aggregate_neighbors = 24;
}
}

Expand Down Expand Up @@ -239,3 +241,38 @@ message AggregateNeighbors {
// Optional storage level for intermediate results
optional StorageLevel storage_level = 14;
}

message RandomWalkEmbeddings {
bool use_edge_direction = 1;
string rw_model = 2;
int32 rw_max_nbrs = 3;
int32 rw_num_walks_per_node = 4;
int32 rw_batch_size = 5;
int32 rw_num_batches = 6;
int64 rw_seed = 7;
double rw_restart_probability = 8;
string rw_temporary_prefix = 9;
string rw_cached_walks = 10;
string sequence_model = 11;
int32 hash2vec_context_size = 12;
int32 hash2vec_num_partitions = 13;
int32 hash2vec_embeddings_dim = 14;
string hash2vec_decay_function = 15;
double hash2vec_gaussian_sigma = 16;
int32 hash2vec_hashing_seed = 17;
int32 hash2vec_sign_seed = 18;
bool hash2vec_do_l2_norm = 19;
bool hash2vec_safe_l2 = 20;
int32 word2vec_max_iter = 21;
int32 word2vec_embeddings_dim = 22;
int32 word2vec_window_size = 23;
int32 word2vec_num_partitions = 24;
int32 word2vec_min_count = 25;
int32 word2vec_max_sentence_length = 26;
int64 word2vec_seed = 27;
double word2vec_step_size = 28;
bool aggregate_neighbors = 29;
int32 aggregate_neighbors_max_nbrs = 30;
int64 aggregate_neighbors_seed = 31;
bool clean_up_after_run = 32;
}
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,45 @@ object GraphFramesConnectUtils {

anBuilder.run()
}

case proto.GraphFramesAPI.MethodCase.RW_EMBEDDINGS => {
val message = apiMessage.getRwEmbeddings()

RandomWalkEmbeddings.pythonAPI(
graph = graphFrame,
useEdgeDirection = message.getUseEdgeDirection(),
rwModel = message.getRwModel(),
rwMaxNbrs = message.getRwMaxNbrs(),
rwNumWalksPerNode = message.getRwNumWalksPerNode(),
rwBatchSize = message.getRwBatchSize(),
rwNumBatches = message.getRwNumBatches(),
rwSeed = message.getRwSeed(),
rwRestartProbability = message.getRwRestartProbability(),
rwTemporaryPrefix = message.getRwTemporaryPrefix(),
rwCachedWalks = message.getRwCachedWalks(),
sequenceModel = message.getSequenceModel(),
hash2vecContextSize = message.getHash2VecContextSize(),
hash2vecNumPartitions = message.getHash2VecNumPartitions(),
hash2vecEmbeddingsDim = message.getHash2VecEmbeddingsDim(),
hash2vecDecayFunction = message.getHash2VecDecayFunction(),
hash2vecGaussianSigma = message.getHash2VecGaussianSigma(),
hash2vecHashingSeed = message.getHash2VecHashingSeed(),
hash2vecSignSeed = message.getHash2VecSignSeed(),
hash2vecDoL2Norm = message.getHash2VecDoL2Norm(),
hash2vecSafeL2 = message.getHash2VecSafeL2(),
word2vecMaxIter = message.getWord2VecMaxIter(),
word2vecEmbeddingsDim = message.getWord2VecEmbeddingsDim(),
word2vecWindowSize = message.getWord2VecWindowSize(),
word2vecNumPartitions = message.getWord2VecNumPartitions(),
word2vecMinCount = message.getWord2VecMinCount(),
word2vecMaxSentenceLength = message.getWord2VecMaxSentenceLength(),
word2vecSeed = message.getWord2VecSeed(),
word2vecStepSize = message.getWord2VecStepSize(),
aggregateNeighbors = message.getAggregateNeighbors(),
aggregateNeighborsMaxNbrs = message.getAggregateNeighborsMaxNbrs(),
aggregateNeighborsSeed = message.getAggregateNeighborsSeed(),
cleanUpAfterRun = message.getCleanUpAfterRun())
}
case _ => throw new GraphFramesUnreachableException() // Unreachable
}
}
Expand Down
45 changes: 45 additions & 0 deletions python/graphframes/classic/graphframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -482,4 +482,49 @@ def aggregate_neighbors(
builder.setIntermediateStorageLevel(storage_level_to_jvm(storage_level, self._spark))

jdf = builder.run()
assert jdf is not None

return DataFrame(jdf, self._spark)

def rw_embeddings(self, params: _RandomWalksEmbeddingsParameters) -> DataFrame:
assert self._jvm is not None
j_rw_embeddings = self._jvm.org.graphframes.embeddings.RandomWalkEmbeddings
assert j_rw_embeddings is not None
jdf: JavaObject = j_rw_embeddings.pythonAPI(
self._jvm_graph,
params.use_edge_direction,
params.rw_model,
params.rw_max_nbrs,
params.rw_num_walks_per_node,
params.rw_batch_size,
params.rw_num_batches,
params.rw_seed,
params.rw_restart_probability,
params.rw_temporary_prefix,
params.rw_cached_walks,
params.sequence_model,
params.hash2vec_context_size,
params.hash2vec_num_partitions,
params.hash2vec_embeddings_dim,
params.hash2vec_decay_function,
params.hash2vec_gaussian_sigma,
params.hash2vec_hashing_seed,
params.hash2vec_sign_seed,
params.hash2vec_do_l2_norm,
params.hash2vec_safe_l2,
params.word2vec_max_iter,
params.word2vec_embeddings_dim,
params.word2vec_window_size,
params.word2vec_num_partitions,
params.word2vec_min_count,
params.word2vec_max_sentence_length,
params.word2vec_seed,
params.word2vec_step_size,
params.aggregate_neighbors,
params.aggregate_neighbors_max_nbrs,
params.aggregate_neighbors_seed,
params.clean_up_after_run,
)
assert jdf is not None

return DataFrame(jdf, self._spark)
58 changes: 58 additions & 0 deletions python/graphframes/connect/graphframes_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -1342,3 +1342,61 @@ def plan(self, session: SparkConnectClient) -> proto.Relation:
),
self._spark,
)

def rw_embeddings(self, params: _RandomWalksEmbeddingsParameters) -> DataFrame:
@final
class RWEmbeddings(LogicalPlan):
def __init__(
self, v: DataFrame, e: DataFrame, params: _RandomWalksEmbeddingsParameters
) -> None:
super().__init__(None)
self.v = v
self.e = e
self.params = params

@override
def plan(self, session: SparkConnectClient) -> proto.Relation:
graphframes_api_call = GraphFrameConnect._get_pb_api_message(
self.v, self.e, session
)
graphframes_api_call.rw_embeddings.CopyFrom(
pb.RandomWalkEmbeddings(
use_edge_direction=self.params.use_edge_direction,
rw_model=self.params.rw_model,
rw_max_nbrs=self.params.rw_max_nbrs,
rw_num_walks_per_node=self.params.rw_num_walks_per_node,
rw_batch_size=self.params.rw_batch_size,
rw_num_batches=self.params.rw_num_batches,
rw_seed=self.params.rw_seed,
rw_restart_probability=self.params.rw_restart_probability,
rw_temporary_prefix=self.params.rw_temporary_prefix,
rw_cached_walks=self.params.rw_cached_walks,
sequence_model=self.params.sequence_model,
hash2vec_context_size=self.params.hash2vec_context_size,
hash2vec_num_partitions=self.params.hash2vec_num_partitions,
hash2vec_embeddings_dim=self.params.hash2vec_embeddings_dim,
hash2vec_decay_function=self.params.hash2vec_decay_function,
hash2vec_gaussian_sigma=self.params.hash2vec_gaussian_sigma,
hash2vec_hashing_seed=self.params.hash2vec_hashing_seed,
hash2vec_sign_seed=self.params.hash2vec_sign_seed,
hash2vec_do_l2_norm=self.params.hash2vec_do_l2_norm,
hash2vec_safe_l2=self.params.hash2vec_safe_l2,
word2vec_max_iter=self.params.word2vec_max_iter,
word2vec_embeddings_dim=self.params.word2vec_embeddings_dim,
word2vec_window_size=self.params.word2vec_window_size,
word2vec_num_partitions=self.params.word2vec_num_partitions,
word2vec_min_count=self.params.word2vec_min_count,
word2vec_max_sentence_length=self.params.word2vec_max_sentence_length,
word2vec_seed=self.params.word2vec_seed,
word2vec_step_size=self.params.word2vec_step_size,
aggregate_neighbors=self.params.aggregate_neighbors,
aggregate_neighbors_max_nbrs=self.params.aggregate_neighbors_max_nbrs,
aggregate_neighbors_seed=self.params.aggregate_neighbors_seed,
clean_up_after_run=self.params.clean_up_after_run,
)
)
plan = self._create_proto_relation()
plan.extension.Pack(graphframes_api_call)
return plan

return _dataframe_from_plan(RWEmbeddings(self._vertices, self._edges, params), self._spark)
Loading
Loading
You are viewing a condensed version of this merge commit. You can view the full changes here.