From 8c4d4c8208bdbedce0a647b5bb80ca4a0cd7cf62 Mon Sep 17 00:00:00 2001 From: Steven Palma Date: Sun, 26 Apr 2026 18:49:52 +0200 Subject: [PATCH] fix(rollout): robot_observation_processor and notify_observation at policy frequency instead of interpolator rate Co-authored-by: Pepijn <138571049+pkooij@users.noreply.github.com> --- src/lerobot/rollout/strategies/base.py | 3 +-- src/lerobot/rollout/strategies/core.py | 27 ++++++++++++++++++++- src/lerobot/rollout/strategies/dagger.py | 19 ++++++++------- src/lerobot/rollout/strategies/highlight.py | 3 +-- src/lerobot/rollout/strategies/sentry.py | 3 +-- 5 files changed, 39 insertions(+), 16 deletions(-) diff --git a/src/lerobot/rollout/strategies/base.py b/src/lerobot/rollout/strategies/base.py index 77632b615..703cf3654 100644 --- a/src/lerobot/rollout/strategies/base.py +++ b/src/lerobot/rollout/strategies/base.py @@ -60,8 +60,7 @@ class BaseStrategy(RolloutStrategy): break obs = robot.get_observation() - obs_processed = ctx.processors.robot_observation_processor(obs) - engine.notify_observation(obs_processed) + obs_processed = self._process_observation_and_notify(ctx.processors, obs) if self._handle_warmup(cfg.use_torch_compile, loop_start, control_interval): continue diff --git a/src/lerobot/rollout/strategies/core.py b/src/lerobot/rollout/strategies/core.py index 3cc8f5f53..8d6bfd930 100644 --- a/src/lerobot/rollout/strategies/core.py +++ b/src/lerobot/rollout/strategies/core.py @@ -32,7 +32,7 @@ from ..inference import InferenceEngine if TYPE_CHECKING: from ..configs import RolloutStrategyConfig - from ..context import HardwareContext, RolloutContext, RuntimeContext + from ..context import HardwareContext, ProcessorContext, RolloutContext, RuntimeContext logger = logging.getLogger(__name__) @@ -50,6 +50,7 @@ class RolloutStrategy(abc.ABC): self._engine: InferenceEngine | None = None self._interpolator: ActionInterpolator | None = None self._warmup_flushed: bool = False + self._cached_obs_processed: dict | None = None def _init_engine(self, ctx: RolloutContext) -> None: """Attach the inference engine and action interpolator, then start the backend. @@ -65,8 +66,32 @@ class RolloutStrategy(abc.ABC): self._engine.reset() self._engine.start() self._warmup_flushed = False + self._cached_obs_processed = None logger.info("Inference engine started") + def _process_observation_and_notify(self, processors: ProcessorContext, obs_raw: dict) -> dict: + """Run the observation processor and notify the engine — throttled to policy ticks. + + Callers are responsible for calling ``robot.get_observation()`` every loop + iteration so ``obs_raw`` stays fresh for the action post-processor. This + helper gates only the comparatively expensive bits — the processor pipeline + and ``engine.notify_observation`` — to fire when the interpolator signals + it needs a new action (once per ``interpolation_multiplier`` ticks). On + interpolated ticks the cached ``obs_processed`` is reused. + + With ``interpolation_multiplier == 1`` this is equivalent to the unthrottled + path: ``needs_new_action()`` is True every tick. + + The cache is implicitly invalidated whenever ``interpolator.reset()`` is + called (warmup completion, DAgger phase transitions back to AUTONOMOUS), + because reset makes ``needs_new_action()`` return True on the next call. + """ + if self._cached_obs_processed is None or self._interpolator.needs_new_action(): + obs_processed = processors.robot_observation_processor(obs_raw) + self._engine.notify_observation(obs_processed) + self._cached_obs_processed = obs_processed + return self._cached_obs_processed + def _handle_warmup(self, use_torch_compile: bool, loop_start: float, control_interval: float) -> bool: """Handle torch.compile warmup phase. diff --git a/src/lerobot/rollout/strategies/dagger.py b/src/lerobot/rollout/strategies/dagger.py index b56478045..48bf2ead4 100644 --- a/src/lerobot/rollout/strategies/dagger.py +++ b/src/lerobot/rollout/strategies/dagger.py @@ -434,19 +434,19 @@ class DAggerStrategy(RolloutStrategy): phase = events.phase obs = robot.get_observation() - obs_processed = ctx.processors.robot_observation_processor(obs) - obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR) # --- CORRECTING: human teleop control --- if phase == DAggerPhase.CORRECTING: + obs_processed = ctx.processors.robot_observation_processor(obs) teleop_action = teleop.get_action() processed_teleop = ctx.processors.teleop_action_processor((teleop_action, obs)) robot_action_to_send = ctx.processors.robot_action_processor((processed_teleop, obs)) robot.send_action(robot_action_to_send) last_action = robot_action_to_send self._log_telemetry(obs_processed, processed_teleop, ctx.runtime) - action_frame = build_dataset_frame(features, processed_teleop, prefix=ACTION) if record_tick % record_stride == 0: + obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR) + action_frame = build_dataset_frame(features, processed_teleop, prefix=ACTION) frame = { **obs_frame, **action_frame, @@ -463,7 +463,7 @@ class DAggerStrategy(RolloutStrategy): # --- AUTONOMOUS: policy control --- else: - engine.notify_observation(obs_processed) + obs_processed = self._process_observation_and_notify(ctx.processors, obs) if self._handle_warmup(cfg.use_torch_compile, loop_start, control_interval): continue @@ -472,8 +472,9 @@ class DAggerStrategy(RolloutStrategy): if action_dict is not None: self._log_telemetry(obs_processed, action_dict, ctx.runtime) last_action = ctx.processors.robot_action_processor((action_dict, obs)) - action_frame = build_dataset_frame(features, action_dict, prefix=ACTION) if record_tick % record_stride == 0: + obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR) + action_frame = build_dataset_frame(features, action_dict, prefix=ACTION) frame = { **obs_frame, **action_frame, @@ -608,10 +609,10 @@ class DAggerStrategy(RolloutStrategy): phase = events.phase obs = robot.get_observation() - obs_processed = ctx.processors.robot_observation_processor(obs) # --- CORRECTING: human teleop control + recording --- if phase == DAggerPhase.CORRECTING: + obs_processed = ctx.processors.robot_observation_processor(obs) teleop_action = teleop.get_action() processed_teleop = ctx.processors.teleop_action_processor((teleop_action, obs)) robot_action_to_send = ctx.processors.robot_action_processor((processed_teleop, obs)) @@ -619,9 +620,9 @@ class DAggerStrategy(RolloutStrategy): last_action = robot_action_to_send self._log_telemetry(obs_processed, processed_teleop, ctx.runtime) - obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR) - action_frame = build_dataset_frame(features, processed_teleop, prefix=ACTION) if record_tick % record_stride == 0: + obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR) + action_frame = build_dataset_frame(features, processed_teleop, prefix=ACTION) dataset.add_frame( { **obs_frame, @@ -639,7 +640,7 @@ class DAggerStrategy(RolloutStrategy): # --- AUTONOMOUS: policy control (no recording) --- else: - engine.notify_observation(obs_processed) + obs_processed = self._process_observation_and_notify(ctx.processors, obs) if self._handle_warmup(cfg.use_torch_compile, loop_start, control_interval): continue diff --git a/src/lerobot/rollout/strategies/highlight.py b/src/lerobot/rollout/strategies/highlight.py index a221de461..de5083eeb 100644 --- a/src/lerobot/rollout/strategies/highlight.py +++ b/src/lerobot/rollout/strategies/highlight.py @@ -135,8 +135,7 @@ class HighlightStrategy(RolloutStrategy): break obs = robot.get_observation() - obs_processed = ctx.processors.robot_observation_processor(obs) - engine.notify_observation(obs_processed) + obs_processed = self._process_observation_and_notify(ctx.processors, obs) if self._handle_warmup(cfg.use_torch_compile, loop_start, control_interval): continue diff --git a/src/lerobot/rollout/strategies/sentry.py b/src/lerobot/rollout/strategies/sentry.py index 4c0ec8af3..d2334195a 100644 --- a/src/lerobot/rollout/strategies/sentry.py +++ b/src/lerobot/rollout/strategies/sentry.py @@ -111,8 +111,7 @@ class SentryStrategy(RolloutStrategy): break obs = robot.get_observation() - obs_processed = ctx.processors.robot_observation_processor(obs) - engine.notify_observation(obs_processed) + obs_processed = self._process_observation_and_notify(ctx.processors, obs) if self._handle_warmup(cfg.use_torch_compile, loop_start, control_interval): continue