diff --git a/src/lerobot/scripts/lerobot_eval.py b/src/lerobot/scripts/lerobot_eval.py index 722763d6e..c4ed35145 100644 --- a/src/lerobot/scripts/lerobot_eval.py +++ b/src/lerobot/scripts/lerobot_eval.py @@ -453,6 +453,9 @@ def eval_policy( raise exc from None start = time.time() + # Preserve the mode for direct callers. eval_policy_all scopes the mode + # around all tasks so parallel evaluations cannot race with each other. + was_training = policy.training policy.eval() # Determine how many batched rollouts we need to get n_episodes. Note that if n_episodes is not evenly @@ -674,6 +677,8 @@ def eval_policy( if save_predicted_video: info["predicted_video_paths"] = predicted_video_paths + policy.train(was_training) + return info @@ -1010,40 +1015,48 @@ def eval_policy_all( recording_private=recording_private, ) - if max_parallel_tasks <= 1: - prefetch_thread: threading.Thread | None = None - for i, (task_group, task_id, env) in enumerate(tasks): - if prefetch_thread is not None: - prefetch_thread.join() - prefetch_thread = None + # Set the shared policy's mode before launching any workers. Restoring it + # inside individual tasks would let one task enable training mode while + # another task is still evaluating. + was_training = policy.training + policy.eval() + try: + if max_parallel_tasks <= 1: + prefetch_thread: threading.Thread | None = None + for i, (task_group, task_id, env) in enumerate(tasks): + if prefetch_thread is not None: + prefetch_thread.join() + prefetch_thread = None - try: - tg, tid, metrics = task_runner(task_group, task_id, env) - _accumulate_to(tg, metrics) - per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics}) - finally: - env.close() - # Prefetch next task's workers *after* closing current env to prevent - # GPU memory overlap between consecutive tasks. - if i + 1 < len(tasks): - next_env = tasks[i + 1][2] - if hasattr(next_env, "_ensure"): - prefetch_thread = threading.Thread(target=next_env._ensure, daemon=True) - prefetch_thread.start() - else: - with cf.ThreadPoolExecutor(max_workers=max_parallel_tasks) as executor: - fut2meta = {} - for task_group, task_id, env in tasks: - fut = executor.submit(task_runner, task_group, task_id, env) - fut2meta[fut] = (task_group, task_id, env) - for fut in cf.as_completed(fut2meta): - tg, tid, env = fut2meta[fut] try: - tg, tid, metrics = fut.result() + tg, tid, metrics = task_runner(task_group, task_id, env) _accumulate_to(tg, metrics) per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics}) finally: env.close() + # Prefetch next task's workers *after* closing current env to prevent + # GPU memory overlap between consecutive tasks. + if i + 1 < len(tasks): + next_env = tasks[i + 1][2] + if hasattr(next_env, "_ensure"): + prefetch_thread = threading.Thread(target=next_env._ensure, daemon=True) + prefetch_thread.start() + else: + with cf.ThreadPoolExecutor(max_workers=max_parallel_tasks) as executor: + fut2meta = {} + for task_group, task_id, env in tasks: + fut = executor.submit(task_runner, task_group, task_id, env) + fut2meta[fut] = (task_group, task_id, env) + for fut in cf.as_completed(fut2meta): + tg, tid, env = fut2meta[fut] + try: + tg, tid, metrics = fut.result() + _accumulate_to(tg, metrics) + per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics}) + finally: + env.close() + finally: + policy.train(was_training) # compute aggregated metrics helper (robust to lists/scalars) def _agg_from_list(xs):