Skip to content

Commit 5db86a2

Browse files
feat: add early stopping to Pregel (#550)
* Add early stopping to Pregel * Make early stopping optional * Forgot to add a condition * From comments
1 parent 8a719fb commit 5db86a2

9 files changed

Lines changed: 194 additions & 55 deletions

File tree

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

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -111,6 +111,7 @@ message Pregel {
111111
string additional_col_name = 6;
112112
ColumnOrExpression additional_col_initial = 7;
113113
ColumnOrExpression additional_col_upd = 8;
114+
optional bool early_stopping = 9;
114115
}
115116

116117
message ShortestPaths {

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

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -174,6 +174,10 @@ object GraphFramesConnectUtils {
174174
.map(parseColumnOrExpression(_, planner))
175175
.foldLeft(pregel)((p, col) => p.sendMsgToDst(col))
176176

177+
if (pregelProto.hasEarlyStopping) {
178+
pregel = pregel.setEarlyStopping(pregelProto.getEarlyStopping)
179+
}
180+
177181
pregel.run()
178182
}
179183
case MethodCase.SHORTEST_PATHS => {

‎python/graphframes/connect/graphframe_client.py‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ def __init__(self, graph: "GraphFrameConnect") -> None:
2424
self._send_msg_to_src = []
2525
self._send_msg_to_dst = []
2626
self._agg_msg = None
27+
self._early_stopping = False
2728

2829
def setMaxIter(self, value: int) -> Self:
2930
self._max_iter = value
@@ -33,6 +34,10 @@ def setCheckpointInterval(self, value: int) -> Self:
3334
self._checkpoint_interval = value
3435
return self
3536

37+
def setEarlyStopping(self, value: bool) -> Self:
38+
self._early_stopping = value
39+
return self
40+
3641
def withVertexColumn(
3742
self,
3843
colName: str,
@@ -62,6 +67,7 @@ def __init__(
6267
self,
6368
max_iter: int,
6469
checkpoint_interval: int,
70+
early_stopping: bool,
6571
vertex_col_name: str,
6672
agg_msg: Column | str,
6773
send2dst: list[Column | str],
@@ -74,6 +80,7 @@ def __init__(
7480
super().__init__(None)
7581
self.max_iter = max_iter
7682
self.checkpoint_interval = checkpoint_interval
83+
self.early_stopping = early_stopping
7784
self.vertex_col_name = vertex_col_name
7885
self.agg_msg = agg_msg
7986
self.send2dst = send2dst
@@ -97,6 +104,7 @@ def plan(self, session: SparkConnectClient) -> proto.Relation:
97104
additional_col_name=self.vertex_col_name,
98105
additional_col_initial=make_column_or_expr(self.vertex_col_init, session),
99106
additional_col_upd=make_column_or_expr(self.vertex_col_upd, session),
107+
early_stopping=self.early_stopping,
100108
)
101109
pb_message = pb.GraphFramesAPI(
102110
vertices=dataframe_to_proto(self.vertices, session),
@@ -129,6 +137,7 @@ def plan(self, session: SparkConnectClient) -> proto.Relation:
129137
send2src=self._send_msg_to_src,
130138
vertices=self.graph._vertices,
131139
edges=self.graph._edges,
140+
early_stopping=self._early_stopping,
132141
),
133142
session=self.graph._spark,
134143
)

‎python/graphframes/connect/proto/graphframes_pb2.py‎

Lines changed: 12 additions & 12 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎python/graphframes/connect/proto/graphframes_pb2.pyi‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -251,6 +251,7 @@ class Pregel(_message.Message):
251251
"additional_col_name",
252252
"additional_col_initial",
253253
"additional_col_upd",
254+
"early_stopping",
254255
)
255256
AGG_MSGS_FIELD_NUMBER: _ClassVar[int]
256257
SEND_MSG_TO_DST_FIELD_NUMBER: _ClassVar[int]
@@ -260,6 +261,7 @@ class Pregel(_message.Message):
260261
ADDITIONAL_COL_NAME_FIELD_NUMBER: _ClassVar[int]
261262
ADDITIONAL_COL_INITIAL_FIELD_NUMBER: _ClassVar[int]
262263
ADDITIONAL_COL_UPD_FIELD_NUMBER: _ClassVar[int]
264+
EARLY_STOPPING_FIELD_NUMBER: _ClassVar[int]
263265
agg_msgs: ColumnOrExpression
264266
send_msg_to_dst: _containers.RepeatedCompositeFieldContainer[ColumnOrExpression]
265267
send_msg_to_src: _containers.RepeatedCompositeFieldContainer[ColumnOrExpression]
@@ -268,6 +270,7 @@ class Pregel(_message.Message):
268270
additional_col_name: str
269271
additional_col_initial: ColumnOrExpression
270272
additional_col_upd: ColumnOrExpression
273+
early_stopping: bool
271274
def __init__(
272275
self,
273276
agg_msgs: _Optional[_Union[ColumnOrExpression, _Mapping]] = ...,
@@ -278,6 +281,7 @@ class Pregel(_message.Message):
278281
additional_col_name: _Optional[str] = ...,
279282
additional_col_initial: _Optional[_Union[ColumnOrExpression, _Mapping]] = ...,
280283
additional_col_upd: _Optional[_Union[ColumnOrExpression, _Mapping]] = ...,
284+
early_stopping: bool = ...,
281285
) -> None: ...
282286

283287
class ShortestPaths(_message.Message):

0 commit comments

Comments
 (0)