From 2939168c336e44fca1bc2234c67246b4d6fe73ac Mon Sep 17 00:00:00 2001 From: Anas <154854317+anasm266@users.noreply.github.com> Date: Thu, 30 Jul 2026 06:20:56 -0700 Subject: [PATCH] fix(envs): use RoboCasa task horizons (#4037) Co-authored-by: Pepijn <138571049+pkooij@users.noreply.github.com> --- docs/source/robocasa.mdx | 2 ++ src/lerobot/envs/configs.py | 2 +- src/lerobot/envs/robocasa.py | 15 ++++++++++- tests/envs/test_robocasa.py | 51 ++++++++++++++++++++++++++++++++++++ 4 files changed, 68 insertions(+), 2 deletions(-) create mode 100644 tests/envs/test_robocasa.py diff --git a/docs/source/robocasa.mdx b/docs/source/robocasa.mdx index 5a335a484..2a3ddb39c 100644 --- a/docs/source/robocasa.mdx +++ b/docs/source/robocasa.mdx @@ -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=` to apply the same explicit horizon to every selected task. + ### Single-task evaluation (recommended for quick iteration) ```bash diff --git a/src/lerobot/envs/configs.py b/src/lerobot/envs/configs.py index 3f6fd75f9..3c210052f 100644 --- a/src/lerobot/envs/configs.py +++ b/src/lerobot/envs/configs.py @@ -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" diff --git a/src/lerobot/envs/robocasa.py b/src/lerobot/envs/robocasa.py index a84a7c766..b127072d9 100644 --- a/src/lerobot/envs/robocasa.py +++ b/src/lerobot/envs/robocasa.py @@ -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=` 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(). diff --git a/tests/envs/test_robocasa.py b/tests/envs/test_robocasa.py new file mode 100644 index 000000000..ab01010a7 --- /dev/null +++ b/tests/envs/test_robocasa.py @@ -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()