This commit is contained in:
Pepijn
2025-10-14 15:22:59 +02:00
parent 4170d1b6f1
commit 9950bfd66f
+4 -2
View File
@@ -341,8 +341,8 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
step += 1 step += 1
train_tracker.step() train_tracker.step()
is_log_step = cfg.log_freq > 0 and step % cfg.log_freq == 0 and is_main_process is_log_step = cfg.log_freq > 0 and step % cfg.log_freq == 0 and is_main_process
is_saving_step = (step % cfg.save_freq == 0 or step == cfg.steps) and is_main_process is_saving_step = (step % cfg.save_freq == 0 or step == cfg.steps)
is_eval_step = cfg.eval_freq > 0 and step % cfg.eval_freq == 0 and is_main_process is_eval_step = cfg.eval_freq > 0 and step % cfg.eval_freq == 0
if is_log_step: if is_log_step:
logging.info(train_tracker) logging.info(train_tracker)
@@ -354,6 +354,7 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
train_tracker.reset_averages() train_tracker.reset_averages()
if cfg.save_checkpoint and is_saving_step: if cfg.save_checkpoint and is_saving_step:
if is_main_process:
logging.info(f"Checkpoint policy after step {step}") logging.info(f"Checkpoint policy after step {step}")
checkpoint_dir = get_step_checkpoint_dir(cfg.output_dir, cfg.steps, step) checkpoint_dir = get_step_checkpoint_dir(cfg.output_dir, cfg.steps, step)
save_checkpoint( save_checkpoint(
@@ -373,6 +374,7 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
accelerator.wait_for_everyone() accelerator.wait_for_everyone()
if cfg.env and is_eval_step: if cfg.env and is_eval_step:
if is_main_process:
step_id = get_step_identifier(step, cfg.steps) step_id = get_step_identifier(step, cfg.steps)
logging.info(f"Eval policy at step {step}") logging.info(f"Eval policy at step {step}")
with torch.no_grad(), accelerator.autocast(): with torch.no_grad(), accelerator.autocast():