mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
fix bug
This commit is contained in:
@@ -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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user