Compare commits

..

5 Commits

Author SHA1 Message Date
Pepijn 4045589246 docs(rtc): use 15-step training delay 2026-08-04 21:09:12 +02:00
Pepijn 80761e8a7a feat(pi05): add training-time RTC to SE3 branch 2026-08-04 21:07:26 +02:00
Pepijn b0cceb2a5f feat(pi05): support continuous 6D SE(3) actions (#4132)
* feat(pi05): add continuous 6D SE3 actions

* fix(processors): align rotation 6D with UMI rows
2026-07-27 12:16:33 +02:00
Pepijn 8e12a5351a feat(pi05): compose relative poses in SE3 2026-07-21 20:43:04 +02:00
Pepijn 7e1077f19a feat(pi05): synthesize relative proprioceptive history 2026-07-18 23:56:19 +02:00
27 changed files with 2014 additions and 97 deletions
+3 -2
View File
@@ -229,11 +229,12 @@ lerobot-rollout \
``` ```
| Flag | Description | | Flag | Description |
| ------------------------------------------- | -------------------------------------------------------------- | | ------------------------------------------- | ------------------------------------------------------------------------------- |
| `--inference.rtc.execution_horizon` | Steps to blend with previous chunk (default: varies by policy) | | `--inference.rtc.execution_horizon` | Steps to blend with previous chunk (default: varies by policy) |
| `--inference.rtc.mode` | `guided` (default) or trained-prefix `trained` for compatible Pi0.5 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.
+57 -1
View File
@@ -1,6 +1,6 @@
# Real-Time Chunking (RTC) # Real-Time Chunking (RTC)
Real-Time Chunking (RTC) is an inference-time method that allows large, flow-matching based robotic policies, such as [Pi0](./pi0), [Pi0.5](./pi05), and [SmolVLA](./smolvla), to produce smooth, continuous, and reactive motion despite having high inference latency. Real-Time Chunking (RTC) allows large, flow-matching based robotic policies, such as [Pi0](./pi0), [Pi0.5](./pi05), and [SmolVLA](./smolvla), to produce smooth, continuous, and reactive motion despite having high inference latency. LeRobot provides the original inference-time guided mode and, for compatible Pi0.5 checkpoints, training-time action conditioning with cheap hard-prefix inference.
These policies generate chunks of future actions (e.g., 50 steps at a time) instead of single actions. These policies generate chunks of future actions (e.g., 50 steps at a time) instead of single actions.
Because the models are large, producing each chunk takes longer than the time it takes the robot to execute it. Because the models are large, producing each chunk takes longer than the time it takes the robot to execute it.
@@ -92,6 +92,15 @@ for step in range(num_steps):
`RTCConfig` has the following parameters to tune: `RTCConfig` has the following parameters to tune:
**`mode`** selects the action-prefix conditioning method:
- `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 Pi0.5 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
@@ -111,6 +120,27 @@ RTCConfig(execution_horizon=10)
**`inference_delay`**: How many timesteps of inference latency your system has. This is passed to `predict_action_chunk()` rather than the config, since it may vary at runtime. **`inference_delay`**: How many timesteps of inference latency your system has. This is passed to `predict_action_chunk()` rather than the config, since it may vary at runtime.
## Training Pi0.5 for Trained RTC
Set the maximum prefix delay when fine-tuning Pi0.5:
```bash
lerobot-train \
--policy.type=pi05 \
--policy.pretrained_path=lerobot/pi05_base \
--policy.rtc_training_max_delay=15 \
--dataset.repo_id=${HF_USERNAME}/dataset_repo_id \
--output_dir=outputs/pi05_training_rtc
```
Each example samples a clean prefix from zero through the configured maximum;
the flow loss is computed only on the remaining postfix. Choose the maximum as
approximately `ceil(p95 end-to-end inference latency * control frequency)` and
keep it smaller than `policy.chunk_size`. Setting the value to zero preserves
ordinary Pi0.5 training and existing checkpoints remain compatible with guided
RTC. At a 50 Hz control rate, 15 steps cover up to 300 ms of end-to-end
inference latency.
## Testing RTC Offline ## Testing RTC Offline
Before running on a real robot, test RTC with dataset samples to visualize how it works: Before running on a real robot, test RTC with dataset samples to visualize how it works:
@@ -124,6 +154,10 @@ python examples/rtc/eval_dataset.py \
--device=cuda --device=cuda
``` ```
Add `--rtc.mode=trained` when evaluating a compatible training-time RTC Pi0.5
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">
@@ -141,6 +175,7 @@ lerobot-rollout \
--strategy.type=base \ --strategy.type=base \
--policy.path=${HF_USERNAME}/policy_repo_id \ --policy.path=${HF_USERNAME}/policy_repo_id \
--inference.type=rtc \ --inference.type=rtc \
--inference.rtc.mode=guided \
--inference.rtc.execution_horizon=10 \ --inference.rtc.execution_horizon=10 \
--inference.rtc.max_guidance_weight=10.0 \ --inference.rtc.max_guidance_weight=10.0 \
--robot.type=so100_follower \ --robot.type=so100_follower \
@@ -151,6 +186,25 @@ lerobot-rollout \
--device=cuda --device=cuda
``` ```
For a training-time RTC Pi0.5 checkpoint, change the mode to `trained`. The
checkpoint records its maximum supported delay, and rollout validates measured
latency against it:
```bash
lerobot-rollout \
--strategy.type=base \
--policy.path=${HF_USERNAME}/pi05_training_rtc \
--inference.type=rtc \
--inference.rtc.mode=trained \
--inference.rtc.execution_horizon=15 \
--inference.queue_threshold=15 \
--robot.type=so100_follower \
--robot.port=/dev/tty.usbmodem58FA0834591 \
--task="Move green small object into the purple platform" \
--duration=120 \
--device=cuda
```
## How It Differs from the Async Inference in LeRobot ## How It Differs from the Async Inference in LeRobot
Both RTC and [async inference](./async) improve real-time robot control, but they solve different problems. Both RTC and [async inference](./async) improve real-time robot control, but they solve different problems.
@@ -189,3 +243,5 @@ See `examples/rtc/eval_dataset.py` for a complete example of offline RTC visuali
- [Smooth-As-Butter Robot Policies](https://alexander-soare.github.io/robotics/2025/08/05/smooth-as-butter-robot-policies.html) - Excellent technical explanation with real robot results - [Smooth-As-Butter Robot Policies](https://alexander-soare.github.io/robotics/2025/08/05/smooth-as-butter-robot-policies.html) - Excellent technical explanation with real robot results
- [Physical Intelligence - Real-Time Chunking](https://www.physicalintelligence.company/research/real_time_chunking) - Original paper and research - [Physical Intelligence - Real-Time Chunking](https://www.physicalintelligence.company/research/real_time_chunking) - Original paper and research
- [Kinetix RTC Implementation](https://github.com/Physical-Intelligence/real-time-chunking-kinetix) - Reference implementation from Physical Intelligence - [Kinetix RTC Implementation](https://github.com/Physical-Intelligence/real-time-chunking-kinetix) - Reference implementation from Physical Intelligence
- [Training-Time Action Conditioning](https://arxiv.org/abs/2512.05964) - Efficient RTC with clean-prefix conditioning during training
- [RLDX-1](https://github.com/RLWRLD/RLDX-1) - PyTorch reference used for the training-time RTC integration
+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,
+93 -5
View File
@@ -18,8 +18,13 @@ from __future__ import annotations
import logging import logging
import numpy as np import numpy as np
import torch
from lerobot.processor import RelativeActionsProcessorStep from lerobot.processor import (
RelativeActionsProcessorStep,
relative_action_output_dim,
to_relative_actions,
)
from lerobot.utils.constants import ACTION, OBS_STATE from lerobot.utils.constants import ACTION, OBS_STATE
from .io_utils import load_image_as_numpy from .io_utils import load_image_as_numpy
@@ -660,17 +665,29 @@ def _compute_relative_chunk_batch(
all_states: np.ndarray, all_states: np.ndarray,
chunk_size: int, chunk_size: int,
relative_mask: np.ndarray, relative_mask: np.ndarray,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> np.ndarray: ) -> np.ndarray:
"""Vectorised relative-action computation for a batch of start indices. """Vectorised relative-action computation for a batch of start indices.
Returns an ``(N * chunk_size, action_dim)`` float32 array. Returns an ``(N * chunk_size, model_action_dim)`` float32 array.
""" """
if len(start_indices) == 0: if len(start_indices) == 0:
return np.empty((0, all_actions.shape[1]), dtype=np.float32) output_dim = relative_action_output_dim(all_actions.shape[1], pose_representation, se3_pose_groups)
return np.empty((0, output_dim), dtype=np.float32)
offsets = np.arange(chunk_size) offsets = np.arange(chunk_size)
frame_idx = start_indices[:, None] + offsets[None, :] frame_idx = start_indices[:, None] + offsets[None, :]
chunks = all_actions[frame_idx].copy() chunks = all_actions[frame_idx].copy()
states = all_states[start_indices] states = all_states[start_indices]
if pose_representation in {"se3", "se3_6d"}:
converted = to_relative_actions(
torch.from_numpy(chunks),
torch.from_numpy(states),
relative_mask.astype(bool).tolist(),
pose_representation=pose_representation,
se3_pose_groups=se3_pose_groups,
)
return converted.numpy().reshape(-1, converted.shape[-1])
mask_dim = len(relative_mask) mask_dim = len(relative_mask)
chunks[:, :, :mask_dim] -= states[:, None, :mask_dim] * relative_mask[None, None, :] chunks[:, :, :mask_dim] -= states[:, None, :mask_dim] * relative_mask[None, None, :]
return chunks.reshape(-1, all_actions.shape[1]) return chunks.reshape(-1, all_actions.shape[1])
@@ -682,6 +699,9 @@ def compute_relative_action_stats(
chunk_size: int, chunk_size: int,
exclude_joints: list[str] | None = None, exclude_joints: list[str] | None = None,
num_workers: int = 0, num_workers: int = 0,
state_from_action: bool = False,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> dict[str, np.ndarray]: ) -> dict[str, np.ndarray]:
"""Compute normalization statistics for relative actions over the full dataset. """Compute normalization statistics for relative actions over the full dataset.
@@ -700,6 +720,9 @@ def compute_relative_action_stats(
num_workers: Number of parallel threads for computation. Values ≤1 num_workers: Number of parallel threads for computation. Values ≤1
mean single-threaded. Numpy releases the GIL so threads give mean single-threaded. Numpy releases the GIL so threads give
real parallelism here. real parallelism here.
state_from_action: Use the current absolute action as state. This is
intended for state-less pose datasets where each action row is the
synchronized measured robot pose.
Returns: Returns:
Statistics dict with keys "mean", "std", "min", "max", "q01", …, "q99". Statistics dict with keys "mean", "std", "min", "max", "q01", …, "q99".
@@ -722,7 +745,7 @@ def compute_relative_action_stats(
logging.info("Loading action/state data for relative action stats...") logging.info("Loading action/state data for relative action stats...")
all_actions = np.array(hf_dataset[ACTION], dtype=np.float32) all_actions = np.array(hf_dataset[ACTION], dtype=np.float32)
all_states = np.array(hf_dataset[OBS_STATE], dtype=np.float32) all_states = all_actions if state_from_action else np.array(hf_dataset[OBS_STATE], dtype=np.float32)
episode_indices = np.array(hf_dataset["episode_index"]) episode_indices = np.array(hf_dataset["episode_index"])
valid_starts = _get_valid_chunk_starts(episode_indices, chunk_size) valid_starts = _get_valid_chunk_starts(episode_indices, chunk_size)
@@ -754,6 +777,8 @@ def compute_relative_action_stats(
all_states, all_states,
chunk_size, chunk_size,
relative_mask, relative_mask,
pose_representation,
se3_pose_groups,
) )
for batch in batches for batch in batches
] ]
@@ -762,7 +787,15 @@ def compute_relative_action_stats(
else: else:
for batch in batches: for batch in batches:
running_stats.update( running_stats.update(
_compute_relative_chunk_batch(batch, all_actions, all_states, chunk_size, relative_mask) _compute_relative_chunk_batch(
batch,
all_actions,
all_states,
chunk_size,
relative_mask,
pose_representation,
se3_pose_groups,
)
) )
stats = running_stats.get_statistics() stats = running_stats.get_statistics()
@@ -777,3 +810,58 @@ def compute_relative_action_stats(
) )
return stats return stats
def compute_state_history_stats(
hf_dataset,
features: dict,
history_steps: int,
exclude_joints: list[str] | None = None,
relative: bool = False,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> dict[str, np.ndarray]:
"""Compute stats for flattened state history synthesized from absolute actions.
History is left-padded with the first action of each episode, matching dataset
boundary padding. When ``relative`` is enabled, every history pose is expressed
relative to its newest pose while excluded dimensions remain absolute.
"""
if history_steps < 1:
raise ValueError("history_steps must be at least 1")
if exclude_joints is None:
exclude_joints = []
actions = np.asarray(hf_dataset[ACTION], dtype=np.float32)
episode_indices = np.asarray(hf_dataset["episode_index"])
sample_indices = np.arange(len(actions))
episode_starts = np.maximum.accumulate(
np.where(
np.concatenate(([True], episode_indices[1:] != episode_indices[:-1])),
sample_indices,
0,
)
)
offsets = np.arange(-(history_steps - 1), 1)
history_indices = np.maximum(sample_indices[:, None] + offsets[None, :], episode_starts[:, None])
history = actions[history_indices].copy()
if relative:
state_dim = actions.shape[-1]
names = features.get(ACTION, {}).get("names")
mask_step = RelativeActionsProcessorStep(
enabled=True,
exclude_joints=exclude_joints,
action_names=names,
)
mask = mask_step._build_mask(state_dim)
history = to_relative_actions(
torch.from_numpy(history),
torch.from_numpy(history[:, -1].copy()),
mask,
pose_representation=pose_representation,
se3_pose_groups=se3_pose_groups,
).numpy()
flattened = history.reshape(len(history), -1)
return get_feature_stats(flattened, axis=0, keepdims=False)
+37 -1
View File
@@ -54,6 +54,7 @@ from .compute_stats import (
aggregate_stats, aggregate_stats,
compute_episode_stats, compute_episode_stats,
compute_relative_action_stats, compute_relative_action_stats,
compute_state_history_stats,
) )
from .dataset_metadata import LeRobotDatasetMetadata from .dataset_metadata import LeRobotDatasetMetadata
from .image_writer import write_image from .image_writer import write_image
@@ -1566,6 +1567,12 @@ def recompute_stats(
relative_exclude_joints: list[str] | None = None, relative_exclude_joints: list[str] | None = None,
chunk_size: int = 50, chunk_size: int = 50,
num_workers: int = 0, num_workers: int = 0,
state_from_action: bool = False,
state_history_steps: int = 1,
relative_state_history: bool = False,
relative_state_exclude_joints: list[str] | None = None,
relative_pose_representation: str = "componentwise",
relative_se3_pose_groups: list[list[int]] | None = None,
) -> LeRobotDataset: ) -> LeRobotDataset:
"""Recompute stats.json from scratch by iterating all episodes. """Recompute stats.json from scratch by iterating all episodes.
@@ -1583,6 +1590,16 @@ def recompute_stats(
``policy.chunk_size``. Only used when ``relative_action=True``. ``policy.chunk_size``. Only used when ``relative_action=True``.
num_workers: Number of parallel threads for relative action stats computation. num_workers: Number of parallel threads for relative action stats computation.
Values ≤1 mean single-threaded. Only used when ``relative_action=True``. Values ≤1 mean single-threaded. Only used when ``relative_action=True``.
state_from_action: Use absolute action rows as synthetic state while
computing relative-action stats, and write their absolute statistics
under ``observation.state``.
state_history_steps: Number of consecutive synthesized state samples.
relative_state_history: Express state history relative to its newest pose.
relative_state_exclude_joints: State dimensions to retain as absolute.
relative_pose_representation: ``componentwise`` for legacy subtraction,
``se3`` for composition with an axis-angle output, or ``se3_6d`` for
composition with a continuous two-column rotation output.
relative_se3_pose_groups: Six-index xyz+rotation-vector pose groups.
Returns: Returns:
The same dataset with updated stats. The same dataset with updated stats.
@@ -1606,7 +1623,21 @@ def recompute_stats(
# (matching what the model sees during training) and skip action in the # (matching what the model sees during training) and skip action in the
# per-episode pass below. # per-episode pass below.
relative_action_stats = None relative_action_stats = None
if relative_action and ACTION in features and OBS_STATE in features: synthetic_state_stats = None
if state_from_action:
if ACTION not in features:
raise ValueError("state_from_action requires an action feature")
synthetic_state_stats = compute_state_history_stats(
dataset.hf_dataset,
features,
history_steps=state_history_steps,
exclude_joints=relative_state_exclude_joints,
relative=relative_state_history,
pose_representation=relative_pose_representation,
se3_pose_groups=relative_se3_pose_groups,
)
if relative_action and ACTION in features and (OBS_STATE in features or state_from_action):
if relative_exclude_joints is None: if relative_exclude_joints is None:
relative_exclude_joints = ["gripper"] relative_exclude_joints = ["gripper"]
relative_action_stats = compute_relative_action_stats( relative_action_stats = compute_relative_action_stats(
@@ -1615,6 +1646,9 @@ def recompute_stats(
chunk_size=chunk_size, chunk_size=chunk_size,
exclude_joints=relative_exclude_joints, exclude_joints=relative_exclude_joints,
num_workers=num_workers, num_workers=num_workers,
state_from_action=state_from_action,
pose_representation=relative_pose_representation,
se3_pose_groups=relative_se3_pose_groups,
) )
features_to_compute.pop(ACTION, None) features_to_compute.pop(ACTION, None)
@@ -1654,6 +1688,8 @@ def recompute_stats(
if relative_action_stats is not None: if relative_action_stats is not None:
new_stats[ACTION] = relative_action_stats new_stats[ACTION] = relative_action_stats
if synthetic_state_stats is not None:
new_stats[OBS_STATE] = synthetic_state_stats
# Merge: keep existing stats for features we didn't recompute # Merge: keep existing stats for features we didn't recompute
if dataset.meta.stats: if dataset.meta.stats:
@@ -55,9 +55,25 @@ class PI05Config(PreTrainedConfig):
relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"]) relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
# Populated at runtime from dataset metadata by make_policy. # Populated at runtime from dataset metadata by make_policy.
action_feature_names: list[str] | None = None action_feature_names: list[str] | None = None
# ``se3`` uses inv(T_current) @ T_target for each xyz+rotation-vector pose group.
# ``se3_6d`` uses the same composition and expands each relative rotation
# vector to the continuous first-two-row 6-D rotation representation.
# ``componentwise`` preserves the legacy action - state behavior.
relative_pose_representation: str = "componentwise"
relative_se3_pose_groups: list[list[int]] = field(default_factory=lambda: [list(range(6))])
# Build proprioception from absolute action samples when the dataset has no
# observation.state. With history_steps=2, training samples request t-1 as
# well as the normal t..t+chunk_size-1 action targets.
state_from_action: bool = False
proprioception_history_steps: int = 1
use_relative_state_history: bool = False
relative_state_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
# Real-Time Chunking (RTC) configuration # Real-Time Chunking (RTC) configuration
rtc_config: RTCConfig | None = None rtc_config: RTCConfig | None = None
# Maximum clean action-prefix length sampled during training. Zero disables trained RTC.
rtc_training_max_delay: int = 0
image_resolution: tuple[int, int] = ( image_resolution: tuple[int, int] = (
DEFAULT_IMAGE_SIZE, DEFAULT_IMAGE_SIZE,
@@ -111,6 +127,11 @@ class PI05Config(PreTrainedConfig):
raise ValueError( raise ValueError(
f"n_action_steps ({self.n_action_steps}) cannot be greater than chunk_size ({self.chunk_size})" f"n_action_steps ({self.n_action_steps}) cannot be greater than chunk_size ({self.chunk_size})"
) )
if not 0 <= self.rtc_training_max_delay < self.chunk_size:
raise ValueError(
"rtc_training_max_delay must satisfy "
f"0 <= delay < chunk_size ({self.chunk_size}), got {self.rtc_training_max_delay}"
)
if self.paligemma_variant not in ["gemma_300m", "gemma_2b"]: if self.paligemma_variant not in ["gemma_300m", "gemma_2b"]:
raise ValueError(f"Invalid paligemma_variant: {self.paligemma_variant}") raise ValueError(f"Invalid paligemma_variant: {self.paligemma_variant}")
@@ -121,6 +142,25 @@ class PI05Config(PreTrainedConfig):
if self.dtype not in ["bfloat16", "float32"]: if self.dtype not in ["bfloat16", "float32"]:
raise ValueError(f"Invalid dtype: {self.dtype}") raise ValueError(f"Invalid dtype: {self.dtype}")
if self.proprioception_history_steps < 1:
raise ValueError("proprioception_history_steps must be at least 1")
if self.relative_pose_representation not in {"componentwise", "se3", "se3_6d"}:
raise ValueError(
"relative_pose_representation must be 'componentwise', 'se3', or 'se3_6d', got "
f"{self.relative_pose_representation!r}"
)
for group in self.relative_se3_pose_groups:
if len(group) != 6 or len(set(group)) != 6 or any(index < 0 for index in group):
raise ValueError(f"Invalid six-index SE(3) pose group: {group}")
if self.relative_pose_representation == "se3_6d" and group != list(range(group[0], group[0] + 6)):
raise ValueError("se3_6d pose groups must contain six contiguous ascending indices")
if self.relative_pose_representation in {"se3", "se3_6d"} and not self.relative_se3_pose_groups:
raise ValueError(
f"relative_pose_representation={self.relative_pose_representation!r} "
"requires relative_se3_pose_groups"
)
def validate_features(self) -> None: def validate_features(self) -> None:
"""Validate and set up input/output features.""" """Validate and set up input/output features."""
for i in range(self.empty_cameras): for i in range(self.empty_cameras):
@@ -131,19 +171,54 @@ class PI05Config(PreTrainedConfig):
) )
self.input_features[key] = empty_camera self.input_features[key] = empty_camera
if OBS_STATE not in self.input_features:
state_feature = PolicyFeature(
type=FeatureType.STATE,
shape=(self.max_state_dim,), # Padded to max_state_dim
)
self.input_features[OBS_STATE] = state_feature
if ACTION not in self.output_features: if ACTION not in self.output_features:
action_feature = PolicyFeature( action_feature = PolicyFeature(
type=FeatureType.ACTION, type=FeatureType.ACTION,
shape=(self.max_action_dim,), # Padded to max_action_dim shape=(self.max_action_dim,), # Padded to max_action_dim
) )
self.output_features[ACTION] = action_feature self.output_features[ACTION] = action_feature
elif self.relative_pose_representation == "se3_6d":
action_feature = self.output_features[ACTION]
source_dim = (
len(self.action_feature_names)
if self.action_feature_names is not None
else action_feature.shape[-1]
)
model_dim = source_dim + 3 * len(self.relative_se3_pose_groups)
if action_feature.shape[-1] == source_dim:
self.output_features[ACTION] = PolicyFeature(
type=action_feature.type,
shape=(model_dim,),
)
elif action_feature.shape[-1] != model_dim:
raise ValueError(
"se3_6d action feature has incompatible width: "
f"source={source_dim}, expected model width={model_dim}, "
f"got={action_feature.shape[-1]}"
)
if model_dim > self.max_action_dim:
raise ValueError(
f"se3_6d action width {model_dim} exceeds max_action_dim={self.max_action_dim}"
)
if OBS_STATE not in self.input_features:
state_shape = (self.max_state_dim,)
if self.state_from_action and ACTION in self.output_features:
state_shape = self.output_features[ACTION].shape
state_feature = PolicyFeature(
type=FeatureType.STATE,
shape=state_shape,
)
self.input_features[OBS_STATE] = state_feature
state_dim = self.input_features[OBS_STATE].shape[-1]
history_state_dim = state_dim * self.proprioception_history_steps
if history_state_dim > self.max_state_dim:
raise ValueError(
"Flattened proprioception history exceeds max_state_dim: "
f"{state_dim} * {self.proprioception_history_steps} = {history_state_dim} > "
f"{self.max_state_dim}"
)
def get_optimizer_preset(self) -> AdamWConfig: def get_optimizer_preset(self) -> AdamWConfig:
return AdamWConfig( return AdamWConfig(
@@ -168,7 +243,8 @@ class PI05Config(PreTrainedConfig):
@property @property
def action_delta_indices(self) -> list: def action_delta_indices(self) -> list:
return list(range(self.chunk_size)) history_prefix = self.proprioception_history_steps - 1 if self.state_from_action else 0
return list(range(-history_prefix, self.chunk_size))
@property @property
def reward_delta_indices(self) -> None: def reward_delta_indices(self) -> None:
+167 -22
View File
@@ -66,6 +66,107 @@ class ActionSelectKwargs(TypedDict, total=False):
execution_horizon: int | None execution_horizon: int | None
def _prepare_trained_rtc_prefix(
x_t: Tensor,
prev_chunk_left_over: Tensor | None,
inference_delay: int,
training_max_delay: int,
) -> tuple[Tensor | None, Tensor | None]:
"""Pad and validate a hard prefix for training-time RTC inference."""
if prev_chunk_left_over is None or inference_delay <= 0:
return None, None
if training_max_delay <= 0:
raise ValueError(
"RTC mode='trained' requires a checkpoint trained with policy.rtc_training_max_delay > 0."
)
if inference_delay > training_max_delay:
raise ValueError(
f"Measured RTC inference delay ({inference_delay}) exceeds the checkpoint's "
f"rtc_training_max_delay ({training_max_delay})."
)
if inference_delay >= x_t.shape[1]:
raise ValueError(
f"RTC inference delay ({inference_delay}) must be smaller than chunk_size ({x_t.shape[1]})."
)
previous = prev_chunk_left_over.to(device=x_t.device, dtype=x_t.dtype)
if not torch.isfinite(previous).all():
raise ValueError("RTC prefix contains NaN or Inf values.")
if previous.ndim == 2:
previous = previous.unsqueeze(0)
if previous.ndim != 3:
raise ValueError(f"Expected RTC prefix shape (B, T, A), got {tuple(previous.shape)}")
if previous.shape[0] == 1 and x_t.shape[0] > 1:
previous = previous.expand(x_t.shape[0], -1, -1)
if previous.shape[0] != x_t.shape[0]:
raise ValueError(
f"RTC prefix batch size ({previous.shape[0]}) does not match policy batch ({x_t.shape[0]})."
)
if previous.shape[1] < inference_delay:
raise ValueError(f"RTC prefix has {previous.shape[1]} steps, but inference_delay={inference_delay}.")
if previous.shape[2] > x_t.shape[2]:
raise ValueError(
f"RTC prefix action dimension ({previous.shape[2]}) exceeds model dimension ({x_t.shape[2]})."
)
padded_prefix = torch.zeros_like(x_t)
padded_prefix[:, :inference_delay, : previous.shape[2]] = previous[:, :inference_delay]
prefix_mask = torch.arange(x_t.shape[1], device=x_t.device) < inference_delay
prefix_mask = prefix_mask[None, :, None].expand(x_t.shape[0], -1, x_t.shape[2])
return padded_prefix, prefix_mask
def _sample_training_rtc_prefix_mask(
batch_size: int,
action_horizon: int,
max_delay: int,
device: torch.device,
) -> Tensor | None:
"""Sample a clean action-prefix length independently for each training example."""
if max_delay <= 0:
return None
delays = torch.randint(0, max_delay + 1, (batch_size,), device=device)
positions = torch.arange(action_horizon, device=device)
return positions.unsqueeze(0) < delays.unsqueeze(1)
def _build_flow_matching_inputs(
actions: Tensor,
noise: Tensor,
time: Tensor,
prefix_mask: Tensor | None,
) -> tuple[Tensor, Tensor]:
"""Keep the sampled RTC prefix clean while noising the remaining action chunk."""
if prefix_mask is None:
model_time = time
expanded_time = time[:, None, None]
else:
model_time = time[:, None].expand_as(prefix_mask)
model_time = torch.where(prefix_mask, torch.zeros_like(model_time), model_time)
expanded_time = model_time.unsqueeze(-1)
x_t = expanded_time * noise + (1 - expanded_time) * actions
return x_t, model_time
def _reduce_training_rtc_loss(
losses: Tensor,
prefix_mask: Tensor | None,
reduction: str,
) -> Tensor:
"""Average flow loss over predicted postfix actions, excluding the clean RTC prefix."""
if reduction not in {"mean", "none"}:
raise ValueError(f"Unsupported loss reduction: {reduction!r}")
if prefix_mask is None:
return losses.mean() if reduction == "mean" else losses.mean(dim=(1, 2))
postfix_mask = (~prefix_mask).unsqueeze(-1).expand_as(losses)
if reduction == "none":
numerator = (losses * postfix_mask).sum(dim=(1, 2))
denominator = postfix_mask.sum(dim=(1, 2))
return numerator / denominator.clamp(min=1)
return (losses * postfix_mask).sum() / postfix_mask.sum().clamp(min=1)
def get_safe_dtype(target_dtype, device_type): def get_safe_dtype(target_dtype, device_type):
"""Get a safe dtype for the given device type.""" """Get a safe dtype for the given device type."""
if device_type == "mps" and target_dtype == torch.float64: if device_type == "mps" and target_dtype == torch.float64:
@@ -82,21 +183,20 @@ def get_safe_dtype(target_dtype, device_type):
def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy) def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu" time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
) -> Tensor: ) -> Tensor:
"""Computes sine-cosine positional embedding vectors for scalar positions.""" """Compute sine-cosine embeddings for scalar or per-action positions."""
if dimension % 2 != 0: if dimension % 2 != 0:
raise ValueError(f"dimension ({dimension}) must be divisible by 2") raise ValueError(f"dimension ({dimension}) must be divisible by 2")
if time.ndim != 1: if time.ndim not in (1, 2):
raise ValueError("The time tensor is expected to be of shape `(batch_size, )`.") raise ValueError("The time tensor must have shape (batch_size,) or (batch_size, action_horizon).")
dtype = get_safe_dtype(torch.float64, device.type) dtype = get_safe_dtype(torch.float64, device.type)
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device) fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
period = min_period * (max_period / min_period) ** fraction period = min_period * (max_period / min_period) ** fraction
# Compute the outer product
scaling_factor = 1.0 / period * 2 * math.pi scaling_factor = 1.0 / period * 2 * math.pi
sin_input = scaling_factor[None, :] * time[:, None] sin_input = time[..., None] * scaling_factor
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1) return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=-1)
def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy) def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy)
@@ -739,14 +839,23 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
return embs, pad_masks, att_masks, adarms_cond return embs, pad_masks, att_masks, adarms_cond
def forward(self, images, img_masks, tokens, masks, actions, noise, time) -> Tensor: def forward(
self,
images,
img_masks,
tokens,
masks,
actions,
noise,
time,
prefix_mask: Tensor | None = None,
) -> Tensor:
"""Do a full training forward pass and compute the loss.""" """Do a full training forward pass and compute the loss."""
time_expanded = time[:, None, None] x_t, model_time = _build_flow_matching_inputs(actions, noise, time, prefix_mask)
x_t = time_expanded * noise + (1 - time_expanded) * actions
u_t = noise - actions u_t = noise - actions
prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(images, img_masks, tokens, masks) prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(images, img_masks, tokens, masks)
suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, time) suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, model_time)
if ( if (
self.paligemma_with_expert.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype self.paligemma_with_expert.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype
@@ -833,11 +942,35 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
dt = -1.0 / num_steps dt = -1.0 / num_steps
x_t = noise x_t = noise
rtc_mode = "guided"
trained_prefix = trained_prefix_mask = None
if self._rtc_enabled():
rtc_mode = self.rtc_processor.rtc_config.mode
if rtc_mode == "trained":
training_max_delay = int(getattr(self.config, "rtc_training_max_delay", 0))
if training_max_delay <= 0:
raise ValueError(
"RTC mode='trained' requires a checkpoint trained with "
"policy.rtc_training_max_delay > 0."
)
trained_prefix, trained_prefix_mask = _prepare_trained_rtc_prefix(
x_t,
kwargs.get("prev_chunk_left_over"),
int(kwargs.get("inference_delay") or 0),
training_max_delay,
)
for step in range(num_steps): for step in range(num_steps):
time = 1.0 + step * dt time = 1.0 + step * dt
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize) time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor): denoise_timestep = time_tensor
if trained_prefix is not None:
x_t = torch.where(trained_prefix_mask, trained_prefix, x_t)
denoise_timestep = time_tensor[:, None].expand(bsize, x_t.shape[1]).clone()
denoise_timestep[trained_prefix_mask[..., 0]] = 0.0
def denoise_step_partial_call(input_x_t, current_timestep=denoise_timestep):
return self.denoise_step( return self.denoise_step(
prefix_pad_masks=prefix_pad_masks, prefix_pad_masks=prefix_pad_masks,
past_key_values=past_key_values, past_key_values=past_key_values,
@@ -845,7 +978,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
timestep=current_timestep, timestep=current_timestep,
) )
if self._rtc_enabled(): if self._rtc_enabled() and rtc_mode == "guided":
inference_delay = kwargs.get("inference_delay") inference_delay = kwargs.get("inference_delay")
prev_chunk_left_over = kwargs.get("prev_chunk_left_over") prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
execution_horizon = kwargs.get("execution_horizon") execution_horizon = kwargs.get("execution_horizon")
@@ -862,6 +995,8 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
v_t = denoise_step_partial_call(x_t) v_t = denoise_step_partial_call(x_t)
x_t = x_t + dt * v_t x_t = x_t + dt * v_t
if trained_prefix is not None:
x_t = torch.where(trained_prefix_mask, trained_prefix, x_t)
if self.rtc_processor is not None and self.rtc_processor.is_debug_enabled(): if self.rtc_processor is not None and self.rtc_processor.is_debug_enabled():
self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t) self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
@@ -1137,7 +1272,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:
@@ -1269,26 +1407,33 @@ class PI05Policy(PreTrainedPolicy):
noise = self.model.sample_noise(actions.shape, actions.device) noise = self.model.sample_noise(actions.shape, actions.device)
time = self.model.sample_time(actions.shape[0], actions.device) time = self.model.sample_time(actions.shape[0], actions.device)
prefix_mask = _sample_training_rtc_prefix_mask(
actions.shape[0],
actions.shape[1],
self.config.rtc_training_max_delay,
actions.device,
)
# Compute loss (no separate state needed for PI05) # Compute loss (no separate state needed for PI05)
losses = self.model.forward(images, img_masks, tokens, masks, actions, noise, time) losses = self.model.forward(images, img_masks, tokens, masks, actions, noise, time, prefix_mask)
# Truncate losses to actual action dimensions # Truncate losses to actual action dimensions
original_action_dim = self.config.output_features[ACTION].shape[0] original_action_dim = self.config.output_features[ACTION].shape[0]
losses = losses[:, :, :original_action_dim] losses = losses[:, :, :original_action_dim]
loss_dict = { if prefix_mask is None:
"loss_per_dim": losses.mean(dim=[0, 1]).detach().cpu().numpy().tolist(), loss_per_dim = losses.mean(dim=(0, 1))
} else:
postfix_mask = (~prefix_mask).unsqueeze(-1).expand_as(losses)
loss_per_dim = (losses * postfix_mask).sum(dim=(0, 1)) / postfix_mask.sum(dim=(0, 1)).clamp(min=1)
loss_dict = {"loss_per_dim": loss_per_dim.detach().cpu().numpy().tolist()}
if reduction == "none": if reduction == "none":
# Return per-sample losses (B,) by averaging over time and action dims per_sample_loss = _reduce_training_rtc_loss(losses, prefix_mask, reduction="none")
per_sample_loss = losses.mean(dim=(1, 2))
loss_dict["loss"] = per_sample_loss.mean().item() loss_dict["loss"] = per_sample_loss.mean().item()
return per_sample_loss, loss_dict return per_sample_loss, loss_dict
else:
# Default: return scalar mean loss loss = _reduce_training_rtc_loss(losses, prefix_mask, reduction="mean")
loss = losses.mean()
loss_dict["loss"] = loss.item() loss_dict["loss"] = loss.item()
return loss, loss_dict return loss, loss_dict
+176 -1
View File
@@ -15,7 +15,7 @@
# limitations under the License. # limitations under the License.
from copy import deepcopy from copy import deepcopy
from dataclasses import dataclass from dataclasses import dataclass, field
from typing import Any from typing import Any
import numpy as np import numpy as np
@@ -36,6 +36,8 @@ from lerobot.processor import (
TokenizerProcessorStep, TokenizerProcessorStep,
UnnormalizerProcessorStep, UnnormalizerProcessorStep,
policy_action_to_transition, policy_action_to_transition,
relative_action_output_dim,
to_relative_actions,
transition_to_policy_action, transition_to_policy_action,
) )
from lerobot.types import EnvTransition, TransitionKey from lerobot.types import EnvTransition, TransitionKey
@@ -48,6 +50,164 @@ from lerobot.utils.constants import (
from .configuration_pi05 import PI05Config from .configuration_pi05 import PI05Config
@ProcessorStepRegistry.register(name="pi05_state_from_action_processor_step")
@dataclass
class Pi05StateFromActionProcessorStep(ProcessorStep):
"""Synthesize proprioception from absolute actions in state-less datasets.
The dataset loader supplies ``history_steps - 1`` actions before the normal
target chunk. Those leading samples and action(t) become state history; only
the leading samples are then removed from the action targets.
"""
enabled: bool = False
history_steps: int = 1
_inference_history: torch.Tensor | None = field(default=None, init=False, repr=False)
def __call__(self, transition: EnvTransition) -> EnvTransition:
if not self.enabled:
return transition
observation = transition.get(TransitionKey.OBSERVATION, {})
observed_state = observation.get(OBS_STATE)
if observed_state is not None:
# At inference the robot normally provides only the current state and
# there is no action target. Build a rolling history in the processor.
if transition.get(TransitionKey.ACTION) is None and observed_state.ndim == 2:
if self._inference_history is None:
self._inference_history = observed_state.unsqueeze(1).repeat(1, self.history_steps, 1)
else:
self._inference_history = torch.cat(
[self._inference_history[:, 1:], observed_state.unsqueeze(1)], dim=1
)
new_transition = transition.copy()
new_observation = dict(observation)
new_observation[OBS_STATE] = self._inference_history.clone()
new_transition[TransitionKey.OBSERVATION] = new_observation
return new_transition
return transition
action = transition.get(TransitionKey.ACTION)
if action is None:
raise ValueError("Cannot synthesize PI0.5 state without action")
if action.ndim != 3:
raise ValueError(f"Expected batched action chunks with shape (B, T, D), got {action.shape}")
if action.shape[1] < self.history_steps:
raise ValueError(
f"Action chunk has {action.shape[1]} steps, fewer than history_steps={self.history_steps}"
)
new_transition = transition.copy()
new_observation = dict(observation)
state = action[:, : self.history_steps].clone()
if self.history_steps == 1:
state = state[:, 0]
new_observation[OBS_STATE] = state
new_transition[TransitionKey.OBSERVATION] = new_observation
new_transition[TransitionKey.ACTION] = action[:, self.history_steps - 1 :]
return new_transition
def get_config(self) -> dict[str, Any]:
return {"enabled": self.enabled, "history_steps": self.history_steps}
def reset(self) -> None:
self._inference_history = None
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
@ProcessorStepRegistry.register(name="pi05_flatten_state_history_processor_step")
@dataclass
class Pi05FlattenStateHistoryProcessorStep(ProcessorStep):
"""Optionally relativize raw state history, then flatten it for PI0.5."""
history_steps: int = 1
max_state_dim: int = 32
relative: bool = False
exclude_joints: list[str] = field(default_factory=list)
state_names: list[str] | None = None
pose_representation: str = "componentwise"
se3_pose_groups: list[list[int]] = field(default_factory=list)
def __call__(self, transition: EnvTransition) -> EnvTransition:
observation = transition.get(TransitionKey.OBSERVATION, {})
state = observation.get(OBS_STATE)
if state is None:
raise ValueError("State is required for PI05")
if self.history_steps == 1 and state.ndim == 2:
state = state.unsqueeze(1)
if state.ndim != 3 or state.shape[1] != self.history_steps:
raise ValueError(
f"Expected state history with shape (B, {self.history_steps}, D), got {state.shape}"
)
processed_state = state.clone()
if self.relative:
mask_step = RelativeActionsProcessorStep(
enabled=True,
exclude_joints=self.exclude_joints,
action_names=self.state_names,
)
processed_state = to_relative_actions(
state,
state[:, -1],
mask_step._build_mask(state.shape[-1]),
pose_representation=self.pose_representation,
se3_pose_groups=self.se3_pose_groups,
)
flattened_dim = processed_state.shape[1] * processed_state.shape[2]
if flattened_dim > self.max_state_dim:
raise ValueError(
f"Flattened state history has {flattened_dim} dimensions, above max_state_dim={self.max_state_dim}"
)
new_transition = transition.copy()
new_observation = dict(observation)
new_observation[OBS_STATE] = processed_state.flatten(start_dim=1)
new_transition[TransitionKey.OBSERVATION] = new_observation
return new_transition
def get_config(self) -> dict[str, Any]:
return {
"history_steps": self.history_steps,
"max_state_dim": self.max_state_dim,
"relative": self.relative,
"exclude_joints": self.exclude_joints,
"state_names": self.state_names,
"pose_representation": self.pose_representation,
"se3_pose_groups": self.se3_pose_groups,
}
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
transformed = deepcopy(features)
for feature_group in transformed.values():
state_feature = feature_group.get(OBS_STATE)
if state_feature is not None:
state_dim = state_feature.shape[-1]
if self.relative:
source_dim = len(self.state_names) if self.state_names is not None else state_dim
model_dim = relative_action_output_dim(
source_dim,
self.pose_representation,
self.se3_pose_groups,
)
if state_dim == source_dim:
state_dim = model_dim
elif state_dim != model_dim:
raise ValueError(
f"Expected source/model state width {source_dim}/{model_dim}, got {state_dim}"
)
state_dim *= self.history_steps
feature_group[OBS_STATE] = PolicyFeature(type=state_feature.type, shape=(state_dim,))
return transformed
@ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step") @ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step")
@dataclass @dataclass
class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep): class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
@@ -133,13 +293,28 @@ def make_pi05_pre_post_processors(
enabled=config.use_relative_actions, enabled=config.use_relative_actions,
exclude_joints=getattr(config, "relative_exclude_joints", []), exclude_joints=getattr(config, "relative_exclude_joints", []),
action_names=getattr(config, "action_feature_names", None), action_names=getattr(config, "action_feature_names", None),
pose_representation=config.relative_pose_representation,
se3_pose_groups=config.relative_se3_pose_groups,
) )
# OpenPI order: raw → relative → normalize → model → unnormalize → absolute # OpenPI order: raw → relative → normalize → model → unnormalize → absolute
input_steps: list[ProcessorStep] = [ input_steps: list[ProcessorStep] = [
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
AddBatchDimensionProcessorStep(), AddBatchDimensionProcessorStep(),
Pi05StateFromActionProcessorStep(
enabled=config.state_from_action,
history_steps=config.proprioception_history_steps,
),
relative_step, relative_step,
Pi05FlattenStateHistoryProcessorStep(
history_steps=config.proprioception_history_steps,
max_state_dim=config.max_state_dim,
relative=config.use_relative_state_history,
exclude_joints=config.relative_state_exclude_joints,
state_names=config.action_feature_names,
pose_representation=config.relative_pose_representation,
se3_pose_groups=config.relative_se3_pose_groups,
),
# NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep # NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep
# because the tokenizer step expects normalized state in [-1, 1] range for discretization # because the tokenizer step expects normalized state in [-1, 1] range for discretization
NormalizerProcessorStep( NormalizerProcessorStep(
@@ -613,14 +613,15 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
device = tokens.device device = tokens.device
lm_head = self.paligemma_with_expert.paligemma.lm_head lm_head = self.paligemma_with_expert.paligemma.lm_head
# NOTE (bug 2 fix): do NOT append a second <bos> here. The language tokens # add bos token after tokens
# already begin with <bos> (standard PaliGemma prefix "[image] <bos> prompt \n"). bos_token = torch.full(
# Appending another <bos> right before decoding pushes the checkpoint into a (bsize, 1), self._paligemma_tokenizer.bos_token_id, dtype=torch.long, device=device
# bos->bos attractor and yields degenerate generation. Generate directly after )
# the prompt instead. tokens = torch.cat([tokens, bos_token], dim=1)
masks = torch.cat([masks, torch.ones((bsize, 1), dtype=torch.bool, device=device)], dim=1)
# 1. Initial Embedding (matches training prefix) # 1. Initial Embedding (matches training prefix)
# prefix_embs will include [Images, Language Prompt] # prefix_embs will include [Images, Language Prompt, BOS]
prefix_embs, prefix_pad_masks, prefix_att_masks, total_t_images, _ = self.embed_prefix_fast( prefix_embs, prefix_pad_masks, prefix_att_masks, total_t_images, _ = self.embed_prefix_fast(
images, img_masks, tokens, masks, fast_action_tokens=None, fast_action_masks=None images, img_masks, tokens, masks, fast_action_tokens=None, fast_action_masks=None
) )
@@ -708,13 +709,14 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
# --- 1. PREFILL PHASE --- # --- 1. PREFILL PHASE ---
# Process Images + Text Prompt + BOS token once to populate the KV cache. # Process Images + Text Prompt + BOS token once to populate the KV cache.
# NOTE (bug 2 fix): do NOT append a second <bos> here. The language tokens # Add BOS token to the prompt
# already begin with <bos> (standard PaliGemma prefix "[image] <bos> prompt \n"). bos_token = torch.full(
# A second <bos> right before decoding causes degenerate bos->bos generation. (bsize, 1), self._paligemma_tokenizer.bos_token_id, dtype=torch.long, device=device
tokens_in = tokens )
masks_in = masks tokens_in = torch.cat([tokens, bos_token], dim=1)
masks_in = torch.cat([masks, torch.ones((bsize, 1), dtype=torch.bool, device=device)], dim=1)
# Embed prefix [Images, Language] # Embed prefix [Images, Language, BOS]
# fast_action_tokens=None means we are just embedding the condition (images+text) # fast_action_tokens=None means we are just embedding the condition (images+text)
prefix_embs, prefix_pad_masks, prefix_att_masks, total_t_images, _ = self.embed_prefix_fast( prefix_embs, prefix_pad_masks, prefix_att_masks, total_t_images, _ = self.embed_prefix_fast(
images, img_masks, tokens_in, masks_in, fast_action_tokens=None, fast_action_masks=None images, img_masks, tokens_in, masks_in, fast_action_tokens=None, fast_action_masks=None
+4 -1
View File
@@ -121,7 +121,10 @@ class PiGemmaRMSNorm(nn.Module):
if cond.shape[-1] != self.cond_dim: if cond.shape[-1] != self.cond_dim:
raise ValueError(f"Expected cond dim {self.cond_dim}, got {cond.shape[-1]}") raise ValueError(f"Expected cond dim {self.cond_dim}, got {cond.shape[-1]}")
modulation = self.dense(cond) modulation = self.dense(cond)
if len(x.shape) == 3: # Per-sample conditioning (B, D) is broadcast across the sequence.
# Training-time RTC supplies per-action conditioning (B, T, D), which
# is already aligned with x and must keep its token dimension intact.
if len(x.shape) == 3 and modulation.dim() == 2:
modulation = modulation.unsqueeze(1) modulation = modulation.unsqueeze(1)
scale, shift, gate = modulation.chunk(3, dim=-1) scale, shift, gate = modulation.chunk(3, dim=-1)
normed = normed * (1 + scale.float()) + shift.float() normed = normed * (1 + scale.float()) + shift.float()
@@ -37,6 +37,10 @@ class RTCConfig:
# Infrastructure # Infrastructure
enabled: bool = True enabled: bool = True
# ``guided`` is the original inference-time Jacobian guidance. ``trained``
# hard-inpaints a prefix and requires a compatible training-time RTC checkpoint.
mode: str = "guided"
# Core RTC settings # Core RTC settings
# Todo change to exp # Todo change to exp
prefix_attention_schedule: RTCAttentionSchedule = RTCAttentionSchedule.LINEAR prefix_attention_schedule: RTCAttentionSchedule = RTCAttentionSchedule.LINEAR
@@ -49,6 +53,8 @@ class RTCConfig:
def __post_init__(self): def __post_init__(self):
"""Validate RTC configuration parameters.""" """Validate RTC configuration parameters."""
if self.mode not in {"guided", "trained"}:
raise ValueError(f"mode must be 'guided' or 'trained', got {self.mode!r}")
if self.max_guidance_weight <= 0: if self.max_guidance_weight <= 0:
raise ValueError(f"max_guidance_weight must be positive, got {self.max_guidance_weight}") raise ValueError(f"max_guidance_weight must be positive, got {self.max_guidance_weight}")
if self.debug_maxlen <= 0: if self.debug_maxlen <= 0:
+6 -1
View File
@@ -42,7 +42,12 @@ 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 PI05-compatible checkpoint trained with "
"rtc_training_max_delay > 0."
)
self.rtc_config = rtc_config self.rtc_config = rtc_config
self.tracker = None self.tracker = None
+7 -1
View File
@@ -49,7 +49,13 @@ def reanchor_relative_rtc_prefix(
action_cpu = prev_actions_absolute.detach().cpu() action_cpu = prev_actions_absolute.detach().cpu()
mask = relative_step._build_mask(action_cpu.shape[-1]) mask = relative_step._build_mask(action_cpu.shape[-1])
relative_actions = to_relative_actions(action_cpu, state, mask) relative_actions = to_relative_actions(
action_cpu,
state,
mask,
pose_representation=relative_step.pose_representation,
se3_pose_groups=relative_step.se3_pose_groups,
)
transition = create_transition(action=relative_actions) transition = create_transition(action=relative_actions)
if normalizer_step is not None: if normalizer_step is not None:
+16 -2
View File
@@ -89,8 +89,15 @@ from .policy_robot_bridge import (
from .relative_action_processor import ( from .relative_action_processor import (
AbsoluteActionsProcessorStep, AbsoluteActionsProcessorStep,
RelativeActionsProcessorStep, RelativeActionsProcessorStep,
relative_action_output_dim,
rotation_6d_to_rotvec,
rotvec_to_rotation_6d,
to_absolute_actions, to_absolute_actions,
to_absolute_se3_pose,
to_absolute_se3_pose_6d,
to_relative_actions, to_relative_actions,
to_relative_se3_pose,
to_relative_se3_pose_6d,
) )
from .rename_processor import RenameObservationsProcessorStep, rename_stats from .rename_processor import RenameObservationsProcessorStep, rename_stats
from .tokenizer_processor import ActionTokenizerProcessorStep, TokenizerProcessorStep from .tokenizer_processor import ActionTokenizerProcessorStep, TokenizerProcessorStep
@@ -135,6 +142,15 @@ __all__ = [
"make_default_robot_observation_processor", "make_default_robot_observation_processor",
"AbsoluteActionsProcessorStep", "AbsoluteActionsProcessorStep",
"RelativeActionsProcessorStep", "RelativeActionsProcessorStep",
"relative_action_output_dim",
"rotation_6d_to_rotvec",
"rotvec_to_rotation_6d",
"to_absolute_actions",
"to_absolute_se3_pose",
"to_absolute_se3_pose_6d",
"to_relative_actions",
"to_relative_se3_pose",
"to_relative_se3_pose_6d",
"MapDeltaActionToRobotActionStep", "MapDeltaActionToRobotActionStep",
"MapTensorToDeltaActionDictStep", "MapTensorToDeltaActionDictStep",
"NewLineTaskProcessorStep", "NewLineTaskProcessorStep",
@@ -168,8 +184,6 @@ __all__ = [
"transition_to_batch", "transition_to_batch",
"TransitionKey", "TransitionKey",
"TruncatedProcessorStep", "TruncatedProcessorStep",
"to_absolute_actions",
"to_relative_actions",
"UnnormalizerProcessorStep", "UnnormalizerProcessorStep",
"VanillaObservationProcessorStep", "VanillaObservationProcessorStep",
] ]
@@ -21,7 +21,7 @@ from torch import Tensor
from lerobot.configs import PipelineFeatureType, PolicyFeature from lerobot.configs import PipelineFeatureType, PolicyFeature
from lerobot.types import EnvTransition, TransitionKey from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import OBS_STATE from lerobot.utils.constants import ACTION, OBS_STATE
from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep
from .pipeline import ProcessorStep, ProcessorStepRegistry from .pipeline import ProcessorStep, ProcessorStepRegistry
@@ -34,57 +34,399 @@ __all__ = [
"AbsoluteActionsProcessorStep", "AbsoluteActionsProcessorStep",
"to_relative_actions", "to_relative_actions",
"to_absolute_actions", "to_absolute_actions",
"to_relative_se3_pose",
"to_absolute_se3_pose",
"to_relative_se3_pose_6d",
"to_absolute_se3_pose_6d",
"rotation_6d_to_rotvec",
"rotvec_to_rotation_6d",
"relative_action_output_dim",
] ]
def to_relative_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor: def _rotvec_to_quaternion(rotvec: Tensor) -> Tensor:
"""Convert absolute actions to relative: relative = action - state (for masked dims). angle = torch.linalg.vector_norm(rotvec, dim=-1, keepdim=True)
angle_sq = angle.square()
small_scale = 0.5 - angle_sq / 48.0 + angle_sq.square() / 3840.0
scale = torch.where(angle > 1e-6, torch.sin(angle / 2.0) / angle.clamp_min(1e-12), small_scale)
return torch.cat((torch.cos(angle / 2.0), rotvec * scale), dim=-1)
def _quaternion_to_rotvec(quaternion: Tensor) -> Tensor:
quaternion = quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
quaternion = quaternion * torch.where(quaternion[..., :1] < 0, -1.0, 1.0)
vector = quaternion[..., 1:]
sin_half_angle = torch.linalg.vector_norm(vector, dim=-1, keepdim=True)
angle = 2.0 * torch.atan2(sin_half_angle, quaternion[..., :1].clamp_min(0.0))
small_scale = 2.0 + sin_half_angle.square() / 3.0
scale = torch.where(
sin_half_angle > 1e-6,
angle / sin_half_angle.clamp_min(1e-12),
small_scale,
)
return vector * scale
def _quaternion_multiply(left: Tensor, right: Tensor) -> Tensor:
left_w, left_xyz = left[..., :1], left[..., 1:]
right_w, right_xyz = right[..., :1], right[..., 1:]
return torch.cat(
(
left_w * right_w - (left_xyz * right_xyz).sum(dim=-1, keepdim=True),
left_w * right_xyz + right_w * left_xyz + torch.linalg.cross(left_xyz, right_xyz, dim=-1),
),
dim=-1,
)
def _quaternion_conjugate(quaternion: Tensor) -> Tensor:
return torch.cat((quaternion[..., :1], -quaternion[..., 1:]), dim=-1)
def _quaternion_rotate(quaternion: Tensor, vector: Tensor) -> Tensor:
quaternion_xyz = quaternion[..., 1:]
uv = torch.linalg.cross(quaternion_xyz, vector, dim=-1)
uuv = torch.linalg.cross(quaternion_xyz, uv, dim=-1)
return vector + 2.0 * (quaternion[..., :1] * uv + uuv)
def _quaternion_to_matrix(quaternion: Tensor) -> Tensor:
quaternion = quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
w, x, y, z = quaternion.unbind(-1)
two_s = 2.0
return torch.stack(
(
1.0 - two_s * (y * y + z * z),
two_s * (x * y - z * w),
two_s * (x * z + y * w),
two_s * (x * y + z * w),
1.0 - two_s * (x * x + z * z),
two_s * (y * z - x * w),
two_s * (x * z - y * w),
two_s * (y * z + x * w),
1.0 - two_s * (x * x + y * y),
),
dim=-1,
).reshape(quaternion.shape[:-1] + (3, 3))
def _matrix_to_quaternion(matrix: Tensor) -> Tensor:
"""Convert proper rotation matrices to normalized ``[w, x, y, z]`` quaternions."""
if matrix.shape[-2:] != (3, 3):
raise ValueError(f"Rotation matrices must have shape (..., 3, 3), got {matrix.shape}")
m00 = matrix[..., 0, 0]
m01 = matrix[..., 0, 1]
m02 = matrix[..., 0, 2]
m10 = matrix[..., 1, 0]
m11 = matrix[..., 1, 1]
m12 = matrix[..., 1, 2]
m20 = matrix[..., 2, 0]
m21 = matrix[..., 2, 1]
m22 = matrix[..., 2, 2]
# Each row is a quaternion candidate scaled by the magnitude of its
# best-conditioned component. Selecting the largest component avoids the
# trace singularity at rotations close to pi.
q_abs = torch.sqrt(
torch.clamp(
torch.stack(
(
1.0 + m00 + m11 + m22,
1.0 + m00 - m11 - m22,
1.0 - m00 + m11 - m22,
1.0 - m00 - m11 + m22,
),
dim=-1,
),
min=0.0,
)
)
quat_by_rijk = torch.stack(
(
torch.stack((q_abs[..., 0].square(), m21 - m12, m02 - m20, m10 - m01), dim=-1),
torch.stack((m21 - m12, q_abs[..., 1].square(), m10 + m01, m02 + m20), dim=-1),
torch.stack((m02 - m20, m10 + m01, q_abs[..., 2].square(), m12 + m21), dim=-1),
torch.stack((m10 - m01, m02 + m20, m12 + m21, q_abs[..., 3].square()), dim=-1),
),
dim=-2,
)
candidates = quat_by_rijk / (2.0 * q_abs[..., :, None].clamp_min(0.1))
best = torch.nn.functional.one_hot(q_abs.argmax(dim=-1), num_classes=4).to(dtype=matrix.dtype)
quaternion = (candidates * best[..., :, None]).sum(dim=-2)
return quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
def rotvec_to_rotation_6d(rotvec: Tensor) -> Tensor:
"""Encode an axis-angle rotation as the first two rotation-matrix rows."""
matrix = _quaternion_to_matrix(_rotvec_to_quaternion(rotvec))
return matrix[..., :2, :].reshape(matrix.shape[:-2] + (6,))
def rotation_6d_to_rotvec(rotation_6d: Tensor) -> Tensor:
"""Decode two predicted 3-D vectors into an axis-angle rotation.
Gram-Schmidt orthonormalization follows the continuous 6-D rotation
representation. Degenerate predictions fail closed instead of producing an
invalid physical rotation.
"""
if rotation_6d.shape[-1] != 6:
raise ValueError(f"6-D rotations must have six values, got {rotation_6d.shape}")
first = rotation_6d[..., :3]
second = rotation_6d[..., 3:]
first_norm = torch.linalg.vector_norm(first, dim=-1, keepdim=True)
first_unit = first / first_norm.clamp_min(1e-12)
second_orthogonal = second - (first_unit * second).sum(dim=-1, keepdim=True) * first_unit
second_norm = torch.linalg.vector_norm(second_orthogonal, dim=-1, keepdim=True)
if bool(torch.any(first_norm <= 1e-8)) or bool(torch.any(second_norm <= 1e-8)):
raise ValueError("Cannot decode a degenerate 6-D rotation prediction")
second_unit = second_orthogonal / second_norm
third_unit = torch.linalg.cross(first_unit, second_unit, dim=-1)
matrix = torch.stack((first_unit, second_unit, third_unit), dim=-2)
return _quaternion_to_rotvec(_matrix_to_quaternion(matrix))
def to_relative_se3_pose(target_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Encode a pose as ``inv(T_reference) @ T_target``.
Poses use ``[x, y, z, rx, ry, rz]`` with an axis-angle rotation vector.
The relative translation is therefore expressed in the reference EE frame.
"""
if target_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
target_quaternion = _rotvec_to_quaternion(target_pose[..., 3:])
inverse_reference_quaternion = _quaternion_conjugate(reference_quaternion)
relative_translation = _quaternion_rotate(
inverse_reference_quaternion, target_pose[..., :3] - reference_pose[..., :3]
)
relative_quaternion = _quaternion_multiply(inverse_reference_quaternion, target_quaternion)
return torch.cat((relative_translation, _quaternion_to_rotvec(relative_quaternion)), dim=-1)
def to_absolute_se3_pose(relative_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Decode a pose with ``T_target = T_reference @ T_relative``."""
if relative_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
relative_quaternion = _rotvec_to_quaternion(relative_pose[..., 3:])
target_translation = reference_pose[..., :3] + _quaternion_rotate(
reference_quaternion, relative_pose[..., :3]
)
target_quaternion = _quaternion_multiply(reference_quaternion, relative_quaternion)
return torch.cat((target_translation, _quaternion_to_rotvec(target_quaternion)), dim=-1)
def to_relative_se3_pose_6d(target_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Encode ``inv(T_reference) @ T_target`` as xyz plus continuous 6-D rotation."""
relative_pose = to_relative_se3_pose(target_pose, reference_pose)
return torch.cat((relative_pose[..., :3], rotvec_to_rotation_6d(relative_pose[..., 3:])), dim=-1)
def to_absolute_se3_pose_6d(relative_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Decode xyz plus continuous 6-D rotation with ``T_target = T_reference @ T_relative``."""
if relative_pose.shape[-1] != 9:
raise ValueError("6-D encoded SE(3) poses must have nine values: xyz plus rotation-6D")
relative_rotvec_pose = torch.cat(
(relative_pose[..., :3], rotation_6d_to_rotvec(relative_pose[..., 3:])), dim=-1
)
return to_absolute_se3_pose(relative_rotvec_pose, reference_pose)
def _broadcast_reference(actions: Tensor, state: Tensor) -> Tensor:
if state.device != actions.device or state.dtype != actions.dtype:
state = state.to(device=actions.device, dtype=actions.dtype)
if actions.ndim == state.ndim + 1:
state = state.unsqueeze(-2)
return state
def _validate_se3_pose_groups(
pose_representation: str,
se3_pose_groups: Sequence[Sequence[int]] | None,
mask: Sequence[bool],
action_dim: int,
) -> list[list[int]]:
if pose_representation not in {"componentwise", "se3", "se3_6d"}:
raise ValueError(
f"Unsupported pose_representation={pose_representation!r}; expected "
"'componentwise', 'se3', or 'se3_6d'"
)
if pose_representation == "componentwise":
return []
if not se3_pose_groups:
raise ValueError(
f"pose_representation={pose_representation!r} requires at least one six-index se3_pose_group"
)
normalized_groups: list[list[int]] = []
used_indices: set[int] = set()
for raw_group in se3_pose_groups:
group = [int(index) for index in raw_group]
if len(group) != 6:
raise ValueError(f"Each SE(3) pose group must contain six indices, got {group}")
if len(set(group)) != 6 or any(index < 0 or index >= action_dim for index in group):
raise ValueError(f"Invalid SE(3) pose group for action_dim={action_dim}: {group}")
if pose_representation == "se3_6d" and group != list(range(group[0], group[0] + 6)):
raise ValueError("se3_6d pose groups must contain six contiguous ascending indices")
if any(index >= len(mask) for index in group):
raise ValueError(f"SE(3) pose group lies outside the relative mask: {group}")
if used_indices.intersection(group):
raise ValueError(f"SE(3) pose groups must not overlap: {group}")
group_mask = [bool(mask[index]) for index in group]
if any(group_mask) and not all(group_mask):
raise ValueError(f"An SE(3) pose group must be wholly relative or wholly absolute: {group}")
used_indices.update(group)
if all(group_mask):
normalized_groups.append(group)
return normalized_groups
def relative_action_output_dim(
source_dim: int,
pose_representation: str,
se3_pose_groups: Sequence[Sequence[int]] | None,
) -> int:
"""Return the model-space action width for a source action width."""
if pose_representation != "se3_6d":
return source_dim
groups = se3_pose_groups or []
return source_dim + 3 * len(groups)
def _expand_se3_6d_actions(
actions: Tensor,
state: Tensor,
groups: Sequence[Sequence[int]],
) -> Tensor:
group_by_start = {group[0]: list(group) for group in groups}
grouped_indices = {index for group in groups for index in group}
parts: list[Tensor] = []
for index in range(actions.shape[-1]):
group = group_by_start.get(index)
if group is not None:
parts.append(to_relative_se3_pose_6d(actions[..., group], state[..., group]))
elif index not in grouped_indices:
parts.append(actions[..., index : index + 1])
return torch.cat(parts, dim=-1)
def _collapse_se3_6d_actions(
actions: Tensor,
state: Tensor,
mask: Sequence[bool],
groups: Sequence[Sequence[int]],
) -> Tensor:
source_dim = len(mask)
expected_dim = relative_action_output_dim(source_dim, "se3_6d", groups)
if actions.shape[-1] != expected_dim:
raise ValueError(
f"Expected se3_6d action width {expected_dim} for source width {source_dim}, "
f"got {actions.shape[-1]}"
)
group_by_start = {group[0]: list(group) for group in groups}
grouped_indices = {index for group in groups for index in group}
parts: list[Tensor] = []
cursor = 0
for index in range(source_dim):
group = group_by_start.get(index)
if group is not None:
parts.append(to_absolute_se3_pose_6d(actions[..., cursor : cursor + 9], state[..., group]))
cursor += 9
elif index not in grouped_indices:
value = actions[..., cursor : cursor + 1]
if mask[index]:
value = value + state[..., index : index + 1]
parts.append(value)
cursor += 1
if cursor != actions.shape[-1]:
raise RuntimeError(f"Consumed {cursor} action values from width {actions.shape[-1]}")
return torch.cat(parts, dim=-1)
def to_relative_actions(
actions: Tensor,
state: Tensor,
mask: Sequence[bool],
*,
pose_representation: str = "componentwise",
se3_pose_groups: Sequence[Sequence[int]] | None = None,
) -> Tensor:
"""Convert absolute actions to a configured relative representation.
Component-wise mode computes ``action - state``. SE(3) modes compute
``inv(T_state) @ T_action`` for each configured pose group. ``se3_6d``
replaces each three-value relative rotation vector with its continuous
six-value encoding, increasing the output width by three per pose group.
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions: (B, T, action_dim) or (B, action_dim).
state: (B, state_dim). Broadcast across time dimension. state: (B, state_dim). Broadcast across time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask: Which dims to convert. Can be shorter than action_dim.
""" """
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, actions.shape[-1])
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
# Align state to the same device/dtype as actions. _last_state is cached before # Align state to the same device/dtype as actions. _last_state is cached before
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA. # DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
if state.device != actions.device or state.dtype != actions.dtype: state = _broadcast_reference(actions, state)
state = state.to(device=actions.device, dtype=actions.dtype) component_mask = mask_t.clone()
state_offset = state[..., :dims] * mask_t for group in groups:
if actions.ndim == 3: component_mask[group] = 0
state_offset = state_offset.unsqueeze(-2) state_offset = state[..., :dims] * component_mask
actions = actions.clone() actions = actions.clone()
actions[..., :dims] -= state_offset actions[..., :dims] -= state_offset
if pose_representation == "se3_6d":
return _expand_se3_6d_actions(actions, state, groups)
for group in groups:
actions[..., group] = to_relative_se3_pose(actions[..., group], state[..., group])
return actions return actions
def to_absolute_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor: def to_absolute_actions(
"""Convert relative actions back to absolute: absolute = relative + state (for masked dims). actions: Tensor,
state: Tensor,
mask: Sequence[bool],
*,
pose_representation: str = "componentwise",
se3_pose_groups: Sequence[Sequence[int]] | None = None,
) -> Tensor:
"""Convert relative actions back to absolute actions.
Component-wise mode computes ``relative + state``. SE(3) mode computes
``T_state @ T_relative`` for each configured pose group.
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions: (B, T, action_dim) or (B, action_dim).
state: (B, state_dim). Broadcast across time dimension. state: (B, state_dim). Broadcast across time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask: Which dims to convert. Can be shorter than action_dim.
""" """
source_dim = len(mask)
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, source_dim)
state = _broadcast_reference(actions, state)
if pose_representation == "se3_6d":
return _collapse_se3_6d_actions(actions, state, mask, groups)
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
# Align state to the same device/dtype as actions. _last_state is cached before # Align state to the same device/dtype as actions. _last_state is cached before
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA. # DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
if state.device != actions.device or state.dtype != actions.dtype: state = _broadcast_reference(actions, state)
state = state.to(device=actions.device, dtype=actions.dtype) component_mask = mask_t.clone()
state_offset = state[..., :dims] * mask_t for group in groups:
if actions.ndim == 3: component_mask[group] = 0
state_offset = state_offset.unsqueeze(-2) state_offset = state[..., :dims] * component_mask
actions = actions.clone() actions = actions.clone()
actions[..., :dims] += state_offset actions[..., :dims] += state_offset
for group in groups:
actions[..., group] = to_absolute_se3_pose(actions[..., group], state[..., group])
return actions return actions
@ProcessorStepRegistry.register("relative_actions_processor") @ProcessorStepRegistry.register("relative_actions_processor")
@dataclass @dataclass
class RelativeActionsProcessorStep(ProcessorStep): class RelativeActionsProcessorStep(ProcessorStep):
"""Converts absolute actions to relative actions (action -= state) for masked dimensions. """Converts absolute actions to the configured relative representation.
Mirrors OpenPI's DeltaActions transform. Applied during preprocessing so the model Mirrors OpenPI's DeltaActions transform. Applied during preprocessing so the model
trains on relative offsets instead of absolute positions. trains on relative offsets instead of absolute positions.
@@ -101,7 +443,10 @@ class RelativeActionsProcessorStep(ProcessorStep):
enabled: bool = False enabled: bool = False
exclude_joints: list[str] = field(default_factory=list) exclude_joints: list[str] = field(default_factory=list)
action_names: list[str] | None = None action_names: list[str] | None = None
pose_representation: str = "componentwise"
se3_pose_groups: list[list[int]] = field(default_factory=list)
_last_state: torch.Tensor | None = field(default=None, init=False, repr=False) _last_state: torch.Tensor | None = field(default=None, init=False, repr=False)
_last_mask: list[bool] | None = field(default=None, init=False, repr=False)
def _build_mask(self, action_dim: int) -> list[bool]: def _build_mask(self, action_dim: int) -> list[bool]:
if not self.exclude_joints or self.action_names is None: if not self.exclude_joints or self.action_names is None:
@@ -126,37 +471,78 @@ class RelativeActionsProcessorStep(ProcessorStep):
observation = transition.get(TransitionKey.OBSERVATION, {}) observation = transition.get(TransitionKey.OBSERVATION, {})
state = observation.get(OBS_STATE) if observation else None state = observation.get(OBS_STATE) if observation else None
# State history has shape (B, H, D). Relative actions are referenced to
# the newest proprioceptive state, not the whole history tensor.
reference_state = state[:, -1] if state is not None and state.ndim == 3 else state
# Always cache state for the paired AbsoluteActionsProcessorStep # Always cache state for the paired AbsoluteActionsProcessorStep
if state is not None: if reference_state is not None:
self._last_state = state self._last_state = reference_state
self._last_mask = self._build_mask(reference_state.shape[-1])
if not self.enabled: if not self.enabled:
return transition return transition
new_transition = transition.copy() new_transition = transition.copy()
action = new_transition.get(TransitionKey.ACTION) action = new_transition.get(TransitionKey.ACTION)
if action is None or state is None: if action is None or reference_state is None:
return new_transition return new_transition
mask = self._build_mask(action.shape[-1]) mask = self._last_mask or self._build_mask(action.shape[-1])
new_transition[TransitionKey.ACTION] = to_relative_actions(action, state, mask) new_transition[TransitionKey.ACTION] = to_relative_actions(
action,
reference_state,
mask,
pose_representation=self.pose_representation,
se3_pose_groups=self.se3_pose_groups,
)
return new_transition return new_transition
def get_cached_state(self) -> torch.Tensor | None: def get_cached_state(self) -> torch.Tensor | None:
"""Return the cached ``observation.state`` used as the reference point for relative/absolute action conversions.""" """Return the cached ``observation.state`` used as the reference point for relative/absolute action conversions."""
return self._last_state return self._last_state
def get_cached_mask(self) -> list[bool] | None:
"""Return the source-space mask cached with the latest state."""
return self._last_mask
def reset(self) -> None:
"""Drop the inference reference so it cannot leak between sessions."""
self._last_state = None
self._last_mask = None
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
return { return {
"enabled": self.enabled, "enabled": self.enabled,
"exclude_joints": self.exclude_joints, "exclude_joints": self.exclude_joints,
"action_names": self.action_names, "action_names": self.action_names,
"pose_representation": self.pose_representation,
"se3_pose_groups": self.se3_pose_groups,
} }
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
if not self.enabled or self.pose_representation != "se3_6d":
return features return features
transformed = {feature_type: dict(feature_group) for feature_type, feature_group in features.items()}
for feature_group in transformed.values():
action_feature = feature_group.get(ACTION)
if action_feature is None:
continue
source_dim = len(self.action_names) if self.action_names is not None else action_feature.shape[-1]
model_dim = relative_action_output_dim(source_dim, self.pose_representation, self.se3_pose_groups)
if action_feature.shape[-1] == source_dim:
feature_group[ACTION] = PolicyFeature(
type=action_feature.type,
shape=(model_dim,),
)
elif action_feature.shape[-1] != model_dim:
raise ValueError(
f"Expected source/model action width {source_dim}/{model_dim}, "
f"got {action_feature.shape[-1]}"
)
return transformed
@ProcessorStepRegistry.register("absolute_actions_processor") @ProcessorStepRegistry.register("absolute_actions_processor")
@@ -198,8 +584,16 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
if action is None: if action is None:
return new_transition return new_transition
mask = self.relative_step._build_mask(action.shape[-1]) mask = self.relative_step.get_cached_mask()
new_transition[TransitionKey.ACTION] = to_absolute_actions(action, cached_state, mask) if mask is None:
mask = self.relative_step._build_mask(cached_state.shape[-1])
new_transition[TransitionKey.ACTION] = to_absolute_actions(
action,
cached_state,
mask,
pose_representation=self.relative_step.pose_representation,
se3_pose_groups=self.relative_step.se3_pose_groups,
)
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
+3 -4
View File
@@ -476,12 +476,11 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
if tokens.dim() > 1: if tokens.dim() > 1:
tokens = tokens.flatten() tokens = tokens.flatten()
# NOTE (bug 2 fix): do NOT prepend a <bos> to the action target. The prompt bos_id = self._paligemma_tokenizer.bos_token_id
# already carries the leading <bos>; a second one before "Action:" mismatches # add bos
# the generation-time prefix (see sample_actions_fast*) and drives degenerate
# bos->bos decoding. Target is "Action: <fast tokens> |".
tokens = torch.cat( tokens = torch.cat(
[ [
torch.tensor([bos_id], device=action.device),
torch.tensor( torch.tensor(
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False), self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
device=action.device, device=action.device,
+68 -3
View File
@@ -46,6 +46,7 @@ from lerobot.processor import (
from lerobot.processor.relative_action_processor import RelativeActionsProcessorStep from lerobot.processor.relative_action_processor import RelativeActionsProcessorStep
from lerobot.robots import make_robot_from_config from lerobot.robots import make_robot_from_config
from lerobot.teleoperators import Teleoperator, make_teleoperator_from_config from lerobot.teleoperators import Teleoperator, make_teleoperator_from_config
from lerobot.utils.constants import OBS_STATE
from lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features from lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features
from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig
@@ -60,6 +61,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 not in {"pi05", "pi052"}:
raise ValueError(
"--inference.rtc.mode=trained currently requires a PI05-compatible 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]:
@@ -80,6 +110,26 @@ def _resolve_action_key_order(
return policy_action_names return policy_action_names
def _align_relative_state_feature_order(
hw_features: dict[str, dict], policy_action_names: list[str] | None
) -> dict[str, dict]:
"""Align policy-facing state with named relative-action dimensions."""
if not policy_action_names or OBS_STATE not in hw_features:
return hw_features
state_feature = hw_features[OBS_STATE]
state_names = state_feature.get("names")
if not state_names or len(state_names) != len(policy_action_names):
return hw_features
if set(state_names) != set(policy_action_names) or state_names == policy_action_names:
return hw_features
aligned = dict(hw_features)
aligned[OBS_STATE] = {**state_feature, "names": list(policy_action_names)}
logger.info("Aligned relative-action state order with checkpoint action names")
return aligned
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Sub-contexts # Sub-contexts
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -178,6 +228,9 @@ 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:
_validate_trained_rtc_rollout_config(policy_config, cfg.inference)
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
@@ -399,10 +452,22 @@ def build_rollout_context(
}, },
) )
if isinstance(cfg.inference, SyncInferenceConfig) and any( relative_action_step = next(
isinstance(step, RelativeActionsProcessorStep) and step.enabled (
step
for step in getattr(preprocessor, "steps", ()) for step in getattr(preprocessor, "steps", ())
): if isinstance(step, RelativeActionsProcessorStep) and step.enabled
),
None,
)
if relative_action_step is not None:
relative_action_names = relative_action_step.action_names or policy_action_names
hw_features = _align_relative_state_feature_order(
hw_features,
list(relative_action_names) if relative_action_names else None,
)
if isinstance(cfg.inference, SyncInferenceConfig) and relative_action_step is not None:
raise NotImplementedError( raise NotImplementedError(
"SyncInferenceEngine does not support policies with relative actions for now." "SyncInferenceEngine does not support policies with relative actions for now."
"Use --inference.type=rtc or remove relative action processor steps from the policy pipeline." "Use --inference.type=rtc or remove relative action processor steps from the policy pipeline."
+95 -2
View File
@@ -57,6 +57,18 @@ _RTC_MAX_CONSECUTIVE_ERRORS: int = 10
_RTC_JOIN_TIMEOUT_S: float = 3.0 _RTC_JOIN_TIMEOUT_S: float = 3.0
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."""
class _TrainedRTCPrefixUnavailableError(_FatalRTCInferenceError):
"""Raised when the queue cannot provide the prefix used for conditioning."""
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# RTC helpers # RTC helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -76,6 +88,50 @@ def _normalize_prev_actions_length(prev_actions: torch.Tensor, target_steps: int
return padded return padded
def _trained_rtc_chunk_can_merge(
*,
conditioned_delay: int,
measured_delay: int,
training_max_delay: int,
has_previous_actions: bool,
) -> bool:
"""Check that a trained RTC chunk covers the overlap observed during inference."""
if not has_previous_actions:
return True
if measured_delay > training_max_delay:
raise _TrainedRTCDelayExceededError(
f"Measured RTC inference delay ({measured_delay}) exceeds the checkpoint's "
f"rtc_training_max_delay ({training_max_delay})."
)
return measured_delay <= conditioned_delay
def _estimate_rtc_delay(
*,
latency: float,
time_per_step: float,
mode: str,
training_max_delay: int,
has_previous_actions: bool,
) -> int:
"""Estimate overlap, using the trained capacity to bootstrap the first transition."""
if latency:
return math.ceil(latency / time_per_step)
if mode == "trained" and has_previous_actions:
return training_max_delay
return 0
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
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -272,9 +328,23 @@ class RTCInferenceEngine(InferenceEngine):
current_time = time.perf_counter() current_time = time.perf_counter()
idx_before = queue.get_action_index() idx_before = queue.get_action_index()
prev_actions = queue.get_left_over() prev_actions = queue.get_left_over()
has_previous_actions = prev_actions is not None and prev_actions.numel() > 0
training_max_delay = int(getattr(self._policy.config, "rtc_training_max_delay", 0))
latency = latency_tracker.max() latency = latency_tracker.max()
delay = math.ceil(latency / time_per_chunk) if latency else 0 delay = _estimate_rtc_delay(
latency=latency,
time_per_step=time_per_chunk,
mode=self._rtc_config.mode,
training_max_delay=training_max_delay,
has_previous_actions=has_previous_actions,
)
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(
@@ -316,11 +386,32 @@ class RTCInferenceEngine(InferenceEngine):
inference_count += 1 inference_count += 1
consecutive_errors = 0 consecutive_errors = 0
is_warmup = self._use_torch_compile and inference_count <= warmup_required is_warmup = self._use_torch_compile and inference_count <= warmup_required
if is_warmup: is_initial_trained_chunk = (
self._rtc_config.mode == "trained" and not has_previous_actions
)
if is_warmup or is_initial_trained_chunk:
latency_tracker.reset() latency_tracker.reset()
else: else:
latency_tracker.add(new_latency) latency_tracker.add(new_latency)
if (
not is_warmup
and self._rtc_config.mode == "trained"
and not _trained_rtc_chunk_can_merge(
conditioned_delay=delay,
measured_delay=new_delay,
training_max_delay=training_max_delay,
has_previous_actions=has_previous_actions,
)
):
logger.warning(
"Discarding trained RTC chunk: measured delay %d exceeded "
"conditioned delay %d; retrying with updated latency",
new_delay,
delay,
)
continue
queue.merge(original, processed, new_delay, idx_before) queue.merge(original, processed, new_delay, idx_before)
if ( if (
@@ -333,6 +424,8 @@ 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 _FatalRTCInferenceError:
raise
except Exception as e: except Exception as e:
consecutive_errors += 1 consecutive_errors += 1
logger.error( logger.error(
@@ -325,6 +325,8 @@ class RecomputeStatsConfig(OperationConfig):
relative_exclude_joints: list[str] | None = None relative_exclude_joints: list[str] | None = None
chunk_size: int = 50 chunk_size: int = 50
num_workers: int = 0 num_workers: int = 0
relative_pose_representation: str = "componentwise"
relative_se3_pose_groups: list[list[int]] | None = None
overwrite: bool = False overwrite: bool = False
@@ -698,6 +700,8 @@ def handle_recompute_stats(cfg: EditDatasetConfig) -> None:
relative_exclude_joints=cfg.operation.relative_exclude_joints, relative_exclude_joints=cfg.operation.relative_exclude_joints,
chunk_size=cfg.operation.chunk_size, chunk_size=cfg.operation.chunk_size,
num_workers=cfg.operation.num_workers, num_workers=cfg.operation.num_workers,
relative_pose_representation=cfg.operation.relative_pose_representation,
relative_se3_pose_groups=cfg.operation.relative_se3_pose_groups,
) )
logging.info(f"Stats written to {dataset.root}") logging.info(f"Stats written to {dataset.root}")
+2
View File
@@ -343,6 +343,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
"enabled": True, "enabled": True,
"exclude_joints": getattr(active_cfg, "relative_exclude_joints", []), "exclude_joints": getattr(active_cfg, "relative_exclude_joints", []),
"action_names": getattr(active_cfg, "action_feature_names", None), "action_names": getattr(active_cfg, "action_feature_names", None),
"pose_representation": getattr(active_cfg, "relative_pose_representation", "componentwise"),
"se3_pose_groups": getattr(active_cfg, "relative_se3_pose_groups", []),
} }
postprocessor_overrides["absolute_actions_processor"] = {"enabled": True} postprocessor_overrides["absolute_actions_processor"] = {"enabled": True}
processor_kwargs["preprocessor_overrides"] = preprocessor_overrides processor_kwargs["preprocessor_overrides"] = preprocessor_overrides
@@ -0,0 +1,94 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
import pytest
import torch
pytest.importorskip("transformers")
from lerobot.policies.pi05.configuration_pi05 import PI05Config # noqa: E402
from lerobot.policies.pi05.modeling_pi05 import ( # noqa: E402
_build_flow_matching_inputs,
_prepare_trained_rtc_prefix,
_reduce_training_rtc_loss,
create_sinusoidal_pos_embedding,
)
from lerobot.policies.pi_gemma import PiGemmaRMSNorm # noqa: E402
def test_pi05_training_rtc_uses_clean_prefix_and_per_token_time():
actions = torch.tensor([[[1.0], [2.0], [3.0], [4.0]]])
noise = torch.tensor([[[10.0], [20.0], [30.0], [40.0]]])
time = torch.tensor([0.25])
prefix_mask = torch.tensor([[True, True, False, False]])
x_t, model_time = _build_flow_matching_inputs(actions, noise, time, prefix_mask)
assert model_time.tolist() == [[0.0, 0.0, 0.25, 0.25]]
assert torch.equal(x_t[:, :2], actions[:, :2])
assert torch.equal(x_t[:, 2:], 0.25 * noise[:, 2:] + 0.75 * actions[:, 2:])
def test_pi05_training_rtc_loss_excludes_clean_prefix():
losses = torch.tensor([[[100.0], [100.0], [2.0], [4.0]]])
prefix_mask = torch.tensor([[True, True, False, False]])
loss = _reduce_training_rtc_loss(losses, prefix_mask, reduction="mean")
assert loss.item() == pytest.approx(3.0)
def test_pi05_training_rtc_adaptive_norm_accepts_per_action_time_conditioning():
norm = PiGemmaRMSNorm(dim=4, cond_dim=3)
hidden = torch.randn(2, 5, 4)
per_action_condition = torch.randn(2, 5, 3)
output, gate = norm(hidden, per_action_condition)
assert output.shape == hidden.shape
assert gate.shape == hidden.shape
def test_pi05_training_rtc_embeds_per_action_timesteps():
per_action_time = torch.tensor([[0.0, 0.0, 0.25, 0.25]])
embedding = create_sinusoidal_pos_embedding(
per_action_time,
dimension=8,
min_period=4e-3,
max_period=4.0,
device=per_action_time.device,
)
assert embedding.shape == (1, 4, 8)
torch.testing.assert_close(embedding[:, 0], embedding[:, 1])
torch.testing.assert_close(embedding[:, 2], embedding[:, 3])
def test_pi05_trained_rtc_prefix_is_padded_to_model_width():
latent = torch.zeros(1, 5, 32)
previous = torch.arange(30, dtype=torch.float32).reshape(1, 3, 10)
prefix, mask = _prepare_trained_rtc_prefix(
latent,
previous,
inference_delay=2,
training_max_delay=4,
)
assert prefix.shape == latent.shape
assert mask.shape == latent.shape
torch.testing.assert_close(prefix[:, :2, :10], previous[:, :2])
assert torch.count_nonzero(prefix[:, :2, 10:]) == 0
assert mask[:, :2].all()
assert not mask[:, 2:].any()
@pytest.mark.parametrize("max_delay", [-1, 5])
def test_pi05_config_rejects_invalid_training_rtc_delay(max_delay):
with pytest.raises(ValueError, match="rtc_training_max_delay"):
PI05Config(chunk_size=5, n_action_steps=5, rtc_training_max_delay=max_delay)
@@ -16,6 +16,8 @@
"""Tests for RTC configuration module.""" """Tests for RTC configuration module."""
import pytest
from lerobot.configs.types import RTCAttentionSchedule from lerobot.configs.types import RTCAttentionSchedule
from lerobot.policies.rtc.configuration_rtc import RTCConfig from lerobot.policies.rtc.configuration_rtc import RTCConfig
@@ -27,6 +29,7 @@ def test_rtc_config_default_initialization():
config = RTCConfig() config = RTCConfig()
assert config.enabled is True assert config.enabled is True
assert config.mode == "guided"
assert config.prefix_attention_schedule == RTCAttentionSchedule.LINEAR assert config.prefix_attention_schedule == RTCAttentionSchedule.LINEAR
assert config.max_guidance_weight == 10.0 assert config.max_guidance_weight == 10.0
assert config.execution_horizon == 10 assert config.execution_horizon == 10
@@ -34,10 +37,16 @@ def test_rtc_config_default_initialization():
assert config.debug_maxlen == 100 assert config.debug_maxlen == 100
def test_rtc_config_rejects_unknown_mode():
with pytest.raises(ValueError, match="mode must be"):
RTCConfig(mode="unknown")
def test_rtc_config_custom_initialization(): def test_rtc_config_custom_initialization():
"""Test RTCConfig initializes with custom values.""" """Test RTCConfig initializes with custom values."""
config = RTCConfig( config = RTCConfig(
enabled=True, enabled=True,
mode="trained",
prefix_attention_schedule=RTCAttentionSchedule.EXP, prefix_attention_schedule=RTCAttentionSchedule.EXP,
max_guidance_weight=5.0, max_guidance_weight=5.0,
execution_horizon=20, execution_horizon=20,
@@ -46,6 +55,7 @@ def test_rtc_config_custom_initialization():
) )
assert config.enabled is True assert config.enabled is True
assert config.mode == "trained"
assert config.prefix_attention_schedule == RTCAttentionSchedule.EXP assert config.prefix_attention_schedule == RTCAttentionSchedule.EXP
assert config.max_guidance_weight == 5.0 assert config.max_guidance_weight == 5.0
assert config.execution_horizon == 20 assert config.execution_horizon == 20
+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 PI05-compatible 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 ======================
@@ -505,6 +505,44 @@ class TestRTCReanchoringWithStateNormalizer:
assert not torch.allclose(cached, post_normalize_state, atol=1e-3) assert not torch.allclose(cached, post_normalize_state, atol=1e-3)
def test_reanchor_se3_6d_prefix_uses_current_camera_frame_and_model_width():
"""A leftover absolute EE chunk is recomposed relative to the latest camera-frame EE pose."""
names = ["pos_x", "pos_y", "pos_z", "rot_x", "rot_y", "rot_z", "gripper"]
relative_step = RelativeActionsProcessorStep(
enabled=True,
exclude_joints=["gripper"],
action_names=names,
pose_representation="se3_6d",
se3_pose_groups=[list(range(6))],
)
current_state = torch.tensor([[0.20, -0.10, 0.40, 0.31, -0.22, 0.17, 0.03]])
previous_absolute = torch.tensor(
[
[0.28, -0.03, 0.46, -0.18, 0.27, 0.41, 0.02],
[0.31, 0.02, 0.50, -0.11, 0.35, 0.52, 0.01],
]
)
result = reanchor_relative_rtc_prefix(
prev_actions_absolute=previous_absolute,
current_state=current_state,
relative_step=relative_step,
normalizer_step=None,
policy_device="cpu",
)
expected = to_relative_actions(
previous_absolute,
current_state,
relative_step._build_mask(previous_absolute.shape[-1]),
pose_representation="se3_6d",
se3_pose_groups=[list(range(6))],
)
assert result.shape == (2, 10)
torch.testing.assert_close(result, expected, atol=1e-6, rtol=1e-6)
torch.testing.assert_close(result[:, -1], previous_absolute[:, -1])
def _detect_relative_actions(preprocessor) -> bool: def _detect_relative_actions(preprocessor) -> bool:
"""Mirror of the helper in lerobot-rollout for testing without importing it.""" """Mirror of the helper in lerobot-rollout for testing without importing it."""
return any(isinstance(step, RelativeActionsProcessorStep) and step.enabled for step in preprocessor.steps) return any(isinstance(step, RelativeActionsProcessorStep) and step.enabled for step in preprocessor.steps)
@@ -0,0 +1,321 @@
#!/usr/bin/env python
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from math import pi
import numpy as np
import pytest
import torch
pytest.importorskip("transformers")
from lerobot.configs import FeatureType, PolicyFeature # noqa: E402
from lerobot.datasets.compute_stats import ( # noqa: E402
compute_relative_action_stats,
compute_state_history_stats,
)
from lerobot.policies.pi05.configuration_pi05 import PI05Config # noqa: E402
from lerobot.policies.pi05.processor_pi05 import ( # noqa: E402
Pi05FlattenStateHistoryProcessorStep,
Pi05StateFromActionProcessorStep,
)
from lerobot.processor.relative_action_processor import ( # noqa: E402
AbsoluteActionsProcessorStep,
RelativeActionsProcessorStep,
)
from lerobot.types import TransitionKey # noqa: E402
from lerobot.utils.constants import ACTION, OBS_STATE # noqa: E402
def _transition(action: torch.Tensor | None, state: torch.Tensor | None = None) -> dict:
observation = {} if state is None else {OBS_STATE: state}
return {
TransitionKey.OBSERVATION: observation,
TransitionKey.ACTION: action,
TransitionKey.REWARD: None,
TransitionKey.DONE: None,
TransitionKey.TRUNCATED: None,
TransitionKey.COMPLEMENTARY_DATA: {},
}
def test_pi05_config_requests_action_history_prefix():
config = PI05Config(
device="cpu",
chunk_size=4,
n_action_steps=4,
state_from_action=True,
proprioception_history_steps=2,
)
assert config.action_delta_indices == [-1, 0, 1, 2, 3]
def test_pi05_config_accepts_se3_6d_action_and_state_with_two_step_history():
names = ["x", "y", "z", "rx", "ry", "rz", "gripper_width"]
config = PI05Config(
device="cpu",
use_relative_actions=True,
state_from_action=True,
proprioception_history_steps=2,
use_relative_state_history=True,
relative_pose_representation="se3_6d",
action_feature_names=names,
output_features={
ACTION: PolicyFeature(type=FeatureType.ACTION, shape=(7,)),
},
)
config.validate_features()
assert config.output_features[ACTION].shape == (10,)
assert config.input_features[OBS_STATE].shape == (10,)
def test_state_from_action_extracts_history_and_preserves_target_horizon():
action = torch.arange(2 * 5 * 3, dtype=torch.float32).reshape(2, 5, 3)
step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2)
result = step(_transition(action))
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], action[:, :2])
torch.testing.assert_close(result[TransitionKey.ACTION], action[:, 1:])
def test_relative_actions_use_newest_state_in_history_and_roundtrip():
state_history = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]])
absolute = torch.tensor([[[3.0, 30.0], [4.0, 40.0]]])
relative_step = RelativeActionsProcessorStep(enabled=True)
absolute_step = AbsoluteActionsProcessorStep(enabled=True, relative_step=relative_step)
relative = relative_step(_transition(absolute, state_history))
expected = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]])
torch.testing.assert_close(relative[TransitionKey.ACTION], expected)
recovered = absolute_step(_transition(relative[TransitionKey.ACTION]))
torch.testing.assert_close(recovered[TransitionKey.ACTION], absolute)
def test_relative_action_reference_is_reset_between_inference_sessions():
step = RelativeActionsProcessorStep(enabled=True)
step(_transition(None, torch.tensor([[1.0, 2.0]])))
step.reset()
assert step.get_cached_state() is None
assert step.get_cached_mask() is None
def test_flatten_state_history_preserves_chronological_order():
state_history = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=4)
result = step(_transition(torch.zeros(1, 2, 2), state_history))
torch.testing.assert_close(
result[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[1.0, 2.0, 3.0, 4.0]])
)
def test_state_history_can_be_relative_with_absolute_gripper():
state_history = torch.tensor([[[1.0, 10.0, 0.2], [3.0, 20.0, 0.4]]])
step = Pi05FlattenStateHistoryProcessorStep(
history_steps=2,
max_state_dim=6,
relative=True,
exclude_joints=["gripper"],
state_names=["x", "y", "gripper_width"],
)
result = step(_transition(torch.zeros(1, 2, 3), state_history))
torch.testing.assert_close(
result[TransitionKey.OBSERVATION][OBS_STATE],
torch.tensor([[-2.0, -10.0, 0.2, 0.0, 0.0, 0.4]]),
)
def test_state_history_can_use_se3_composition_with_absolute_gripper():
state_history = torch.tensor(
[[[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.2], [0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.4]]]
)
step = Pi05FlattenStateHistoryProcessorStep(
history_steps=2,
max_state_dim=14,
relative=True,
exclude_joints=["gripper"],
state_names=["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
pose_representation="se3",
se3_pose_groups=[list(range(6))],
)
result = step(_transition(torch.zeros(1, 2, 7), state_history))
expected = torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.2, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.4]])
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], expected, atol=1e-6, rtol=1e-6)
def test_state_history_can_use_se3_6d_rotation_with_absolute_gripper():
state_history = torch.tensor(
[[[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.2], [0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.4]]]
)
step = Pi05FlattenStateHistoryProcessorStep(
history_steps=2,
max_state_dim=20,
relative=True,
exclude_joints=["gripper"],
state_names=["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
pose_representation="se3_6d",
se3_pose_groups=[list(range(6))],
)
result = step(_transition(torch.zeros(1, 2, 7), state_history))
identity_6d = [1.0, 0.0, 0.0, 0.0, 1.0, 0.0]
expected = torch.tensor([[1.0, 0.0, 0.0, *identity_6d, 0.2, 0.0, 0.0, 0.0, *identity_6d, 0.4]])
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], expected, atol=1e-6, rtol=1e-6)
def test_inference_state_history_is_rolled_and_reset():
step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2)
first = step(_transition(None, torch.tensor([[1.0, 2.0]])))
second = step(_transition(None, torch.tensor([[3.0, 4.0]])))
torch.testing.assert_close(
first[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [1.0, 2.0]]])
)
torch.testing.assert_close(
second[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
)
step.reset()
reset = step(_transition(None, torch.tensor([[5.0, 6.0]])))
torch.testing.assert_close(
reset[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[5.0, 6.0], [5.0, 6.0]]])
)
def test_flatten_state_history_checks_max_state_dim():
step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=3)
with pytest.raises(ValueError, match="above max_state_dim"):
step(_transition(torch.zeros(1, 2, 2), torch.zeros(1, 2, 2)))
def test_relative_stats_can_use_absolute_action_as_state():
actions = np.asarray([[0.0, 0.0], [1.0, 2.0], [2.0, 4.0], [3.0, 6.0]], dtype=np.float32)
dataset = {"action": actions, "episode_index": np.zeros(4, dtype=np.int64)}
features = {"action": {"shape": [2], "names": ["x", "y"]}}
stats = compute_relative_action_stats(
dataset,
features,
chunk_size=2,
state_from_action=True,
)
np.testing.assert_allclose(stats["mean"], [0.5, 1.0])
def test_relative_state_history_stats_match_processor_representation():
actions = np.asarray(
[[0.0, 0.1], [1.0, 0.2], [3.0, 0.3]],
dtype=np.float32,
)
dataset = {"action": actions, "episode_index": np.zeros(3, dtype=np.int64)}
features = {"action": {"shape": [2], "names": ["x", "gripper_width"]}}
stats = compute_state_history_stats(
dataset,
features,
history_steps=2,
exclude_joints=["gripper"],
relative=True,
)
expected = np.asarray([[0.0, 0.1, 0.0, 0.1], [-1.0, 0.1, 0.0, 0.2], [-2.0, 0.2, 0.0, 0.3]])
np.testing.assert_allclose(stats["mean"], expected.mean(axis=0))
def test_se3_relative_action_stats_use_reference_frame():
actions = np.asarray(
[
[0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.2],
[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.3],
],
dtype=np.float32,
)
dataset = {"action": actions, "episode_index": np.zeros(2, dtype=np.int64)}
features = {
"action": {
"shape": [7],
"names": ["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
}
}
stats = compute_relative_action_stats(
dataset,
features,
chunk_size=2,
exclude_joints=["gripper"],
state_from_action=True,
pose_representation="se3",
se3_pose_groups=[list(range(6))],
)
np.testing.assert_allclose(stats["mean"][:3], [0.5, 0.0, 0.0], atol=1e-6)
np.testing.assert_allclose(stats["mean"][6], 0.25, atol=1e-6)
def test_se3_6d_stats_expand_action_and_state_history():
actions = np.asarray(
[
[0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.2],
[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.3],
],
dtype=np.float32,
)
dataset = {"action": actions, "episode_index": np.zeros(2, dtype=np.int64)}
features = {
"action": {
"shape": [7],
"names": ["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
}
}
action_stats = compute_relative_action_stats(
dataset,
features,
chunk_size=2,
exclude_joints=["gripper"],
state_from_action=True,
pose_representation="se3_6d",
se3_pose_groups=[list(range(6))],
)
state_stats = compute_state_history_stats(
dataset,
features,
history_steps=2,
exclude_joints=["gripper"],
relative=True,
pose_representation="se3_6d",
se3_pose_groups=[list(range(6))],
)
assert action_stats["mean"].shape == (10,)
assert state_stats["mean"].shape == (20,)
np.testing.assert_allclose(action_stats["mean"][:3], [0.5, 0.0, 0.0], atol=1e-6)
np.testing.assert_allclose(action_stats["mean"][9], 0.25, atol=1e-6)
@@ -0,0 +1,144 @@
import math
import pytest
import torch
from lerobot.processor.relative_action_processor import (
rotation_6d_to_rotvec,
rotvec_to_rotation_6d,
to_absolute_actions,
to_absolute_se3_pose,
to_absolute_se3_pose_6d,
to_relative_actions,
to_relative_se3_pose,
to_relative_se3_pose_6d,
)
POSE_GROUP = [list(range(6))]
def test_se3_translation_is_expressed_in_reference_frame():
reference = torch.tensor([[1.0, 2.0, 3.0, 0.0, 0.0, math.pi / 2]])
target = torch.tensor([[1.0, 3.0, 3.0, 0.0, 0.0, math.pi / 2]])
relative = to_relative_se3_pose(target, reference)
torch.testing.assert_close(relative, torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0]]), atol=1e-6, rtol=1e-6)
def test_se3_pose_roundtrip_for_batched_chunks():
torch.manual_seed(0)
reference = torch.randn(4, 6)
reference[:, 3:] *= 0.8
target = torch.randn(4, 11, 6)
target[..., 3:] *= 0.8
relative = to_relative_se3_pose(target, reference.unsqueeze(1))
recovered = to_absolute_se3_pose(relative, reference.unsqueeze(1))
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_mixed_se3_pose_and_absolute_gripper_roundtrip():
reference = torch.tensor([[0.2, -0.1, 0.4, 0.1, 0.2, -0.3, 0.06]])
target = torch.tensor([[[0.3, 0.2, 0.5, -0.2, 0.1, 0.4, 0.03], [0.1, -0.3, 0.2, 0.5, -0.1, 0.2, 0.05]]])
mask = [True, True, True, True, True, True, False]
relative = to_relative_actions(
target,
reference,
mask,
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)
recovered = to_absolute_actions(
relative,
reference,
mask,
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)
torch.testing.assert_close(relative[..., 6], target[..., 6])
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_se3_pose_group_cannot_be_partially_relative():
with pytest.raises(ValueError, match="wholly relative or wholly absolute"):
to_relative_actions(
torch.zeros(1, 7),
torch.zeros(1, 7),
[True, True, True, False, False, False, False],
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)
@pytest.mark.parametrize(
"rotvec",
[
[0.0, 0.0, 0.0],
[0.2, -0.5, 0.8],
[math.pi - 1e-4, 0.0, 0.0],
],
)
def test_rotation_6d_roundtrip(rotvec):
source = torch.tensor([rotvec], dtype=torch.float64)
recovered = rotation_6d_to_rotvec(rotvec_to_rotation_6d(source))
torch.testing.assert_close(recovered, source, atol=2e-6, rtol=2e-6)
def test_rotation_6d_uses_umi_first_two_rows():
source = torch.tensor([[0.0, 0.0, math.pi / 2]], dtype=torch.float64)
encoded = rotvec_to_rotation_6d(source)
expected = torch.tensor([[0.0, -1.0, 0.0, 1.0, 0.0, 0.0]], dtype=torch.float64)
torch.testing.assert_close(encoded, expected, atol=1e-7, rtol=1e-7)
torch.testing.assert_close(rotation_6d_to_rotvec(expected), source, atol=1e-7, rtol=1e-7)
def test_se3_6d_pose_roundtrip_for_batched_chunks():
torch.manual_seed(1)
reference = torch.randn(4, 6)
reference[:, 3:] *= 0.8
target = torch.randn(4, 11, 6)
target[..., 3:] *= 0.8
relative = to_relative_se3_pose_6d(target, reference.unsqueeze(1))
recovered = to_absolute_se3_pose_6d(relative, reference.unsqueeze(1))
assert relative.shape == (4, 11, 9)
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_mixed_se3_6d_pose_and_absolute_gripper_roundtrip():
reference = torch.tensor([[0.2, -0.1, 0.4, 0.1, 0.2, -0.3, 0.06]])
target = torch.tensor([[[0.3, 0.2, 0.5, -0.2, 0.1, 0.4, 0.03], [0.1, -0.3, 0.2, 0.5, -0.1, 0.2, 0.05]]])
mask = [True, True, True, True, True, True, False]
relative = to_relative_actions(
target,
reference,
mask,
pose_representation="se3_6d",
se3_pose_groups=POSE_GROUP,
)
recovered = to_absolute_actions(
relative,
reference,
mask,
pose_representation="se3_6d",
se3_pose_groups=POSE_GROUP,
)
assert relative.shape == (1, 2, 10)
torch.testing.assert_close(relative[..., 9], target[..., 6])
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_rotation_6d_rejects_degenerate_prediction():
with pytest.raises(ValueError, match="degenerate"):
rotation_6d_to_rotvec(torch.zeros(1, 6))
+126
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
@@ -98,6 +99,131 @@ def test_inference_config_types():
assert rtc.rtc is not None assert rtc.rtc is not None
def test_trained_rtc_retries_chunk_when_measured_delay_exceeds_conditioning():
from lerobot.rollout.inference.rtc import _trained_rtc_chunk_can_merge
assert not _trained_rtc_chunk_can_merge(
conditioned_delay=2,
measured_delay=3,
training_max_delay=4,
has_previous_actions=True,
)
assert _trained_rtc_chunk_can_merge(
conditioned_delay=2,
measured_delay=5,
training_max_delay=4,
has_previous_actions=False,
)
def test_trained_rtc_bootstraps_first_overlap_with_checkpoint_capacity():
from lerobot.rollout.inference.rtc import _estimate_rtc_delay
assert (
_estimate_rtc_delay(
latency=0,
time_per_step=1 / 30,
mode="trained",
training_max_delay=10,
has_previous_actions=False,
)
== 0
)
assert (
_estimate_rtc_delay(
latency=0,
time_per_step=1 / 30,
mode="trained",
training_max_delay=10,
has_previous_actions=True,
)
== 10
)
def test_trained_rtc_rejects_measured_delay_above_checkpoint_support():
from lerobot.rollout.inference.rtc import (
_trained_rtc_chunk_can_merge,
_TrainedRTCDelayExceededError,
)
with pytest.raises(_TrainedRTCDelayExceededError, match="rtc_training_max_delay"):
_trained_rtc_chunk_can_merge(
conditioned_delay=3,
measured_delay=5,
training_max_delay=4,
has_previous_actions=True,
)
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="pi05", 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_relative_state_order_follows_checkpoint_action_names():
from lerobot.rollout.context import _align_relative_state_feature_order
from lerobot.utils.constants import OBS_STATE
from lerobot.utils.feature_utils import build_dataset_frame
hw_features = {
OBS_STATE: {
"dtype": "float32",
"shape": (4,),
"names": ["left_joint.pos", "left_gripper.pos", "right_joint.pos", "right_gripper.pos"],
}
}
checkpoint_order = [
"right_joint.pos",
"right_gripper.pos",
"left_joint.pos",
"left_gripper.pos",
]
aligned = _align_relative_state_feature_order(hw_features, checkpoint_order)
frame = build_dataset_frame(
aligned,
{
"left_joint.pos": 1.0,
"left_gripper.pos": 2.0,
"right_joint.pos": 3.0,
"right_gripper.pos": 4.0,
},
prefix="observation",
)
assert aligned[OBS_STATE]["names"] == checkpoint_order
assert frame[OBS_STATE].tolist() == [3.0, 4.0, 1.0, 2.0]
assert hw_features[OBS_STATE]["names"][0] == "left_joint.pos"
def test_sentry_config_defaults(): def test_sentry_config_defaults():
from lerobot.rollout import SentryStrategyConfig from lerobot.rollout import SentryStrategyConfig