@@ -674,6 +674,7 @@ def training_loop_iteration(
674674 eval_interval = immutable_data ["eval_interval" ]
675675 eval_steps = immutable_data ["eval_steps" ]
676676 start_step = immutable_data ["start_step" ]
677+ eval_start_step = immutable_data ["eval_start_step" ]
677678
678679 # HLO dump config
679680 dump_hlo = immutable_data ["dump_hlo" ]
@@ -717,7 +718,7 @@ def training_loop_iteration(
717718 all_host_upload = dump_hlo_upload_all ,
718719 )
719720
720- if eval_interval > 0 and step > start_step and (step + 1 ) % eval_interval == 0 :
721+ if eval_interval > 0 and step >= start_step and step >= eval_start_step and (step + 1 ) % eval_interval == 0 :
721722 assert eval_data_iterator
722723 # Explicitly reset the eval iterator and counters before starting the eval loop
723724 eval_data_iterator .reset ()
@@ -871,6 +872,7 @@ def train_loop(config, recorder, state=None):
871872 "steps" : config .steps ,
872873 "eval_interval" : config .eval_interval ,
873874 "eval_steps" : config .eval_steps ,
875+ "eval_start_step" : config .eval_start_step ,
874876 "save_checkpoint_on_completion" : config .save_checkpoint_on_completion ,
875877 "start_step" : start_step ,
876878 "dump_hlo" : config .dump_hlo ,
0 commit comments