@@ -669,6 +669,7 @@ def training_loop_iteration(
669669 eval_interval = immutable_data ["eval_interval" ]
670670 eval_steps = immutable_data ["eval_steps" ]
671671 start_step = immutable_data ["start_step" ]
672+ eval_start_step = immutable_data ["eval_start_step" ]
672673
673674 # HLO dump config
674675 dump_hlo = immutable_data ["dump_hlo" ]
@@ -712,7 +713,7 @@ def training_loop_iteration(
712713 all_host_upload = dump_hlo_upload_all ,
713714 )
714715
715- if eval_interval > 0 and step > start_step and (step + 1 ) % eval_interval == 0 :
716+ if eval_interval > 0 and step >= start_step and step >= eval_start_step and (step + 1 ) % eval_interval == 0 :
716717 assert eval_data_iterator
717718 # Explicitly reset the eval iterator and counters before starting the eval loop
718719 eval_data_iterator .reset ()
@@ -863,6 +864,7 @@ def train_loop(config, recorder, state=None):
863864 "steps" : config .steps ,
864865 "eval_interval" : config .eval_interval ,
865866 "eval_steps" : config .eval_steps ,
867+ "eval_start_step" : config .eval_start_step ,
866868 "save_checkpoint_on_completion" : config .save_checkpoint_on_completion ,
867869 "start_step" : start_step ,
868870 "dump_hlo" : config .dump_hlo ,
0 commit comments