diff --git a/src/lerobot/rollout/configs.py b/src/lerobot/rollout/configs.py index 14b0c1864..e877ce253 100644 --- a/src/lerobot/rollout/configs.py +++ b/src/lerobot/rollout/configs.py @@ -209,6 +209,12 @@ class RolloutConfig: # Rename map for mapping robot/dataset observation keys to policy keys rename_map: dict[str, str] = field(default_factory=dict) + # Hardware teardown + # When True (default), smoothly interpolate the robot back to the joint + # positions captured at startup before disconnecting. Set to False to + # leave the robot in its final achieved pose at shutdown. + return_to_initial_position: bool = True + # Torch compile use_torch_compile: bool = False torch_compile_backend: str = "inductor" diff --git a/src/lerobot/rollout/strategies/base.py b/src/lerobot/rollout/strategies/base.py index 703cf3654..e47b65209 100644 --- a/src/lerobot/rollout/strategies/base.py +++ b/src/lerobot/rollout/strategies/base.py @@ -78,5 +78,8 @@ class BaseStrategy(RolloutStrategy): def teardown(self, ctx: RolloutContext) -> None: """Disconnect hardware and stop inference.""" - self._teardown_hardware(ctx.hardware) + self._teardown_hardware( + ctx.hardware, + return_to_initial_position=ctx.runtime.cfg.return_to_initial_position, + ) logger.info("Base strategy teardown complete") diff --git a/src/lerobot/rollout/strategies/core.py b/src/lerobot/rollout/strategies/core.py index ffddff636..b775afb93 100644 --- a/src/lerobot/rollout/strategies/core.py +++ b/src/lerobot/rollout/strategies/core.py @@ -116,16 +116,20 @@ class RolloutStrategy(abc.ABC): engine.resume() return False - def _teardown_hardware(self, hw: HardwareContext) -> None: - """Stop the inference engine, return robot to initial position, and disconnect hardware.""" + def _teardown_hardware(self, hw: HardwareContext, return_to_initial_position: bool = True) -> None: + """Stop the inference engine, optionally return robot to initial position, and disconnect hardware.""" if self._engine is not None: logger.info("Stopping inference engine...") self._engine.stop() robot = hw.robot_wrapper.inner if robot.is_connected: - if hw.initial_position: + if return_to_initial_position and hw.initial_position: logger.info("Returning robot to initial position before shutdown...") self._return_to_initial_position(hw) + elif not return_to_initial_position: + logger.info( + "Skipping return-to-initial-position (disabled by config); leaving robot in final pose." + ) logger.info("Disconnecting robot...") robot.disconnect() teleop = hw.teleop diff --git a/src/lerobot/rollout/strategies/dagger.py b/src/lerobot/rollout/strategies/dagger.py index 48bf2ead4..6af3edbe0 100644 --- a/src/lerobot/rollout/strategies/dagger.py +++ b/src/lerobot/rollout/strategies/dagger.py @@ -371,7 +371,10 @@ class DAggerStrategy(RolloutStrategy): logger.info("Dataset uploaded to hub") log_say("Dataset uploaded to hub", play_sounds) - self._teardown_hardware(ctx.hardware) + self._teardown_hardware( + ctx.hardware, + return_to_initial_position=ctx.runtime.cfg.return_to_initial_position, + ) logger.info("DAgger strategy teardown complete") # ------------------------------------------------------------------ diff --git a/src/lerobot/rollout/strategies/highlight.py b/src/lerobot/rollout/strategies/highlight.py index de5083eeb..91816e0bf 100644 --- a/src/lerobot/rollout/strategies/highlight.py +++ b/src/lerobot/rollout/strategies/highlight.py @@ -227,7 +227,10 @@ class HighlightStrategy(RolloutStrategy): logger.info("Dataset uploaded to hub") log_say("Dataset uploaded to hub", play_sounds) - self._teardown_hardware(ctx.hardware) + self._teardown_hardware( + ctx.hardware, + return_to_initial_position=ctx.runtime.cfg.return_to_initial_position, + ) logger.info("Highlight strategy teardown complete") def _setup_keyboard(self, shutdown_event: ThreadingEvent) -> None: diff --git a/src/lerobot/rollout/strategies/sentry.py b/src/lerobot/rollout/strategies/sentry.py index d2334195a..61e38aa68 100644 --- a/src/lerobot/rollout/strategies/sentry.py +++ b/src/lerobot/rollout/strategies/sentry.py @@ -196,7 +196,10 @@ class SentryStrategy(RolloutStrategy): logger.info("Dataset uploaded to hub") log_say("Dataset uploaded to hub", play_sounds) - self._teardown_hardware(ctx.hardware) + self._teardown_hardware( + ctx.hardware, + return_to_initial_position=ctx.runtime.cfg.return_to_initial_position, + ) logger.info("Sentry strategy teardown complete") def _background_push(self, dataset, cfg) -> None: diff --git a/tests/test_rollout.py b/tests/test_rollout.py index df6963717..2dae82921 100644 --- a/tests/test_rollout.py +++ b/tests/test_rollout.py @@ -254,7 +254,7 @@ def test_create_inference_engine_sync(): def test_estimate_max_episode_seconds_no_video(): from lerobot.rollout.strategies import estimate_max_episode_seconds - assert estimate_max_episode_seconds({}, fps=30.0) == 600.0 + assert estimate_max_episode_seconds({}, fps=30.0) == 300.0 def test_estimate_max_episode_seconds_with_video(): @@ -264,7 +264,7 @@ def test_estimate_max_episode_seconds_with_video(): result = estimate_max_episode_seconds(features, fps=30.0) assert result > 0 # With a real camera, duration should differ from the fallback - assert result != 600.0 + assert result != 300.0 def test_safe_push_to_hub():