@@ -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 )
0 commit comments