fix(robotwin): close scene between episodes

This commit is contained in:
Pepijn
2026-07-28 15:48:14 +02:00
parent d613c0cd74
commit 8fc6aa0acf
2 changed files with 21 additions and 2 deletions
+10 -2
View File
@@ -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 --------------------------------------------------------
+11
View File
@@ -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()