fix(rtc): validate trained prefix capacity

This commit is contained in:
Pepijn
2026-07-16 12:30:40 +02:00
parent fa477fca5a
commit c395286f3c
10 changed files with 125 additions and 18 deletions
+1 -1
View File
@@ -280,7 +280,7 @@ lerobot-rollout \
| `--inference.rtc.mode` | `guided` (default) or trained-prefix `trained` for compatible Pi052 checkpoints | | `--inference.rtc.mode` | `guided` (default) or trained-prefix `trained` for compatible Pi052 checkpoints |
| `--inference.rtc.max_guidance_weight` | Consistency enforcement strength (default: varies by policy) | | `--inference.rtc.max_guidance_weight` | Consistency enforcement strength (default: varies by policy) |
| `--inference.rtc.prefix_attention_schedule` | Blend schedule: `LINEAR`, `EXP`, `ONES`, `ZEROS` | | `--inference.rtc.prefix_attention_schedule` | Blend schedule: `LINEAR`, `EXP`, `ONES`, `ZEROS` |
| `--inference.queue_threshold` | Max queue size before backpressure (default: 30) | | `--inference.queue_threshold` | Backpressure threshold; trained RTC requires at least its maximum delay |
See the [Real-Time Chunking](./rtc) guide for details on tuning RTC parameters. See the [Real-Time Chunking](./rtc) guide for details on tuning RTC parameters.
+3 -2
View File
@@ -152,8 +152,9 @@ than `chunk_size`. Choose it to cover the largest inference latency expected at
deployment: at 50 Hz, for example, 10 steps correspond to 200 ms. A delay of deployment: at 50 Hz, for example, 10 steps correspond to 200 ms. A delay of
zero is included in the uniform sampling distribution, so the checkpoint also zero is included in the uniform sampling distribution, so the checkpoint also
continues to receive ordinary flow-matching examples. Set rollout's continues to receive ordinary flow-matching examples. Set rollout's
`inference.rtc.execution_horizon` to at least this maximum so the previous `inference.rtc.execution_horizon` and `inference.queue_threshold` to at least
chunk cache retains enough actions to construct every supported prefix. this maximum so inference starts early enough and the previous chunk retains
every action needed for the committed prefix.
Run the resulting checkpoint with the asynchronous `lerobot-rollout` backend Run the resulting checkpoint with the asynchronous `lerobot-rollout` backend
and select the trained-prefix path explicitly: and select the trained-prefix path explicitly:
+8
View File
@@ -97,6 +97,10 @@ for step in range(num_steps):
- `guided` (default) applies the original Jacobian guidance during denoising and works with ordinary flow-matching checkpoints. - `guided` (default) applies the original Jacobian guidance during denoising and works with ordinary flow-matching checkpoints.
- `trained` hard-inpaints the previous chunk's prefix with per-action flow timesteps. It currently requires a Pi052 checkpoint trained with `policy.rtc_training_max_delay > 0` and avoids the guidance backward pass. - `trained` hard-inpaints the previous chunk's prefix with per-action flow timesteps. It currently requires a Pi052 checkpoint trained with `policy.rtc_training_max_delay > 0` and avoids the guidance backward pass.
For trained mode, both `execution_horizon` and the rollout backend's
`inference.queue_threshold` must be at least the checkpoint's
`rtc_training_max_delay`; rollout validates this before connecting the robot.
**`execution_horizon`**: How many timesteps from the previous chunk to maintain consistency with. Higher values mean smoother transitions but potentially less reactivity. **`execution_horizon`**: How many timesteps from the previous chunk to maintain consistency with. Higher values mean smoother transitions but potentially less reactivity.
Typical values: 8-12 steps Typical values: 8-12 steps
@@ -129,6 +133,10 @@ python examples/rtc/eval_dataset.py \
--device=cuda --device=cuda
``` ```
Add `--rtc.mode=trained` when evaluating a compatible training-time RTC Pi052
checkpoint. Unsupported policies reject trained mode instead of falling back to
guided RTC.
The script generates a visualization of the denoising process, comparing standard generation (left) with RTC (right). In the RTC plots, you can see how the first few steps (blue/purple lines) are guided to match the red ground truth trajectory (previous chunk's tail), ensuring a smooth transition between chunks. The script generates a visualization of the denoising process, comparing standard generation (left) with RTC (right). In the RTC plots, you can see how the first few steps (blue/purple lines) are guided to match the red ground truth trajectory (previous chunk's tail), ensuring a smooth transition between chunks.
<p align="center"> <p align="center">
+1
View File
@@ -306,6 +306,7 @@ class RTCEvaluator:
# Configure RTC # Configure RTC
rtc_config = RTCConfig( rtc_config = RTCConfig(
enabled=rtc_enabled, enabled=rtc_enabled,
mode=self.cfg.rtc.mode,
execution_horizon=self.cfg.rtc.execution_horizon, execution_horizon=self.cfg.rtc.execution_horizon,
max_guidance_weight=self.cfg.rtc.max_guidance_weight, max_guidance_weight=self.cfg.rtc.max_guidance_weight,
prefix_attention_schedule=self.cfg.rtc.prefix_attention_schedule, prefix_attention_schedule=self.cfg.rtc.prefix_attention_schedule,
+4 -1
View File
@@ -1284,7 +1284,10 @@ class PI05Policy(PreTrainedPolicy):
# Create processor if config provided # Create processor if config provided
# If RTC is not enabled - we can still track the denoising data # If RTC is not enabled - we can still track the denoising data
if self.config.rtc_config is not None: if self.config.rtc_config is not None:
self.rtc_processor = RTCProcessor(self.config.rtc_config) self.rtc_processor = RTCProcessor(
self.config.rtc_config,
trained_mode_supported=int(getattr(self.config, "rtc_training_max_delay", 0)) > 0,
)
model_value = getattr(self, "model", None) model_value = getattr(self, "model", None)
if model_value is not None: if model_value is not None:
+5 -1
View File
@@ -42,7 +42,11 @@ class RTCProcessor:
prefix attention, and adaptive chunk processing. prefix attention, and adaptive chunk processing.
""" """
def __init__(self, rtc_config: RTCConfig): def __init__(self, rtc_config: RTCConfig, *, trained_mode_supported: bool = False):
if rtc_config.enabled and rtc_config.mode == "trained" and not trained_mode_supported:
raise ValueError(
"RTC mode='trained' requires a Pi052 checkpoint trained with rtc_training_max_delay > 0."
)
self.rtc_config = rtc_config self.rtc_config = rtc_config
self.tracker = None self.tracker = None
+31 -11
View File
@@ -60,6 +60,35 @@ from .robot_wrapper import ThreadSafeRobot
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _validate_trained_rtc_rollout_config(policy_config, inference_config: RTCInferenceConfig) -> None:
"""Fail fast when rollout cannot retain every trained RTC prefix."""
rtc = inference_config.rtc
if not rtc.enabled or rtc.mode != "trained":
return
if policy_config.type != "pi052":
raise ValueError(
"--inference.rtc.mode=trained currently requires a Pi052 checkpoint; "
f"got policy type {policy_config.type!r}."
)
training_max_delay = int(getattr(policy_config, "rtc_training_max_delay", 0))
if training_max_delay <= 0:
raise ValueError(
"--inference.rtc.mode=trained requires a checkpoint trained with "
"--policy.rtc_training_max_delay > 0."
)
if rtc.execution_horizon < training_max_delay:
raise ValueError(
f"--inference.rtc.execution_horizon ({rtc.execution_horizon}) must be at least the "
f"checkpoint's rtc_training_max_delay ({training_max_delay})."
)
if inference_config.queue_threshold < training_max_delay:
raise ValueError(
f"--inference.queue_threshold ({inference_config.queue_threshold}) must be at least the "
f"checkpoint's rtc_training_max_delay ({training_max_delay})."
)
def _resolve_action_key_order( def _resolve_action_key_order(
policy_action_names: list[str] | None, dataset_action_names: list[str] policy_action_names: list[str] | None, dataset_action_names: list[str]
) -> list[str]: ) -> list[str]:
@@ -178,17 +207,8 @@ def build_rollout_context(
policy_config = cfg.policy policy_config = cfg.policy
policy_class = get_policy_class(policy_config.type) policy_class = get_policy_class(policy_config.type)
if is_rtc and cfg.inference.rtc.enabled and cfg.inference.rtc.mode == "trained": if is_rtc:
if policy_config.type != "pi052": _validate_trained_rtc_rollout_config(policy_config, cfg.inference)
raise ValueError(
"--inference.rtc.mode=trained currently requires a Pi052 checkpoint; "
f"got policy type {policy_config.type!r}."
)
if int(getattr(policy_config, "rtc_training_max_delay", 0)) <= 0:
raise ValueError(
"--inference.rtc.mode=trained requires a checkpoint trained with "
"--policy.rtc_training_max_delay > 0."
)
if hasattr(policy_config, "compile_model"): if hasattr(policy_config, "compile_model"):
policy_config.compile_model = cfg.use_torch_compile policy_config.compile_model = cfg.use_torch_compile
+26 -2
View File
@@ -57,10 +57,18 @@ _RTC_MAX_CONSECUTIVE_ERRORS: int = 10
_RTC_JOIN_TIMEOUT_S: float = 3.0 _RTC_JOIN_TIMEOUT_S: float = 3.0
class _TrainedRTCDelayExceededError(RuntimeError): class _FatalRTCInferenceError(RuntimeError):
"""Base class for RTC errors that cannot become valid after a retry."""
class _TrainedRTCDelayExceededError(_FatalRTCInferenceError):
"""Raised when measured latency exceeds a trained RTC checkpoint's support.""" """Raised when measured latency exceeds a trained RTC checkpoint's support."""
class _TrainedRTCPrefixUnavailableError(_FatalRTCInferenceError):
"""Raised when the queue cannot provide the prefix used for conditioning."""
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# RTC helpers # RTC helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -96,6 +104,16 @@ def _trained_rtc_chunk_can_merge(
return not has_previous_actions or measured_delay <= conditioned_delay return not has_previous_actions or measured_delay <= conditioned_delay
def _validate_trained_rtc_prefix_available(*, conditioned_delay: int, available_steps: int) -> None:
"""Reject hard-prefix inference when the real queue is shorter than its delay."""
if conditioned_delay > available_steps:
raise _TrainedRTCPrefixUnavailableError(
f"Trained RTC needs {conditioned_delay} committed prefix actions, but the queue has "
f"only {available_steps}. Increase --inference.queue_threshold and "
"--inference.rtc.execution_horizon."
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# RTCInferenceEngine # RTCInferenceEngine
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -296,6 +314,12 @@ class RTCInferenceEngine(InferenceEngine):
latency = latency_tracker.max() latency = latency_tracker.max()
delay = math.ceil(latency / time_per_chunk) if latency else 0 delay = math.ceil(latency / time_per_chunk) if latency else 0
if self._rtc_config.mode == "trained" and delay > 0:
available_steps = 0 if prev_actions is None else prev_actions.shape[0]
_validate_trained_rtc_prefix_available(
conditioned_delay=delay,
available_steps=available_steps,
)
obs_batch = build_dataset_frame(self._hw_features, obs, prefix="observation") obs_batch = build_dataset_frame(self._hw_features, obs, prefix="observation")
obs_batch = prepare_observation_for_inference( obs_batch = prepare_observation_for_inference(
@@ -372,7 +396,7 @@ class RTCInferenceEngine(InferenceEngine):
logger.debug("RTC inference latency=%.2fs, queue=%d", new_latency, queue.qsize()) logger.debug("RTC inference latency=%.2fs, queue=%d", new_latency, queue.qsize())
except _TrainedRTCDelayExceededError: except _FatalRTCInferenceError:
raise raise
except Exception as e: except Exception as e:
consecutive_errors += 1 consecutive_errors += 1
+13
View File
@@ -93,6 +93,19 @@ def test_rtc_processor_initialization_without_debug(rtc_config_debug_disabled):
assert processor.tracker is None assert processor.tracker is None
def test_rtc_processor_rejects_trained_mode_when_policy_does_not_support_it():
config = RTCConfig(mode="trained")
with pytest.raises(ValueError, match="requires a Pi052 checkpoint"):
RTCProcessor(config)
processor = RTCProcessor(config, trained_mode_supported=True)
assert processor.rtc_config.mode == "trained"
disabled = RTCProcessor(RTCConfig(enabled=False, mode="trained"))
assert disabled.rtc_config.enabled is False
# ====================== Tracker Proxy Methods Tests ====================== # ====================== Tracker Proxy Methods Tests ======================
+33
View File
@@ -17,6 +17,7 @@
from __future__ import annotations from __future__ import annotations
import dataclasses import dataclasses
from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock
import pytest import pytest
@@ -130,6 +131,38 @@ def test_trained_rtc_rejects_measured_delay_above_checkpoint_support():
) )
def test_trained_rtc_rejects_prefix_shorter_than_conditioned_delay():
from lerobot.rollout.inference.rtc import (
_TrainedRTCPrefixUnavailableError,
_validate_trained_rtc_prefix_available,
)
with pytest.raises(_TrainedRTCPrefixUnavailableError, match="only 2"):
_validate_trained_rtc_prefix_available(conditioned_delay=4, available_steps=2)
@pytest.mark.parametrize(
("execution_horizon", "queue_threshold", "match"),
[
(3, 4, "execution_horizon"),
(4, 3, "queue_threshold"),
],
)
def test_trained_rtc_rollout_requires_capacity_for_max_delay(execution_horizon, queue_threshold, match):
from lerobot.policies.rtc.configuration_rtc import RTCConfig
from lerobot.rollout.context import _validate_trained_rtc_rollout_config
from lerobot.rollout.inference import RTCInferenceConfig
policy_config = SimpleNamespace(type="pi052", rtc_training_max_delay=4)
inference_config = RTCInferenceConfig(
rtc=RTCConfig(mode="trained", execution_horizon=execution_horizon),
queue_threshold=queue_threshold,
)
with pytest.raises(ValueError, match=match):
_validate_trained_rtc_rollout_config(policy_config, inference_config)
def test_sentry_config_defaults(): def test_sentry_config_defaults():
from lerobot.rollout import SentryStrategyConfig from lerobot.rollout import SentryStrategyConfig