mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4045589246 | |||
| 80761e8a7a | |||
| b0cceb2a5f | |||
| 8e12a5351a | |||
| 7e1077f19a |
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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]:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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."
|
||||||
|
|||||||
@@ -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}")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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))
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user