This commit is contained in:
Pepijn
2025-10-14 15:22:59 +02:00
parent 4170d1b6f1
commit 9950bfd66f
+60 -58
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,67 +354,69 @@ 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:
logging.info(f"Checkpoint policy after step {step}") if is_main_process:
checkpoint_dir = get_step_checkpoint_dir(cfg.output_dir, cfg.steps, step) logging.info(f"Checkpoint policy after step {step}")
save_checkpoint( checkpoint_dir = get_step_checkpoint_dir(cfg.output_dir, cfg.steps, step)
checkpoint_dir=checkpoint_dir, save_checkpoint(
step=step, checkpoint_dir=checkpoint_dir,
cfg=cfg, step=step,
policy=accelerator.unwrap_model(policy), cfg=cfg,
optimizer=optimizer, policy=accelerator.unwrap_model(policy),
scheduler=lr_scheduler, optimizer=optimizer,
preprocessor=preprocessor, scheduler=lr_scheduler,
postprocessor=postprocessor, preprocessor=preprocessor,
) postprocessor=postprocessor,
update_last_checkpoint(checkpoint_dir) )
if wandb_logger: update_last_checkpoint(checkpoint_dir)
wandb_logger.log_policy(checkpoint_dir) if wandb_logger:
wandb_logger.log_policy(checkpoint_dir)
accelerator.wait_for_everyone() accelerator.wait_for_everyone()
if cfg.env and is_eval_step: if cfg.env and is_eval_step:
step_id = get_step_identifier(step, cfg.steps) if is_main_process:
logging.info(f"Eval policy at step {step}") step_id = get_step_identifier(step, cfg.steps)
with torch.no_grad(), accelerator.autocast(): logging.info(f"Eval policy at step {step}")
eval_info = eval_policy_all( with torch.no_grad(), accelerator.autocast():
envs=eval_env, # dict[suite][task_id] -> vec_env eval_info = eval_policy_all(
policy=accelerator.unwrap_model(policy), envs=eval_env, # dict[suite][task_id] -> vec_env
preprocessor=preprocessor, policy=accelerator.unwrap_model(policy),
postprocessor=postprocessor, preprocessor=preprocessor,
n_episodes=cfg.eval.n_episodes, postprocessor=postprocessor,
videos_dir=cfg.output_dir / "eval" / f"videos_step_{step_id}", n_episodes=cfg.eval.n_episodes,
max_episodes_rendered=4, videos_dir=cfg.output_dir / "eval" / f"videos_step_{step_id}",
start_seed=cfg.seed, max_episodes_rendered=4,
max_parallel_tasks=cfg.env.max_parallel_tasks, start_seed=cfg.seed,
max_parallel_tasks=cfg.env.max_parallel_tasks,
)
# overall metrics (suite-agnostic)
aggregated = eval_info["overall"]
# optional: per-suite logging
for suite, suite_info in eval_info.items():
logging.info("Suite %s aggregated: %s", suite, suite_info)
# meters/tracker
eval_metrics = {
"avg_sum_reward": AverageMeter("∑rwrd", ":.3f"),
"pc_success": AverageMeter("success", ":.1f"),
"eval_s": AverageMeter("eval_s", ":.3f"),
}
eval_tracker = MetricsTracker(
cfg.batch_size,
dataset.num_frames,
dataset.num_episodes,
eval_metrics,
initial_step=step,
accelerator=accelerator,
) )
# overall metrics (suite-agnostic) eval_tracker.eval_s = aggregated.pop("eval_s")
aggregated = eval_info["overall"] eval_tracker.avg_sum_reward = aggregated.pop("avg_sum_reward")
eval_tracker.pc_success = aggregated.pop("pc_success")
# optional: per-suite logging if wandb_logger:
for suite, suite_info in eval_info.items(): wandb_log_dict = {**eval_tracker.to_dict(), **eval_info}
logging.info("Suite %s aggregated: %s", suite, suite_info) wandb_logger.log_dict(wandb_log_dict, step, mode="eval")
wandb_logger.log_video(eval_info["overall"]["video_paths"][0], step, mode="eval")
# meters/tracker
eval_metrics = {
"avg_sum_reward": AverageMeter("∑rwrd", ":.3f"),
"pc_success": AverageMeter("success", ":.1f"),
"eval_s": AverageMeter("eval_s", ":.3f"),
}
eval_tracker = MetricsTracker(
cfg.batch_size,
dataset.num_frames,
dataset.num_episodes,
eval_metrics,
initial_step=step,
accelerator=accelerator,
)
eval_tracker.eval_s = aggregated.pop("eval_s")
eval_tracker.avg_sum_reward = aggregated.pop("avg_sum_reward")
eval_tracker.pc_success = aggregated.pop("pc_success")
if wandb_logger:
wandb_log_dict = {**eval_tracker.to_dict(), **eval_info}
wandb_logger.log_dict(wandb_log_dict, step, mode="eval")
wandb_logger.log_video(eval_info["overall"]["video_paths"][0], step, mode="eval")
accelerator.wait_for_everyone() accelerator.wait_for_everyone()