fix(envs): use RoboCasa task horizons (#4037)

Co-authored-by: Pepijn <138571049+pkooij@users.noreply.github.com>
This commit is contained in:
Anas
2026-07-30 06:20:56 -07:00
committed by GitHub
parent 40a5e70352
commit 2939168c33
4 changed files with 68 additions and 2 deletions
+2
View File
@@ -82,6 +82,8 @@ By default the env samples objects only from the `lightwheel` registry (what `--
All eval snippets below mirror the CI command (see `.github/workflows/benchmark_tests.yml`). The `--rename_map` argument maps RoboCasa's native camera keys (`robot0_agentview_left` / `robot0_eye_in_hand` / `robot0_agentview_right`) onto the three-camera (`camera1` / `camera2` / `camera3`) input layout the released `smolvla_robocasa` policy was trained on.
By default, each task uses the rollout horizon registered by RoboCasa. Set `--env.episode_length=<steps>` to apply the same explicit horizon to every selected task.
### Single-task evaluation (recommended for quick iteration)
```bash
+1 -1
View File
@@ -507,7 +507,7 @@ class MetaworldEnv(EnvConfig):
class RoboCasaEnv(EnvConfig):
task: str = "CloseFridge"
fps: int = 20
episode_length: int = 1000
episode_length: int | None = None
obs_type: str = "pixels_agent_pos"
render_mode: str = "rgb_array"
camera_name: str = "robot0_agentview_left,robot0_eye_in_hand,robot0_agentview_right"
+14 -1
View File
@@ -98,6 +98,19 @@ def _resolve_tasks(task: str) -> tuple[list[str], str | None]:
return names, None
def _get_task_horizon(task: str) -> int:
"""Return the rollout horizon registered by RoboCasa for a task."""
from robocasa.utils.dataset_registry_utils import get_task_horizon
try:
return int(get_task_horizon(task))
except ValueError as exc:
raise ValueError(
f"No RoboCasa horizon is registered for task '{task}'. "
"Set `--env.episode_length=<steps>` explicitly."
) from exc
def convert_action(flat_action: np.ndarray) -> dict[str, Any]:
"""Split a flat (12,) action vector into a RoboCasa action dict.
@@ -154,7 +167,7 @@ class RoboCasaEnv(gym.Env):
self.camera_name = parse_camera_names(camera_name)
self._max_episode_steps = episode_length if episode_length is not None else 1000
self._max_episode_steps = episode_length if episode_length is not None else _get_task_horizon(task)
# Deferred — created on first reset() inside the worker subprocess
# to avoid inheriting stale GPU/EGL contexts across fork().
+51
View File
@@ -0,0 +1,51 @@
from __future__ import annotations
from collections.abc import Callable, Sequence
from unittest.mock import Mock, call
import pytest
from lerobot.envs import robocasa
from lerobot.envs.configs import RoboCasaEnv as RoboCasaEnvConfig
def _instantiate_envs(
factories: Sequence[Callable[[], robocasa.RoboCasaEnv]],
) -> list[robocasa.RoboCasaEnv]:
return [factory() for factory in factories]
def test_robocasa_config_uses_registered_horizon_by_default() -> None:
assert RoboCasaEnvConfig().episode_length is None
def test_multi_task_envs_use_registered_horizons(monkeypatch: pytest.MonkeyPatch) -> None:
horizons = {"CloseFridge": 900, "SearingMeat": 4350}
get_task_horizon = Mock(side_effect=horizons.__getitem__)
monkeypatch.setattr(robocasa, "_get_task_horizon", get_task_horizon)
envs = robocasa.create_robocasa_envs(
task="CloseFridge,SearingMeat",
n_envs=1,
env_cls=_instantiate_envs,
)
assert envs["CloseFridge"][0][0]._max_episode_steps == 900
assert envs["SearingMeat"][0][0]._max_episode_steps == 4350
assert get_task_horizon.call_args_list == [call("CloseFridge"), call("SearingMeat")]
def test_explicit_episode_length_overrides_registered_horizons(monkeypatch: pytest.MonkeyPatch) -> None:
get_task_horizon = Mock()
monkeypatch.setattr(robocasa, "_get_task_horizon", get_task_horizon)
envs = robocasa.create_robocasa_envs(
task="CloseFridge,SearingMeat",
n_envs=1,
env_cls=_instantiate_envs,
episode_length=1234,
)
assert envs["CloseFridge"][0][0]._max_episode_steps == 1234
assert envs["SearingMeat"][0][0]._max_episode_steps == 1234
get_task_horizon.assert_not_called()