feat(pi052): add training-time RTC

This commit is contained in:
Pepijn
2026-07-16 10:38:28 +02:00
parent 5b8e6ffe8e
commit 18e02ded4f
10 changed files with 396 additions and 25 deletions
+2 -1
View File
@@ -275,8 +275,9 @@ 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 Pi052 checkpoints |
| `--inference.rtc.max_guidance_weight` | Consistency enforcement strength (default: varies by policy) | | `--inference.rtc.max_guidance_weight` | Consistency enforcement strength (default: varies by policy) |
| `--inference.rtc.prefix_attention_schedule` | Blend schedule: `LINEAR`, `EXP`, `ONES`, `ZEROS` | | `--inference.rtc.prefix_attention_schedule` | Blend schedule: `LINEAR`, `EXP`, `ONES`, `ZEROS` |
| `--inference.queue_threshold` | Max queue size before backpressure (default: 30) | | `--inference.queue_threshold` | Max queue size before backpressure (default: 30) |
+59
View File
@@ -116,12 +116,71 @@ the expected prompt, text target, and action endpoints before scaling up.
| `policy.fast_action_loss_weight` | `1.0` | FAST cross-entropy weight | | `policy.fast_action_loss_weight` | `1.0` | FAST cross-entropy weight |
| `policy.knowledge_insulation` | `true` | Blocks action-loss gradients through the VLM K/V path | | `policy.knowledge_insulation` | `true` | Blocks action-loss gradients through the VLM K/V path |
| `policy.flow_num_repeats` | `5` | Reuses one VLM prefix for independent denoising targets | | `policy.flow_num_repeats` | `5` | Reuses one VLM prefix for independent denoising targets |
| `policy.rtc_training_max_delay` | `0` | Maximum clean-prefix delay; `0` disables training-time RTC |
| `policy.lm_head_lr_scale` | `1.0` | Scales language-head learning rate; `1.0` uses the base rate | | `policy.lm_head_lr_scale` | `1.0` | Scales language-head learning rate; `1.0` uses the base rate |
The loss weights are starting points, not dataset-independent constants. Track The loss weights are starting points, not dataset-independent constants. Track
flow loss and text/FAST losses separately, and inspect generated subtasks rather flow loss and text/FAST losses separately, and inspect generated subtasks rather
than selecting a checkpoint from total loss alone. than selecting a checkpoint from total loss alone.
### Training-time RTC
Pi052 optionally supports training-time action conditioning from
[Training-Time Action Conditioning for Efficient Real-Time Chunking](https://arxiv.org/abs/2512.05964).
It simulates inference latency by sampling a clean action prefix for every flow
draw, passing a per-action flow timestep to the action expert, and computing the
flow loss only on the remaining postfix. The default value of `0` leaves the
standard Pi052 objective unchanged.
```bash
lerobot-train \
--dataset.repo_id=${HF_USER}/my_language_annotated_dataset \
--policy.type=pi052 \
--policy.pretrained_path=lerobot/pi05_base \
--policy.recipe_path=recipes/subtask_mem.yaml \
--policy.rtc_training_max_delay=10 \
--policy.dtype=bfloat16 \
--policy.device=cuda \
--batch_size=8 \
--steps=30000 \
--output_dir=outputs/pi052_rtc \
--job_name=pi052_rtc
```
`rtc_training_max_delay` is measured in controller steps and must be smaller
than `chunk_size`. Choose it to cover the largest inference latency expected at
deployment: at 50 Hz, for example, 10 steps correspond to 200 ms. A delay of
zero is included in the uniform sampling distribution, so the checkpoint also
continues to receive ordinary flow-matching examples. Set rollout's
`inference.rtc.execution_horizon` to at least this maximum so the previous
chunk cache retains enough actions to construct every supported prefix.
Run the resulting checkpoint with the asynchronous `lerobot-rollout` backend
and select the trained-prefix path explicitly:
```bash
lerobot-rollout \
--strategy.type=base \
--policy.path=outputs/pi052_rtc/checkpoints/last/pretrained_model \
--inference.type=rtc \
--inference.rtc.mode=trained \
--inference.rtc.execution_horizon=10 \
--robot.type=so100_follower \
--robot.port=/dev/ttyACM0 \
--task="pick up the cube" \
--fps=50 \
--device=cuda
```
The rollout engine measures latency continuously, carries the still-unexecuted
actions from the previous chunk into the next prediction, and discards the
prefix that elapsed during inference. If the measured delay exceeds the
checkpoint's `rtc_training_max_delay`, rollout stops with an explicit error
instead of silently extrapolating beyond the training distribution. Use
`--inference.rtc.mode=guided` for the original Jacobian-guided RTC path; it does
not require a training-time RTC checkpoint but adds backward-pass work during
denoising.
### Dataset-specific FAST tokenizer ### Dataset-specific FAST tokenizer
The universal FAST tokenizer works out of the box. For a large or The universal FAST tokenizer works out of the box. For a large or
+27 -1
View File
@@ -1,6 +1,6 @@
# Real-Time Chunking (RTC) # Real-Time Chunking (RTC)
Real-Time Chunking (RTC) is an inference-time method that allows large, flow-matching based robotic policies, such as [Pi0](./pi0), [Pi0.5](./pi05), and [SmolVLA](./smolvla), to produce smooth, continuous, and reactive motion despite having high inference latency. Real-Time Chunking (RTC) allows large, flow-matching based robotic policies, such as [Pi0](./pi0), [Pi0.5](./pi05), and [SmolVLA](./smolvla), to produce smooth, continuous, and reactive motion despite having high inference latency. LeRobot provides the original inference-time guided mode and, for compatible Pi052 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,11 @@ 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 Pi052 checkpoint trained with `policy.rtc_training_max_delay > 0` and avoids the guidance backward pass.
**`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
@@ -141,6 +146,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 +157,24 @@ lerobot-rollout \
--device=cuda --device=cuda
``` ```
For a training-time RTC Pi052 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}/pi052_training_rtc \
--inference.type=rtc \
--inference.rtc.mode=trained \
--inference.rtc.execution_horizon=10 \
--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 +213,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
+82 -8
View File
@@ -70,6 +70,54 @@ 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 Pi052 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 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
_SAFETENSORS_FILE = "model.safetensors" _SAFETENSORS_FILE = "model.safetensors"
_SAFETENSORS_INDEX = "model.safetensors.index.json" _SAFETENSORS_INDEX = "model.safetensors.index.json"
@@ -164,21 +212,21 @@ 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.""" """Computes 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 # Broadcast the frequency dimension over either (B,) or (B, H).
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)
@@ -948,6 +996,24 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
) )
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 Pi052 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
if times is None: if times is None:
@@ -955,7 +1021,13 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
else: else:
time_tensor = times[step].expand(bsize) time_tensor = times[step].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,
@@ -963,7 +1035,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")
@@ -980,6 +1052,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)
@@ -121,6 +121,15 @@ class PI052Config(PI05Config):
# Reuse each VLM prefix across independent denoising draws; 1 restores single-draw flow. # Reuse each VLM prefix across independent denoising draws; 1 restores single-draw flow.
flow_num_repeats: int = 5 flow_num_repeats: int = 5
# Training-time RTC (arXiv:2512.05964). Zero preserves standard flow matching.
rtc_training_max_delay: int = 0
"""Largest clean action-prefix length sampled during training.
A value greater than zero enables training-time action conditioning. Each
flow draw samples a delay uniformly from ``[0, rtc_training_max_delay]``;
the corresponding action prefix stays clean and is excluded from the loss.
"""
# PaLM-style z-loss stabilizes large-vocabulary CE; 0 disables it. # PaLM-style z-loss stabilizes large-vocabulary CE; 0 disables it.
text_ce_z_loss_weight: float = 1e-4 text_ce_z_loss_weight: float = 1e-4
@@ -151,6 +160,11 @@ class PI052Config(PI05Config):
self.train_expert_only = False self.train_expert_only = False
if self.flow_num_repeats < 1: if self.flow_num_repeats < 1:
raise ValueError(f"flow_num_repeats must be >= 1, got {self.flow_num_repeats}") raise ValueError(f"flow_num_repeats must be >= 1, got {self.flow_num_repeats}")
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.manual_attention_scope not in {"all", "action"}: if self.manual_attention_scope not in {"all", "action"}:
raise ValueError( raise ValueError(
f"manual_attention_scope must be 'all' or 'action', got {self.manual_attention_scope!r}" f"manual_attention_scope must be 'all' or 'action', got {self.manual_attention_scope!r}"
+73 -10
View File
@@ -187,6 +187,54 @@ def _reduce_action_loss(per_sample: Tensor, predict_actions_t: Tensor | None, re
return (per_sample * mask).sum() / mask.sum().clamp(min=1.0) return (per_sample * mask).sum() / mask.sum().clamp(min=1.0)
def _sample_training_rtc_prefix_mask(
batch_size: int,
action_horizon: int,
max_delay: int,
device: torch.device,
) -> Tensor | None:
"""Sample per-draw clean prefixes for training-time RTC."""
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]:
"""Build noisy actions and scalar/per-token flow times.
LeRobot's PI0.5 flow uses ``t=0`` for clean data and ``t=1`` for
noise, the reverse of the notation in arXiv:2512.05964. Consequently,
clean RTC prefix tokens receive ``t=0`` here.
"""
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 _flow_loss_per_sample(flow_per_dim: Tensor, prefix_mask: Tensor | None) -> Tensor:
"""Average each draw over postfix action positions and dimensions only."""
if prefix_mask is None:
return flow_per_dim.flatten(start_dim=1).mean(dim=1)
postfix = (~prefix_mask).unsqueeze(-1).expand_as(flow_per_dim)
reduce_dims = tuple(range(1, flow_per_dim.ndim))
numerator = (flow_per_dim * postfix).sum(dim=reduce_dims)
denominator = postfix.sum(dim=reduce_dims).clamp(min=1)
return numerator / denominator
# Materialized logits win at VLA token counts; larger dense targets use Liger. # Materialized logits win at VLA token counts; larger dense targets use Liger.
_LOGITS_CE_MAX_POSITIONS = 2048 _LOGITS_CE_MAX_POSITIONS = 2048
@@ -928,6 +976,7 @@ class PI052Policy(PI05Policy):
text_labels is None text_labels is None
and predict_actions_t is None and predict_actions_t is None
and not getattr(self.config, "enable_fast_action_loss", False) and not getattr(self.config, "enable_fast_action_loss", False)
and self.config.rtc_training_max_delay == 0
): ):
return super().forward(batch, reduction=reduction) return super().forward(batch, reduction=reduction)
@@ -1118,12 +1167,17 @@ class PI052Policy(PI05Policy):
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)
time_expanded = time[:, None, None] prefix_mask = _sample_training_rtc_prefix_mask(
x_t = time_expanded * noise + (1 - time_expanded) * actions actions.shape[0],
actions.shape[1],
self.config.rtc_training_max_delay,
actions.device,
)
x_t, model_time = _build_flow_matching_inputs(actions, noise, time, prefix_mask)
u_t = noise - actions u_t = noise - actions
# ---- suffix: noisy actions ---------------------------------- # ---- suffix: noisy actions ----------------------------------
suffix_embs, suffix_pad, suffix_att, adarms_cond = self.model.embed_suffix(x_t, time) suffix_embs, suffix_pad, suffix_att, adarms_cond = self.model.embed_suffix(x_t, model_time)
# ---- bf16 alignment (mirrors PI05Pytorch.forward) ----------- # ---- bf16 alignment (mirrors PI05Pytorch.forward) -----------
first_layer = self.model.paligemma_with_expert.paligemma.model.language_model.layers[0] first_layer = self.model.paligemma_with_expert.paligemma.model.language_model.layers[0]
@@ -1169,7 +1223,7 @@ class PI052Policy(PI05Policy):
# internally to max_action_dim). # internally to max_action_dim).
original_action_dim = self.config.output_features[ACTION].shape[0] original_action_dim = self.config.output_features[ACTION].shape[0]
flow_per_dim = flow_per_dim[:, :, :original_action_dim] flow_per_dim = flow_per_dim[:, :, :original_action_dim]
per_sample_flow = flow_per_dim.mean(dim=(1, 2)) per_sample_flow = _flow_loss_per_sample(flow_per_dim, prefix_mask)
flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction) flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
return prefix_out, flow_loss return prefix_out, flow_loss
@@ -1216,10 +1270,15 @@ class PI052Policy(PI05Policy):
noise = model.sample_noise((k * batch_size, *actions.shape[1:]), actions.device) noise = model.sample_noise((k * batch_size, *actions.shape[1:]), actions.device)
time = model.sample_time(k * batch_size, actions.device) time = model.sample_time(k * batch_size, actions.device)
actions_rep = actions.repeat(k, 1, 1) # (k*B, chunk, motor_dim) actions_rep = actions.repeat(k, 1, 1) # (k*B, chunk, motor_dim)
time_expanded = time[:, None, None] prefix_mask_flat = _sample_training_rtc_prefix_mask(
x_t = time_expanded * noise + (1 - time_expanded) * actions_rep k * batch_size,
chunk,
self.config.rtc_training_max_delay,
actions.device,
)
x_t, model_time = _build_flow_matching_inputs(actions_rep, noise, time, prefix_mask_flat)
u_t = (noise - actions_rep).view(k, batch_size, chunk, -1).transpose(0, 1) # (B, k, chunk, motor) u_t = (noise - actions_rep).view(k, batch_size, chunk, -1).transpose(0, 1) # (B, k, chunk, motor)
s_embs, suffix_pad, suffix_att, adarms = model.embed_suffix(x_t, time) s_embs, suffix_pad, suffix_att, adarms = model.embed_suffix(x_t, model_time)
if use_bf16: if use_bf16:
s_embs = s_embs.to(dtype=torch.bfloat16) s_embs = s_embs.to(dtype=torch.bfloat16)
suffix_pad = suffix_pad[:batch_size] suffix_pad = suffix_pad[:batch_size]
@@ -1228,9 +1287,10 @@ class PI052Policy(PI05Policy):
s_embs.view(k, batch_size, chunk, -1).transpose(0, 1).reshape(batch_size, k * chunk, -1) s_embs.view(k, batch_size, chunk, -1).transpose(0, 1).reshape(batch_size, k * chunk, -1)
) # (B, k*chunk, D) ) # (B, k*chunk, D)
# Broadcast each draw's AdaRMS condition over its action chunk. # Broadcast each draw's AdaRMS condition over its action chunk.
if adarms.ndim == 2:
adarms = adarms[:, None, :].expand(-1, chunk, -1)
adarms_cond = ( adarms_cond = (
adarms.view(k, batch_size, 1, adarms.shape[-1]) adarms.view(k, batch_size, chunk, adarms.shape[-1])
.expand(k, batch_size, chunk, adarms.shape[-1])
.transpose(0, 1) .transpose(0, 1)
.reshape(batch_size, k * chunk, adarms.shape[-1]) .reshape(batch_size, k * chunk, adarms.shape[-1])
) # (B, k*chunk, cond_dim) ) # (B, k*chunk, cond_dim)
@@ -1300,7 +1360,10 @@ class PI052Policy(PI05Policy):
v_t = model.action_out_proj(suffix_out.to(dtype=torch.float32)) v_t = model.action_out_proj(suffix_out.to(dtype=torch.float32))
v_t = v_t.view(batch_size, k, chunk, -1) # (B, k, chunk, motor) v_t = v_t.view(batch_size, k, chunk, -1) # (B, k, chunk, motor)
flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")[..., :original_action_dim] flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")[..., :original_action_dim]
per_sample_flow = flow_per_dim.mean(dim=(1, 2, 3)) prefix_mask = None
if prefix_mask_flat is not None:
prefix_mask = prefix_mask_flat.view(k, batch_size, chunk).transpose(0, 1)
per_sample_flow = _flow_loss_per_sample(flow_per_dim, prefix_mask)
flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction) flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
return prefix_out, flow_loss return prefix_out, flow_loss
@@ -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:
+12
View File
@@ -178,6 +178,18 @@ def build_rollout_context(
policy_config = cfg.policy policy_config = cfg.policy
policy_class = get_policy_class(policy_config.type) policy_class = get_policy_class(policy_config.type)
if is_rtc and cfg.inference.rtc.enabled and cfg.inference.rtc.mode == "trained":
if policy_config.type != "pi052":
raise ValueError(
"--inference.rtc.mode=trained currently requires a Pi052 checkpoint; "
f"got policy type {policy_config.type!r}."
)
if int(getattr(policy_config, "rtc_training_max_delay", 0)) <= 0:
raise ValueError(
"--inference.rtc.mode=trained requires a checkpoint trained with "
"--policy.rtc_training_max_delay > 0."
)
if hasattr(policy_config, "compile_model"): if hasattr(policy_config, "compile_model"):
policy_config.compile_model = cfg.use_torch_compile policy_config.compile_model = cfg.use_torch_compile
@@ -0,0 +1,106 @@
#!/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.
"""Unit coverage for Pi052 training-time RTC conditioning."""
from types import SimpleNamespace
import pytest
import torch
from torch import nn
pytest.importorskip("transformers")
from lerobot.policies.pi05.modeling_pi05 import ( # noqa: E402
_prepare_trained_rtc_prefix,
create_sinusoidal_pos_embedding,
)
from lerobot.policies.pi052.configuration_pi052 import PI052Config # noqa: E402
from lerobot.policies.pi052.modeling_pi052 import ( # noqa: E402
PI05Pytorch as PI052Pytorch,
_build_flow_matching_inputs,
_flow_loss_per_sample,
)
def test_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_training_rtc_loss_averages_over_postfix_only():
flow_loss = torch.tensor([[[100.0], [100.0], [2.0], [4.0]]])
prefix_mask = torch.tensor([[True, True, False, False]])
per_sample = _flow_loss_per_sample(flow_loss, prefix_mask)
assert per_sample.tolist() == [3.0]
def test_per_token_time_embedding_preserves_action_axis():
time = torch.tensor([[0.0, 0.5, 1.0]])
embedding = create_sinusoidal_pos_embedding(time, 8, 4e-3, 4.0, time.device)
assert embedding.shape == (1, 3, 8)
assert not torch.equal(embedding[:, 0], embedding[:, 1])
def test_action_expert_embeds_per_token_flow_times():
model = PI052Pytorch.__new__(PI052Pytorch)
nn.Module.__init__(model)
model.config = SimpleNamespace(chunk_size=3, min_period=4e-3, max_period=4.0)
model.gradient_checkpointing_enabled = False
model.action_in_proj = nn.Linear(2, 8)
model.time_mlp_in = nn.Linear(8, 8)
model.time_mlp_out = nn.Linear(8, 8)
suffix, _, _, adarms_cond = model.embed_suffix(
torch.randn(2, 3, 2),
torch.tensor([[0.0, 0.5, 0.5], [0.0, 0.0, 0.5]]),
)
assert suffix.shape == (2, 3, 8)
assert adarms_cond.shape == (2, 3, 8)
def test_trained_rtc_prefix_is_padded_and_masked():
x_t = torch.randn(1, 5, 4)
previous = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
prefix, mask = _prepare_trained_rtc_prefix(x_t, previous, inference_delay=2, training_max_delay=3)
assert prefix.shape == x_t.shape
assert mask.shape == x_t.shape
assert torch.equal(prefix[0, :2, :2], previous[:2])
assert torch.count_nonzero(prefix[0, :2, 2:]) == 0
assert mask[0, :2].all()
assert not mask[0, 2:].any()
def test_trained_rtc_rejects_delay_outside_training_distribution():
with pytest.raises(ValueError, match="exceeds the checkpoint"):
_prepare_trained_rtc_prefix(
torch.randn(1, 5, 4),
torch.randn(4, 2),
inference_delay=4,
training_max_delay=3,
)
@pytest.mark.parametrize("max_delay", [-1, 5])
def test_pi052_config_rejects_invalid_training_rtc_delay(max_delay):
with pytest.raises(ValueError, match="rtc_training_max_delay"):
PI052Config(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