mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 09:46:00 +00:00
feat(rollout): add episode success labeling to DAgger strategy
This commit is contained in:
@@ -121,6 +121,8 @@ class DAggerPedalConfig:
|
|||||||
pause_resume: str = "KEY_A"
|
pause_resume: str = "KEY_A"
|
||||||
correction: str = "KEY_B"
|
correction: str = "KEY_B"
|
||||||
upload: str = "KEY_C"
|
upload: str = "KEY_C"
|
||||||
|
success: str = "KEY_D"
|
||||||
|
failure: str = "KEY_E"
|
||||||
|
|
||||||
|
|
||||||
@RolloutStrategyConfig.register_subclass("episodic")
|
@RolloutStrategyConfig.register_subclass("episodic")
|
||||||
|
|||||||
@@ -118,6 +118,9 @@ class DAggerEvents:
|
|||||||
# Episode success labeling
|
# Episode success labeling
|
||||||
self._episode_success: bool | None = None
|
self._episode_success: bool | None = None
|
||||||
|
|
||||||
|
# Episode success labeling
|
||||||
|
self._episode_success: bool | None = None
|
||||||
|
|
||||||
# -- Thread-safe phase access ------------------------------------------
|
# -- Thread-safe phase access ------------------------------------------
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -181,6 +184,23 @@ class DAggerEvents:
|
|||||||
self._episode_success = None
|
self._episode_success = None
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
def mark_success(self) -> None:
|
||||||
|
"""Mark the current episode as successful (called from input threads)."""
|
||||||
|
with self._lock:
|
||||||
|
self._episode_success = True
|
||||||
|
|
||||||
|
def mark_failure(self) -> None:
|
||||||
|
"""Mark the current episode as failed (called from input threads)."""
|
||||||
|
with self._lock:
|
||||||
|
self._episode_success = False
|
||||||
|
|
||||||
|
def consume_episode_success(self) -> bool | None:
|
||||||
|
"""Consume and reset the episode success label. Returns None if unlabeled."""
|
||||||
|
with self._lock:
|
||||||
|
result = self._episode_success
|
||||||
|
self._episode_success = None
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Input device handlers
|
# Input device handlers
|
||||||
@@ -243,6 +263,12 @@ def _init_dagger_pedal(events: DAggerEvents, cfg: DAggerPedalConfig):
|
|||||||
events.request_transition(code_to_event[code])
|
events.request_transition(code_to_event[code])
|
||||||
if code == cfg.upload:
|
if code == cfg.upload:
|
||||||
events.upload_requested.set()
|
events.upload_requested.set()
|
||||||
|
if code == cfg.success:
|
||||||
|
events.mark_success()
|
||||||
|
logger.info("Episode marked as SUCCESS (pedal)")
|
||||||
|
if code == cfg.failure:
|
||||||
|
events.mark_failure()
|
||||||
|
logger.info("Episode marked as FAILURE (pedal)")
|
||||||
|
|
||||||
logger.info("Initializing DAgger foot pedal listener (device=%s)", cfg.device_path)
|
logger.info("Initializing DAgger foot pedal listener (device=%s)", cfg.device_path)
|
||||||
return start_pedal_listener(on_press, device_path=cfg.device_path)
|
return start_pedal_listener(on_press, device_path=cfg.device_path)
|
||||||
@@ -365,7 +391,6 @@ class DAggerStrategy(RolloutStrategy):
|
|||||||
return
|
return
|
||||||
|
|
||||||
label = self._events.consume_episode_success()
|
label = self._events.consume_episode_success()
|
||||||
logger.info("_stamp_episode_success: label=%s, buffer_len=%d", label, len(success_buf))
|
|
||||||
|
|
||||||
if label:
|
if label:
|
||||||
success_buf[-1] = np.array([True], dtype=bool)
|
success_buf[-1] = np.array([True], dtype=bool)
|
||||||
|
|||||||
@@ -338,6 +338,103 @@ def test_dagger_events_reset():
|
|||||||
assert not events.upload_requested.is_set()
|
assert not events.upload_requested.is_set()
|
||||||
|
|
||||||
|
|
||||||
|
def test_dagger_mark_success():
|
||||||
|
"""mark_success sets the episode label to True."""
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
events = DAggerEvents()
|
||||||
|
assert events.consume_episode_success() is None
|
||||||
|
|
||||||
|
events.mark_success()
|
||||||
|
assert events.consume_episode_success() is True
|
||||||
|
# Consuming clears the label
|
||||||
|
assert events.consume_episode_success() is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_dagger_mark_failure():
|
||||||
|
"""mark_failure sets the episode label to False."""
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
events = DAggerEvents()
|
||||||
|
events.mark_failure()
|
||||||
|
assert events.consume_episode_success() is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_dagger_success_overrides_failure():
|
||||||
|
"""Last label wins — success after failure overrides."""
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
events = DAggerEvents()
|
||||||
|
events.mark_failure()
|
||||||
|
events.mark_success()
|
||||||
|
assert events.consume_episode_success() is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_dagger_reset_clears_success_label():
|
||||||
|
"""reset() clears any pending episode success label."""
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
events = DAggerEvents()
|
||||||
|
events.mark_success()
|
||||||
|
events.reset()
|
||||||
|
assert events.consume_episode_success() is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_stamp_episode_success_labels_terminal_frame():
|
||||||
|
"""_stamp_episode_success sets last frame's next.success to True."""
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from lerobot.rollout.strategies.dagger import DAggerStrategy
|
||||||
|
|
||||||
|
strategy = DAggerStrategy.__new__(DAggerStrategy)
|
||||||
|
strategy.config = MagicMock()
|
||||||
|
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
strategy._events = DAggerEvents()
|
||||||
|
strategy._events.mark_success()
|
||||||
|
|
||||||
|
dataset = MagicMock()
|
||||||
|
dataset.writer.episode_buffer = {
|
||||||
|
"next.success": [
|
||||||
|
np.array([False], dtype=bool),
|
||||||
|
np.array([False], dtype=bool),
|
||||||
|
np.array([False], dtype=bool),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
strategy._stamp_episode_success(dataset)
|
||||||
|
|
||||||
|
assert dataset.writer.episode_buffer["next.success"][-1].item() is True
|
||||||
|
assert dataset.writer.episode_buffer["next.success"][0].item() is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_stamp_episode_success_no_label_stays_false():
|
||||||
|
"""Without a label, all frames remain False."""
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from lerobot.rollout.strategies.dagger import DAggerStrategy
|
||||||
|
|
||||||
|
strategy = DAggerStrategy.__new__(DAggerStrategy)
|
||||||
|
strategy.config = MagicMock()
|
||||||
|
|
||||||
|
from lerobot.rollout.strategies import DAggerEvents
|
||||||
|
|
||||||
|
strategy._events = DAggerEvents()
|
||||||
|
|
||||||
|
dataset = MagicMock()
|
||||||
|
dataset.writer.episode_buffer = {
|
||||||
|
"next.success": [
|
||||||
|
np.array([False], dtype=bool),
|
||||||
|
np.array([False], dtype=bool),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
strategy._stamp_episode_success(dataset)
|
||||||
|
|
||||||
|
assert all(v.item() is False for v in dataset.writer.episode_buffer["next.success"])
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Context dataclass
|
# Context dataclass
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user