fix(eval): address review comments

- Wrap rollout loop in try/finally so finalize() runs on crash/interrupt
- Guard push_to_hub with num_episodes > 0 to avoid pushing empty datasets
- Hoist loop-invariant multi_env and base_repo_id out of creation loop
This commit is contained in:
Khalil Meftah
2026-06-22 19:03:33 +02:00
parent a130a9db39
commit e46c6c1912
+15 -9
View File
@@ -216,10 +216,10 @@ def rollout(
features = _env_features_to_dataset_features(env_features) features = _env_features_to_dataset_features(env_features)
fps = env.unwrapped.metadata.get("render_fps", 30) fps = env.unwrapped.metadata.get("render_fps", 30)
recording_datasets = [] recording_datasets = []
for i in range(env.num_envs):
multi_env = env.num_envs > 1 multi_env = env.num_envs > 1
root = str(recording_dir / f"env_{i}") if multi_env else str(recording_dir)
base_repo_id = recording_repo_id or "eval_recording" base_repo_id = recording_repo_id or "eval_recording"
for i in range(env.num_envs):
root = str(recording_dir / f"env_{i}") if multi_env else str(recording_dir)
repo_id = f"{base_repo_id}_env_{i}" if multi_env else base_repo_id repo_id = f"{base_repo_id}_env_{i}" if multi_env else base_repo_id
recording_datasets.append( recording_datasets.append(
LeRobotDataset.create( LeRobotDataset.create(
@@ -253,6 +253,7 @@ def rollout(
leave=False, leave=False,
) )
check_env_attributes_and_types(env) check_env_attributes_and_types(env)
try:
while not np.all(done) and step < max_steps: while not np.all(done) and step < max_steps:
# Numpy array to tensor and changing dictionary keys to LeRobot policy format. # Numpy array to tensor and changing dictionary keys to LeRobot policy format.
observation = preprocess_observation(observation) observation = preprocess_observation(observation)
@@ -303,7 +304,9 @@ def rollout(
elif "is_success" in info: elif "is_success" in info:
is_success = info["is_success"] is_success = info["is_success"]
successes = ( successes = (
is_success.tolist() if hasattr(is_success, "tolist") else [bool(is_success)] * env.num_envs is_success.tolist()
if hasattr(is_success, "tolist")
else [bool(is_success)] * env.num_envs
) )
else: else:
successes = [False] * env.num_envs successes = [False] * env.num_envs
@@ -347,6 +350,15 @@ def rollout(
) )
progbar.set_postfix({"running_success_rate": f"{running_success_rate.item() * 100:.1f}%"}) progbar.set_postfix({"running_success_rate": f"{running_success_rate.item() * 100:.1f}%"})
progbar.update() progbar.update()
finally:
if recording_datasets is not None:
for ds in recording_datasets:
ds.finalize()
if recording_repo_id is not None:
if ds.num_episodes > 0:
ds.push_to_hub(private=recording_private)
else:
logging.warning("No episodes recorded for %s — skipping push to hub.", ds.repo_id)
# Track the final observation. # Track the final observation.
if return_observations: if return_observations:
@@ -366,12 +378,6 @@ def rollout(
stacked_observations[key] = torch.stack([obs[key] for obs in all_observations], dim=1) stacked_observations[key] = torch.stack([obs[key] for obs in all_observations], dim=1)
ret[OBS_STR] = stacked_observations ret[OBS_STR] = stacked_observations
if recording_datasets is not None:
for ds in recording_datasets:
ds.finalize()
if recording_repo_id is not None:
ds.push_to_hub(private=recording_private)
if hasattr(policy, "use_original_modules"): if hasattr(policy, "use_original_modules"):
policy.use_original_modules() policy.use_original_modules()