From 8fc6aa0acf419f9157532dab5ac2bc3b18351bef Mon Sep 17 00:00:00 2001 From: Pepijn Date: Tue, 28 Jul 2026 15:48:14 +0200 Subject: [PATCH] fix(robotwin): close scene between episodes --- src/lerobot/envs/robotwin.py | 12 ++++++++++-- tests/envs/test_robotwin.py | 11 +++++++++++ 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/src/lerobot/envs/robotwin.py b/src/lerobot/envs/robotwin.py index 5b03f337b..010e77f3e 100644 --- a/src/lerobot/envs/robotwin.py +++ b/src/lerobot/envs/robotwin.py @@ -383,8 +383,11 @@ class RoboTwinEnv(gym.Env): self.render_mode = render_mode self._env: Any | None = None # deferred — created on first reset() inside worker + self._episode_active = False self._step_count: int = 0 - self._black_frame = np.zeros((self.observation_height, self.observation_width, 3), dtype=np.uint8) + self._black_frame: np.ndarray = np.zeros( + (self.observation_height, self.observation_width, 3), dtype=np.uint8 + ) image_spaces = { cam: spaces.Box( @@ -464,12 +467,16 @@ class RoboTwinEnv(gym.Env): self._ensure_env() super().reset(seed=seed) assert self._env is not None # set by _ensure_env() above + if self._episode_active and hasattr(self._env, "close_env"): + self._env.close_env() + self._episode_active = False actual_seed = self.episode_index if seed is None else seed setup_kwargs = _load_robotwin_setup_kwargs(self.task_name) setup_kwargs.update(seed=actual_seed, is_test=True) with torch.enable_grad(): self._env.setup_demo(**setup_kwargs) + self._episode_active = True self.episode_index += self._reset_stride self._step_count = 0 @@ -496,6 +503,7 @@ class RoboTwinEnv(gym.Env): with torch.enable_grad(): if self.action_mode == "ee": + assert self._init_eef_pose is not None, "EEF pose must be initialized during reset()." ee_action = _add_init_eef_pose(np.asarray(action, dtype=np.float64), self._init_eef_pose) self._env.take_action(ee_action, action_type="ee") elif hasattr(self._env, "take_action"): @@ -524,7 +532,6 @@ class RoboTwinEnv(gym.Env): "task": self.task_name, "is_success": is_success, } - self.reset() return obs, reward, terminated, truncated, info @@ -544,6 +551,7 @@ class RoboTwinEnv(gym.Env): with contextlib.suppress(TypeError): self._env.close_env() self._env = None + self._episode_active = False # ---- Multi-task factory -------------------------------------------------------- diff --git a/tests/envs/test_robotwin.py b/tests/envs/test_robotwin.py index fcd45adbf..b982a67ed 100644 --- a/tests/envs/test_robotwin.py +++ b/tests/envs/test_robotwin.py @@ -143,6 +143,16 @@ class TestRoboTwinEnv: assert call_kwargs["seed"] == 42 assert call_kwargs["is_test"] is True + def test_repeated_reset_closes_previous_scene(self): + mock_task = _make_mock_task_env() + env = RoboTwinEnv(task_name="beat_block_hammer") + with _patch_runtime(mock_task): + env.reset(seed=1) + env.reset(seed=2) + + mock_task.close_env.assert_called_once() + assert mock_task.setup_demo.call_count == 2 + def test_step_returns_correct_types(self): mock_task = _make_mock_task_env() env = RoboTwinEnv(task_name="beat_block_hammer") @@ -176,6 +186,7 @@ class TestRoboTwinEnv: _, _, terminated, _, info = env.step(action) assert terminated is True assert info["is_success"] is True + assert mock_task.setup_demo.call_count == 1 def test_truncation_after_episode_length(self): mock_task = _make_mock_task_env()