mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 01:19:46 +00:00
fix(robotwin): close scene between episodes
This commit is contained in:
@@ -383,8 +383,11 @@ class RoboTwinEnv(gym.Env):
|
|||||||
self.render_mode = render_mode
|
self.render_mode = render_mode
|
||||||
|
|
||||||
self._env: Any | None = None # deferred — created on first reset() inside worker
|
self._env: Any | None = None # deferred — created on first reset() inside worker
|
||||||
|
self._episode_active = False
|
||||||
self._step_count: int = 0
|
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 = {
|
image_spaces = {
|
||||||
cam: spaces.Box(
|
cam: spaces.Box(
|
||||||
@@ -464,12 +467,16 @@ class RoboTwinEnv(gym.Env):
|
|||||||
self._ensure_env()
|
self._ensure_env()
|
||||||
super().reset(seed=seed)
|
super().reset(seed=seed)
|
||||||
assert self._env is not None # set by _ensure_env() above
|
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
|
actual_seed = self.episode_index if seed is None else seed
|
||||||
setup_kwargs = _load_robotwin_setup_kwargs(self.task_name)
|
setup_kwargs = _load_robotwin_setup_kwargs(self.task_name)
|
||||||
setup_kwargs.update(seed=actual_seed, is_test=True)
|
setup_kwargs.update(seed=actual_seed, is_test=True)
|
||||||
with torch.enable_grad():
|
with torch.enable_grad():
|
||||||
self._env.setup_demo(**setup_kwargs)
|
self._env.setup_demo(**setup_kwargs)
|
||||||
|
self._episode_active = True
|
||||||
self.episode_index += self._reset_stride
|
self.episode_index += self._reset_stride
|
||||||
self._step_count = 0
|
self._step_count = 0
|
||||||
|
|
||||||
@@ -496,6 +503,7 @@ class RoboTwinEnv(gym.Env):
|
|||||||
|
|
||||||
with torch.enable_grad():
|
with torch.enable_grad():
|
||||||
if self.action_mode == "ee":
|
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)
|
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")
|
self._env.take_action(ee_action, action_type="ee")
|
||||||
elif hasattr(self._env, "take_action"):
|
elif hasattr(self._env, "take_action"):
|
||||||
@@ -524,7 +532,6 @@ class RoboTwinEnv(gym.Env):
|
|||||||
"task": self.task_name,
|
"task": self.task_name,
|
||||||
"is_success": is_success,
|
"is_success": is_success,
|
||||||
}
|
}
|
||||||
self.reset()
|
|
||||||
|
|
||||||
return obs, reward, terminated, truncated, info
|
return obs, reward, terminated, truncated, info
|
||||||
|
|
||||||
@@ -544,6 +551,7 @@ class RoboTwinEnv(gym.Env):
|
|||||||
with contextlib.suppress(TypeError):
|
with contextlib.suppress(TypeError):
|
||||||
self._env.close_env()
|
self._env.close_env()
|
||||||
self._env = None
|
self._env = None
|
||||||
|
self._episode_active = False
|
||||||
|
|
||||||
|
|
||||||
# ---- Multi-task factory --------------------------------------------------------
|
# ---- Multi-task factory --------------------------------------------------------
|
||||||
|
|||||||
@@ -143,6 +143,16 @@ class TestRoboTwinEnv:
|
|||||||
assert call_kwargs["seed"] == 42
|
assert call_kwargs["seed"] == 42
|
||||||
assert call_kwargs["is_test"] is True
|
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):
|
def test_step_returns_correct_types(self):
|
||||||
mock_task = _make_mock_task_env()
|
mock_task = _make_mock_task_env()
|
||||||
env = RoboTwinEnv(task_name="beat_block_hammer")
|
env = RoboTwinEnv(task_name="beat_block_hammer")
|
||||||
@@ -176,6 +186,7 @@ class TestRoboTwinEnv:
|
|||||||
_, _, terminated, _, info = env.step(action)
|
_, _, terminated, _, info = env.step(action)
|
||||||
assert terminated is True
|
assert terminated is True
|
||||||
assert info["is_success"] is True
|
assert info["is_success"] is True
|
||||||
|
assert mock_task.setup_demo.call_count == 1
|
||||||
|
|
||||||
def test_truncation_after_episode_length(self):
|
def test_truncation_after_episode_length(self):
|
||||||
mock_task = _make_mock_task_env()
|
mock_task = _make_mock_task_env()
|
||||||
|
|||||||
Reference in New Issue
Block a user