Compare commits

...

36 Commits

Author SHA1 Message Date
Khalil Meftah be59464e7e fix(rewards): simplify VF architecture, remove CLS token, use bidirectional mean-pool or last_token readout 2026-07-20 18:43:28 +02:00
Khalil Meftah c8e32d1afe fix(rewards): normalize returns by (H - 1 + C_fail) 2026-07-19 13:04:15 +02:00
Khalil Meftah 62597032b9 refactor(rewards): add get_optim_params() with vision_encoder_lr_multiplier for differential LR 2026-07-18 20:18:01 +02:00
Khalil Meftah d0348b1803 refactor(rewards): rewrite distributional VF with SigLIP2 + Gemma3 backbone
Replace PaliGemma-based value function with a monolithic VLM architecture
using SigLIP2-so400m (vision encoder) and Gemma3-270M (shared backbone),
matching the pi*0.6 paper.

Key changes:
- SigLIP2-so400m as vision encoder, Gemma3-270M as unified transformer
- [CLS] token readout with bidirectional prefix attention
- 2-layer MLP value head (Linear -> LayerNorm -> GELU -> Dropout -> Linear)
- Multi-camera support with per-camera validity masks
- Image preprocessing in processor (resize with pad, normalize to [-1,1])
- Freeze controls for vision encoder and language model independently
2026-07-18 19:31:03 +02:00
Khalil Meftah 535371a5b8 feat(aggregate): initialize language type columns in aggregate_data function 2026-07-13 14:58:53 +02:00
Khalil Meftah 0be63969f3 feat(annotation): global advantage threshold
- Add global_threshold config option (default True)
- Add precompute_global_threshold() method with caching
2026-07-13 14:12:55 +02:00
Khalil Meftah d0f3619ef0 feat(dataset): merge datasets with different features 2026-07-13 11:57:33 +02:00
Khalil Meftah d0cb001b9c feat(compute-returns): push to hub 2026-07-13 11:42:55 +02:00
Khalil Meftah 348efac2bd fix(so100): add retry parameter to sync_read and sync_write methods 2026-07-12 17:31:09 +02:00
Khalil Meftah 81b6ea1669 add recipe for advantage annotation with dropout for cfg 2026-07-11 11:11:13 +02:00
Khalil Meftah 235e88c743 fix(annotation): remove dropout when doing annotation 2026-07-11 11:09:24 +02:00
Khalil Meftah 5fccaf0477 fix(annotate): add frame_provider field to AdvantageModule dataclass 2026-07-10 14:06:12 +02:00
Khalil Meftah 235bc3a78a fix(processor): add lazy import fallback for unregistered processor steps 2026-07-10 14:05:29 +02:00
Khalil Meftah c043a6c418 fix(processor): make RenderMessagesStep.recipe optional for non-recipe pipelines 2026-07-10 14:05:08 +02:00
Khalil Meftah f5c2ee1753 fix(annotate): remove partial download 2026-07-10 14:04:43 +02:00
Khalil Meftah 6adb74b05f feat(recap): implement CFGRL 2026-07-08 14:57:09 +02:00
Khalil Meftah 407a8c1d7d feat(annotate): support video datasets in VF advantage scoring 2026-07-08 12:16:01 +02:00
Khalil Meftah 3c3f3bdf61 fix(processor): pass RL keys through pipeline 2026-07-08 12:13:20 +02:00
Khalil Meftah 582e953676 fix(train): fix VF scheduler config and reward model hub push
Add missing peak_lr and decay_lr to DistributionalVFConfig scheduler
preset. Fix push_model_to_hub call for reward models.
2026-07-08 12:08:17 +02:00
Khalil Meftah 9a846c4fca fix(advantage): update frame count calculation in constant mode 2026-07-03 17:39:23 +02:00
Khalil Meftah ad32d3e00d fix(annotation): skip vlm initialization when using advantage module 2026-07-03 17:33:38 +02:00
Khalil Meftah 1cd1ec468e feat(molmoact2): add RECAP advantage conditioning via recipe system for MolmoAct2
- Add recipe_path, advantage_prefix, cfg_beta to MolmoAct2Config
- Place advantage clause in the assistant section of _build_robot_text
- Add MolmoAct2NormalizeTaskStep for consistent task normalization
- Parse recipe-rendered advantage in MolmoAct2PackInputsProcessorStep
- Insert RenderMessagesStep pipeline when recipe_path is configured
- Add recap_advantage_molmoact2.yaml recipe
2026-07-03 16:58:12 +02:00
Khalil Meftah 79b7f992b4 feat(annotate): add constant advantage labeling for RECAP SFT phase
- Add constant_value and seed fields to AdvantageConfig
- Implement _run_constant_mode in AdvantageModule with CFG dropout
- Use deterministic seeding (config.seed + episode_index) for reproducibility
2026-07-03 16:58:12 +02:00
Khalil Meftah 04a39d419d feat(pi05): implement Classifier-Free Guidance (CFG) inference
Add dual-path denoising with configurable cfg_beta scale for language-
conditioned action generation. When cfg_beta > 1.0, VLM prefills both
conditioned and unconditional prompts, and action expert velocities are
interpolated via v = v_uncond + β*(v_cond - v_uncond).
2026-07-03 16:58:11 +02:00
Khalil Meftah b63a714ae9 feat(pi05): integrate RenderMessagesStep for advantage conditioning
Add RenderedMessagesToTaskStep adapter that bridges recipe-rendered chat
messages back into PI05's task-string prompt format. When recipe_path is
set on PI05Config, the preprocessor inserts RenderMessagesStep + adapter
before prompt construction, enabling RECAP advantage text to flow
end-to-end through the recipe YAML system.
2026-07-03 16:58:11 +02:00
Khalil Meftah 2ded9ba783 feat(rollout): add episode success labeling to DAgger strategy 2026-07-03 16:58:08 +02:00
Khalil Meftah 194a6379ea feat(recap): add advantage conditioning recipe YAMLs 2026-07-03 16:55:41 +02:00
Khalil Meftah cc782e3589 feat(recap): add advantage scoring annotation module
Implement the RECAP advantage scoring module as a new phase in
lerobot-annotate. Uses a frozen distributional VF to compute per-frame
advantages, binarizes into positive/negative indicators with per-task
threshold, and writes style=advantage persistent rows for policy
conditioning. Skips VF inference on intervention frames as an optimization.
2026-07-03 16:55:40 +02:00
Khalil Meftah b90ccd283b feat(recap): add lerobot-compute-returns script to compute MC returns 2026-07-03 16:55:40 +02:00
Khalil Meftah f8fa8ba394 test(rewards): add unit tests for distributional value function model 2026-07-03 16:55:40 +02:00
Khalil Meftah 6663cac584 feat(rewards): introduce distributional value function model
- Added a new distributional value function (DistributionalVF) model for RECAP, including its configuration, modeling, and processor components.
- Updated the rewards factory to support the new model type.
- Updated  to include the new model in the dependencies.
2026-07-03 16:55:28 +02:00
Khalil Meftah 4af7095693 Merge branch 'main' into feat/rollout/dagger-episode-save 2026-07-03 16:50:10 +02:00
Khalil Meftah 46d4ddc698 chore(rollout): log episode success label and buffer length 2026-07-02 19:12:10 +02:00
Khalil Meftah b29ba27977 fix(rollout): guard empty buffer save 2026-07-02 18:02:59 +02:00
Khalil Meftah 599e2432e5 fix(rollout): clear last_action after return_to_initial 2026-07-02 18:02:36 +02:00
Khalil Meftah 44f76dbbf0 feat(rollout): add episode success/failure labeling to DAgger strategy
Enable operators to mark episodes as success or failure during DAgger
data collection. Pressing 's' or 'f' immediately saves the episode
with the appropriate label and returns the robot to its initial position.

- Add success/failure key bindings to DAggerKeyboardConfig
- Add save_episode_requested event and episode_success state to DAggerEvents
- Stamp next.success=True on terminal frame for successful episodes
- Pause and return to initial position after manual save for env reset
- Add num_episodes target to stop continuous recording automatically
- Defer save during corrections to avoid splitting mid-intervention
2026-07-02 17:48:02 +02:00
52 changed files with 6374 additions and 1587 deletions
+3
View File
@@ -228,6 +228,7 @@ groot = [
sarm = ["lerobot[transformers-dep]", "pydantic>=2.0.0,<3.0.0", "faker>=33.0.0,<35.0.0", "lerobot[matplotlib-dep]", "lerobot[qwen-vl-utils-dep]"]
robometer = ["lerobot[transformers-dep]", "lerobot[qwen-vl-utils-dep]", "lerobot[peft-dep]"]
topreward = ["lerobot[transformers-dep]"]
recap = ["lerobot[transformers-dep]"]
xvla = ["lerobot[transformers-dep]"]
eo1 = ["lerobot[transformers-dep]", "lerobot[qwen-vl-utils-dep]"]
fastwam = [
@@ -332,6 +333,7 @@ all = [
"lerobot[sarm]",
"lerobot[robometer]",
"lerobot[topreward]",
"lerobot[recap]",
"lerobot[peft]",
# "lerobot[unitree_g1]", TODO: Unitree requires specific installation instructions for unitree_sdk2
]
@@ -355,6 +357,7 @@ lerobot-edit-dataset="lerobot.scripts.lerobot_edit_dataset:main"
lerobot-setup-can="lerobot.scripts.lerobot_setup_can:main"
lerobot-annotate="lerobot.scripts.lerobot_annotate:main"
lerobot-rollout="lerobot.scripts.lerobot_rollout:main"
lerobot-compute-returns="lerobot.scripts.lerobot_compute_returns:main"
# ---------------- Tool Configurations ----------------
@@ -169,6 +169,49 @@ class ExecutorConfig:
episode_parallelism: int = 16
@dataclass
class AdvantageConfig:
"""``advantage`` module: RECAP advantage scoring via frozen value function."""
enabled: bool = True
# Constant advantage label for all frames (e.g. "positive" for SFT iteration 0).
# Skips VF inference.
constant_value: str | None = None
# Trained value function checkpoint (local path or Hub repo ID).
# Ignored when constant_value is set.
value_function_path: str = ""
# Device to run the value function on.
device: str = "cuda"
# N-step lookahead for advantage estimation.
# None = MC (N=T): A_t = R_t - V(s_t), using mc_return from dataset.
# 50 = fine-tuning mode: A_t = Σ r_{t:t+N} + V(s_{t+N}) - V(s_t).
n_step: int | None = None
# Per-task percentile for binarization threshold ε_.
# Actions with advantage > ε_ get I_t = True (positive).
threshold_percentile: float = 0.3
# When True, compute a single global threshold across all episodes (paper behavior).
# When False, compute threshold per-episode (faster but less accurate).
global_threshold: bool = True
# Force I_t = True for frames marked as human interventions.
force_positive_on_intervention: bool = True
# Column name in dataset for intervention flag.
intervention_key: str = "intervention"
# Column name for pre-computed MC returns (from lerobot-compute-returns).
mc_return_key: str = "mc_return"
# Batch size for value function inference.
batch_size: int = 32
@dataclass
class AnnotationPipelineConfig:
"""Top-level config for ``lerobot-annotate`` (rewrites data shards in place)."""
@@ -190,6 +233,7 @@ class AnnotationPipelineConfig:
plan: PlanConfig = field(default_factory=PlanConfig)
interjections: InterjectionsConfig = field(default_factory=InterjectionsConfig)
vqa: VqaConfig = field(default_factory=VqaConfig)
advantage: AdvantageConfig = field(default_factory=AdvantageConfig)
vlm: VlmConfig = field(default_factory=VlmConfig)
executor: ExecutorConfig = field(default_factory=ExecutorConfig)
@@ -15,20 +15,24 @@
# limitations under the License.
"""In-process executor that runs the annotation phases.
The executor runs **six phases** in dependency order:
The executor runs **seven phases** in dependency order:
phase 1: ``plan`` module (plan + subtasks + memory)
phase 2: ``interjections`` module (interjections + speech)
phase 3: ``plan`` plan-update pass — re-runs plan emission at every
interjection timestamp produced by phase 2
phase 4: ``vqa`` module (VQA)
phase 5: validator
phase 6: writer
phase 5: ``advantage`` module (advantage scoring via frozen VF)
phase 6: validator
phase 7: writer
Phase 3 is why the ``plan`` module must be re-entered after the
``interjections`` module — to refresh ``plan`` rows at interjection
timestamps.
Phase 5 (advantage) does not depend on the VLM modules, it uses a frozen
distributional value function to compute per-frame advantage indicators.
Distributed execution is provided by Hugging Face Jobs (see
``examples/annotations/run_hf_job.py``); the runner inside the job
invokes ``lerobot-annotate`` which uses this in-process executor.
@@ -74,7 +78,7 @@ class PipelineRunSummary:
@dataclass
class Executor:
"""Run all six phases over a dataset root in-process.
"""Run all seven phases over a dataset root in-process.
Episode-level concurrency comes from ``ExecutorConfig.episode_parallelism``
(a thread pool); cluster-level concurrency comes from running this
@@ -86,6 +90,7 @@ class Executor:
plan: Any # PlanSubtasksMemoryModule
interjections: Any # InterjectionsAndSpeechModule
vqa: Any # GeneralVqaModule
advantage: Any # AdvantageModule
writer: LanguageColumnsWriter
validator: StagingValidator
@@ -112,6 +117,12 @@ class Executor:
phases.append(self._run_plan_update_phase(records, staging_dir))
# Phase 4: ``vqa`` module (VQA)
phases.append(self._run_module_phase("vqa", records, staging_dir, self.vqa))
# Phase 5: ``advantage`` module (advantage scoring via frozen VF)
# Two-pass global threshold: compute advantages across all episodes first,
# then apply the single threshold uniformly (matches paper Section V-D).
if self.advantage.enabled and self.advantage.config.global_threshold:
self.advantage.precompute_global_threshold(records)
phases.append(self._run_module_phase("advantage", records, staging_dir, self.advantage))
print("[annotate] running validator...", flush=True)
report = self.validator.validate(records, staging_dir)
@@ -179,7 +190,7 @@ class Executor:
staging_dir: Path,
module: Any,
) -> PhaseResult:
if not module.enabled:
if module is None or not module.enabled:
print(f"[annotate] phase={name} skipped (module disabled)", flush=True)
return PhaseResult(name=name, episodes_processed=0, episodes_skipped=len(records))
n = len(records)
@@ -231,7 +242,7 @@ class Executor:
``plan`` module with the interjection timestamps so its existing
prompt path is reused.
"""
if not self.plan.enabled or not self.interjections.enabled:
if not self.plan or not self.plan.enabled or not self.interjections or not self.interjections.enabled:
return PhaseResult(name="plan_update", episodes_processed=0, episodes_skipped=len(records))
processed = 0
for record in records:
@@ -14,11 +14,13 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from .advantage import AdvantageModule
from .general_vqa import GeneralVqaModule
from .interjections_and_speech import InterjectionsAndSpeechModule
from .plan_subtasks_memory import PlanSubtasksMemoryModule
__all__ = [
"AdvantageModule",
"GeneralVqaModule",
"InterjectionsAndSpeechModule",
"PlanSubtasksMemoryModule",
@@ -0,0 +1,378 @@
#!/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.
# 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.
"""Advantage scoring module for RECAP.
Computes per-frame advantage values using a frozen distributional value function,
binarizes them into improvement indicators (I_t), and emits ``style="advantage"``
persistent rows for policy conditioning.
Paper reference: pi*0.6, Section IV-B and Appendix F.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any
import numpy as np
import torch
from ..config import AdvantageConfig
from ..frames import VideoFrameProvider, null_provider
from ..reader import EpisodeRecord
from ..staging import EpisodeStaging
logger = logging.getLogger(__name__)
@dataclass
class AdvantageModule:
"""Compute advantage indicators and emit persistent annotation rows.
The module loads a frozen distributional value function and scores each
frame in an episode. Advantages are binarized into ``positive``/``negative``
indicators using a per-task threshold, then written as ``style="advantage"``
persistent rows into the staging area.
Requires ``mc_return`` column in the dataset (from lerobot-compute-returns).
"""
config: AdvantageConfig
frame_provider: Any = None
_model: Any = field(default=None, init=False, repr=False)
_preprocessor: Any = field(default=None, init=False, repr=False)
_threshold: float | None = field(default=None, init=False, repr=False)
_cache: dict = field(default_factory=dict, init=False, repr=False)
@property
def enabled(self) -> bool:
return self.config.enabled
def _ensure_model_loaded(self) -> None:
"""Lazy-load the frozen value function on first use."""
if self._model is not None:
return
from lerobot.rewards import (
make_reward_model,
make_reward_model_config,
make_reward_pre_post_processors,
)
cfg = make_reward_model_config(
"distributional_value_function",
pretrained_path=self.config.value_function_path,
device=self.config.device,
)
self._model = make_reward_model(cfg)
self._model.eval()
for p in self._model.parameters():
p.requires_grad_(False)
self._preprocessor, _ = make_reward_pre_post_processors(cfg)
logger.info("Loaded frozen VF from %s on %s", self.config.value_function_path, self.config.device)
def compute_advantages_for_episode(self, record: EpisodeRecord) -> tuple[np.ndarray, np.ndarray]:
"""Compute raw advantage values for all frames in an episode.
Returns:
(advantages, intervention_mask) both shape [num_frames].
advantages[t] = A_t, intervention_mask[t] = True if frame is intervention.
"""
self._ensure_model_loaded()
df = record.frames_df()
num_frames = len(df)
mc_return_key = self.config.mc_return_key
if mc_return_key not in df.columns:
raise KeyError(
f"Column '{mc_return_key}' not found in episode {record.episode_index}. "
"Run lerobot-compute-returns first."
)
mc_returns = df[mc_return_key].values.astype(np.float32)
intervention_mask = np.zeros(num_frames, dtype=bool)
if self.config.intervention_key in df.columns:
intervention_mask = df[self.config.intervention_key].values.astype(bool)
# Skip VF inference on intervention frames — they're always "positive"
# regardless of advantage value, so V(s_t) is never used for them.
skip_mask = intervention_mask if self.config.force_positive_on_intervention else None
values = self._compute_values(record, skip_mask=skip_mask)
if self.config.n_step is None:
advantages = mc_returns - values
else:
advantages = self._compute_n_step_advantages(mc_returns, values, record, n=self.config.n_step)
return advantages, intervention_mask
def _compute_values(self, record: EpisodeRecord, skip_mask: np.ndarray | None = None) -> np.ndarray:
"""Run frozen VF over all frames to get V(s_t) predictions.
Supports both image datasets (columns in parquet) and video datasets
(frames decoded from .mp4 via the shared VideoFrameProvider).
Args:
record: Episode data.
skip_mask: Optional boolean mask [num_frames]. Frames where True are
skipped (left as 0.0) to avoid unnecessary inference.
"""
df = record.frames_df()
num_frames = len(df)
values = np.zeros(num_frames, dtype=np.float32)
# Determine which frame indices actually need inference
infer_indices = np.where(~skip_mask)[0] if skip_mask is not None else np.arange(num_frames)
if len(infer_indices) == 0:
return values
# Try parquet image columns first, fall back to video decoding
image_key = self._resolve_image_key(df)
video_frames = None
if image_key is None:
image_key, video_frames = self._decode_video_frames(record, infer_indices)
if image_key is None:
logger.warning(
"No image/video key found for episode %d; returning zero values.", record.episode_index
)
return values
task_text = record.episode_task
for batch_start in range(0, len(infer_indices), self.config.batch_size):
batch_end = min(batch_start + self.config.batch_size, len(infer_indices))
batch_indices = infer_indices[batch_start:batch_end]
batch_images = []
for local_i in range(len(batch_indices)):
if video_frames is not None:
img_tensor = video_frames[batch_start + local_i].float()
else:
idx = batch_indices[local_i]
img_val = df.iloc[idx][image_key]
if isinstance(img_val, np.ndarray):
img_tensor = torch.from_numpy(img_val).float()
elif isinstance(img_val, torch.Tensor):
img_tensor = img_val.float()
else:
img_tensor = torch.zeros(3, 224, 224)
batch_images.append(img_tensor)
batch_images_tensor = torch.stack(batch_images)
batch_size = batch_images_tensor.shape[0]
raw_batch = {
image_key: batch_images_tensor,
"task": [task_text] * batch_size,
}
processed = self._preprocessor(raw_batch)
with torch.no_grad():
v_values = self._model.compute_reward(processed)
values[batch_indices] = v_values.cpu().numpy()
return values
def _decode_video_frames(
self, record: EpisodeRecord, infer_indices: np.ndarray
) -> tuple[str | None, torch.Tensor | None]:
"""Decode video frames using the existing VideoFrameProvider infrastructure.
Returns (image_key, decoded_frames_tensor) or (None, None) on failure.
"""
dataset_root = record.data_path.parent.parent.parent
if not hasattr(self, "_frame_provider") or self._frame_provider is None:
try:
self._frame_provider = VideoFrameProvider(root=dataset_root)
except Exception:
self._frame_provider = null_provider()
if not self._frame_provider.camera_keys:
return None, None
camera_key = self._frame_provider.camera_keys[0]
timestamps = [float(record.frame_timestamps[i]) for i in infer_indices]
frames = self._frame_provider.frames_at(record, timestamps, camera_key=camera_key)
if not frames:
return None, None
frames_tensor = torch.stack(frames)
return camera_key, frames_tensor
def _compute_n_step_advantages(
self, mc_returns: np.ndarray, values: np.ndarray, record: EpisodeRecord, n: int
) -> np.ndarray:
"""Compute N-step advantage: A_t = Σ r_{t:t+N-1} + V(s_{t+N}) - V(s_t).
When t+N exceeds episode length, truncates to MC (uses mc_return directly).
"""
num_frames = len(values)
advantages = np.zeros(num_frames, dtype=np.float32)
for t in range(num_frames):
if t + n >= num_frames:
advantages[t] = mc_returns[t] - values[t]
else:
n_step_return = mc_returns[t] - mc_returns[t + n]
advantages[t] = n_step_return + values[t + n] - values[t]
return advantages
def _resolve_image_key(self, df) -> str | None:
"""Find the first image observation key in the dataframe columns."""
for col in df.columns:
if col.startswith("observation.images."):
return col
return None
def precompute_global_threshold(self, records: list[EpisodeRecord]) -> None:
"""Two-pass: compute advantages for all episodes and set a single global threshold.
This matches the paper (pi*0.6, Section V-D / Appendix F):
'We set ε_ to the Nth percentile of values predicted by the value function for the task .'
The threshold is computed across ALL non-intervention frames in the dataset,
so successful episodes naturally get more 'positive' labels and failed episodes
get more 'negative' labels.
"""
if self.config.constant_value:
return
if not self.config.value_function_path:
return
logger.info("Computing global advantage threshold (two-pass mode)...")
all_advantages: list[float] = []
for record in records:
advantages, intervention_mask = self.compute_advantages_for_episode(record)
self._cache[record.episode_index] = (advantages, intervention_mask)
non_intervention = advantages[~intervention_mask] if intervention_mask.any() else advantages
all_advantages.extend(non_intervention.tolist())
if not all_advantages:
self._threshold = 0.0
else:
self._threshold = float(np.percentile(all_advantages, self.config.threshold_percentile * 100))
num_positive = sum(1 for a in all_advantages if a > self._threshold)
logger.info(
"Global threshold: %.4f (percentile=%.0f%%, %d/%d frames positive = %.1f%%)",
self._threshold,
self.config.threshold_percentile * 100,
num_positive,
len(all_advantages),
100 * num_positive / max(len(all_advantages), 1),
)
def run_episode(self, record: EpisodeRecord, staging: EpisodeStaging) -> None:
"""Score one episode and write advantage rows to staging."""
if self.config.constant_value:
self._run_constant_mode(record, staging)
return
if not self.config.value_function_path:
logger.warning("No value_function_path or constant_value configured; skipping advantage scoring.")
return
if record.episode_index in self._cache:
advantages, intervention_mask = self._cache.pop(record.episode_index)
else:
advantages, intervention_mask = self.compute_advantages_for_episode(record)
num_frames = len(advantages)
if self._threshold is not None:
threshold = self._threshold
else:
threshold = self._compute_threshold(advantages, intervention_mask)
rows: list[dict[str, Any]] = []
for t in range(num_frames):
if (
self.config.force_positive_on_intervention
and intervention_mask[t]
or advantages[t] > threshold
):
indicator = "positive"
else:
indicator = "negative"
timestamp = float(record.frame_timestamps[t]) if t < len(record.frame_timestamps) else 0.0
rows.append(
{
"role": "user",
"content": indicator,
"style": "advantage",
"timestamp": timestamp,
"camera": None,
"tool_calls": None,
}
)
staging.write("advantage", rows)
logger.debug(
"Episode %d: %d/%d frames scored (threshold=%.4f, %d positive, %d negative)",
record.episode_index,
len(rows),
num_frames,
threshold,
sum(1 for r in rows if r["content"] == "positive"),
sum(1 for r in rows if r["content"] == "negative"),
)
def _run_constant_mode(self, record: EpisodeRecord, staging: EpisodeStaging) -> None:
"""Emit a fixed advantage value for every frame."""
num_frames = len(record.frame_timestamps)
rows: list[dict[str, Any]] = []
for t in range(num_frames):
rows.append(
{
"role": "user",
"content": self.config.constant_value,
"style": "advantage",
"timestamp": float(record.frame_timestamps[t]),
"camera": None,
"tool_calls": None,
}
)
staging.write("advantage", rows)
logger.debug(
"Episode %d: %d/%d frames labeled constant '%s'",
record.episode_index,
len(rows),
num_frames,
self.config.constant_value,
)
def _compute_threshold(self, advantages: np.ndarray, intervention_mask: np.ndarray) -> float:
"""Compute the binarization threshold as the configured percentile of advantages."""
non_intervention = advantages[~intervention_mask] if intervention_mask.any() else advantages
if len(non_intervention) == 0:
return 0.0
return float(np.percentile(non_intervention, self.config.threshold_percentile * 100))
@@ -31,6 +31,7 @@ rows into memory at once.
from __future__ import annotations
import functools
from collections.abc import Iterator, Sequence
from dataclasses import dataclass, field
from pathlib import Path
@@ -42,6 +43,18 @@ from lerobot.datasets.io_utils import load_tasks
from lerobot.datasets.utils import DEFAULT_TASKS_PATH
@functools.lru_cache(maxsize=8)
def _read_parquet_as_pandas(path: Path): # type: ignore[no-untyped-def]
"""Read a parquet shard once and cache the pandas DataFrame.
Multiple EpisodeRecords from the same shard share this single read.
The LRU cache (keyed by path) avoids re-reading the same file
across 100+ episodes that all live in one chunk.
"""
return pq.read_table(path).to_pandas()
@dataclass
class EpisodeRecord:
"""Per-episode record yielded by the reader."""
@@ -61,10 +74,7 @@ class EpisodeRecord:
def frames_df(self): # type: ignore[no-untyped-def]
"""Lazy-load the pandas slice for this episode (memoized)."""
if self._frames_df_cache is None:
import pandas as pd # noqa: PLC0415 - deferred for optional dataset extra
table = pq.read_table(self.data_path)
df: pd.DataFrame = table.to_pandas()
df = _read_parquet_as_pandas(self.data_path)
self._frames_df_cache = df.iloc[self.row_offset : self.row_offset + self.row_count].reset_index(
drop=True
)
@@ -39,6 +39,7 @@ _MODULES: tuple[ModuleName, ...] = (
"plan",
"interjections",
"vqa",
"advantage",
)
+1
View File
@@ -32,6 +32,7 @@ DEFAULT_BINDINGS = {
"interjection": "emitted_at(t, style=interjection)",
"vqa": "emitted_at(t, style=vqa, role=assistant)",
"vqa_query": "emitted_at(t, style=vqa, role=user)",
"advantage": "active_at(t, style=advantage)",
}
PLACEHOLDER_RE = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}")
@@ -0,0 +1,30 @@
# RECAP advantage-conditioned recipe.
#
# Composes task + advantage indicator into the prompt for conditional SFT.
# The advantage binding resolves to "positive" or "negative" from the
# language_persistent column (written by lerobot-annotate --advantage).
# When advantage is absent (30% dropout), the advantage turn is skipped
# entirely via if_present, training the unconditional branch for CFG.
#
# This recipe is policy-agnostic: any VLA that consumes chat-style messages
# can use it. Override bindings or add blend components for task-specific needs.
#
# Paper: pi*0.6, Section IV-B (conditional policy training with I_t).
bindings:
advantage: "active_at(t, style=advantage)"
messages:
- role: user
content: "${task}"
stream: high_level
- role: user
content: "Advantage: ${advantage}"
stream: high_level
if_present: advantage
- role: assistant
content: "${subtask}"
stream: low_level
target: true
@@ -0,0 +1,41 @@
# RECAP full recipe with advantage conditioning and subtask blending.
#
# Blend of two training modes:
# 1. advantage_conditioned (70%): Task + advantage indicator → action
# 2. unconditional (30%): Task only → action (no advantage, trains CFG baseline)
#
# This achieves the same effect as per-frame dropout in the annotation module
# but at the recipe level, giving explicit control over the conditioning ratio.
# Use this instead of annotation-level dropout if you want a fixed split.
#
# Paper: pi*0.6, Appendix E (classifier-free guidance requires both branches).
blend:
advantage_conditioned:
weight: 0.7
messages:
- role: user
content: "${task}\nAdvantage: ${advantage}"
stream: high_level
if_present: advantage
- role: user
content: "${task}"
stream: high_level
- role: assistant
content: "${subtask}"
stream: low_level
target: true
unconditional:
weight: 0.3
messages:
- role: user
content: "${task}"
stream: high_level
- role: assistant
content: "${subtask}"
stream: low_level
target: true
@@ -0,0 +1,28 @@
# RECAP advantage recipe for MolmoAct2.
#
# Renders task + advantage into the task field as "<task> Advantage: <value>".
# MolmoAct2PackInputsProcessorStep parses this, extracts the advantage value,
# and places it AFTER the full user prompt but BEFORE action tokens — matching
# the RECAP paper (Section V-B): "The advantage indicator appears in the training
# sequence after ˆℓ but before the actions, such that only the action
# log-likelihoods are affected."
#
# Final prompt layout:
# <images><|im_start|>user\n{prompt}<|im_end|>\n<|im_start|>assistant\nAdvantage: positive. <action_output>...
#
# When advantage is absent (CFG dropout), if_present guard skips this message
# and RenderedMessagesToTaskStep leaves the task unchanged — no advantage clause.
bindings:
advantage: "active_at(t, style=advantage)"
messages:
- role: user
content: "${task} Advantage: ${advantage}"
stream: high_level
if_present: advantage
- role: assistant
content: ""
stream: low_level
target: true
@@ -0,0 +1,43 @@
# RECAP advantage recipe for MolmoAct2 with CFG blend (training-time dropout).
#
# Two components selected per sample:
# 1. advantage_conditioned (70%): Task + advantage indicator → action
# 2. unconditional (30%): Task only → action (no advantage, trains CFG baseline)
#
# At inference, classifier-free guidance combines both:
# action = action_uncond + w * (action_cond - action_uncond)
#
# Paper: pi*0.6, Appendix E & F.
bindings:
advantage: "active_at(t, style=advantage)"
blend:
advantage_conditioned:
weight: 0.7
messages:
- role: user
content: "${task} Advantage: ${advantage}"
stream: high_level
if_present: advantage
- role: user
content: "${task}"
stream: high_level
- role: assistant
content: ""
stream: low_level
target: true
unconditional:
weight: 0.3
messages:
- role: user
content: "${task}"
stream: high_level
- role: assistant
content: ""
stream: low_level
target: true
+52 -8
View File
@@ -92,7 +92,7 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta
return merged_info
def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]):
def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata], lenient: bool = False):
"""Validates that all dataset metadata have consistent properties.
Ensures all datasets have the same fps, robot_type, and features to guarantee
@@ -101,13 +101,16 @@ def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]):
Args:
all_metadata: List of LeRobotDatasetMetadata objects to validate.
lenient: If True, allow feature mismatches and return the union of all features.
Missing columns will be filled with default values during aggregation.
Returns:
tuple: A tuple containing (fps, robot_type, features) from the first metadata.
tuple: A tuple containing (fps, robot_type, features) from the first metadata
(or union of features if lenient=True).
Raises:
ValueError: If any metadata has different fps, robot_type, or features
than the first metadata in the list.
than the first metadata in the list (unless lenient=True for features).
"""
fps = all_metadata[0].fps
@@ -122,9 +125,15 @@ def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]):
f"Same robot_type is expected, but got robot_type={meta.robot_type} instead of {robot_type}."
)
if not features_equal_for_merge(features, meta.features):
raise ValueError(
f"Same features is expected, but got features={meta.features} instead of {features}."
)
if not lenient:
raise ValueError(
f"Same features is expected, but got features={meta.features} instead of {features}."
)
# Union: add any features present in this dataset but not the first
for key, feat_def in meta.features.items():
if key not in features:
features[key] = feat_def
logging.info(f"Lenient merge: adding missing feature '{key}' from {meta.repo_id}")
return fps, robot_type, features
@@ -289,6 +298,7 @@ def aggregate_datasets(
chunk_size: int | None = None,
concatenate_videos: bool = True,
concatenate_data: bool = True,
lenient: bool = False,
):
"""Aggregates multiple LeRobot datasets into a single unified dataset.
@@ -325,8 +335,17 @@ def aggregate_datasets(
LeRobotDatasetMetadata(repo_id, root=root) for repo_id, root in zip(repo_ids, roots, strict=False)
]
)
fps, robot_type, _ = validate_all_metadata(all_metadata)
features = merge_video_feature_info_for_aggregate(all_metadata)
fps, robot_type, union_features = validate_all_metadata(all_metadata, lenient=lenient)
if lenient:
# Use union features as the base, then merge video encoder info on top
features = copy.deepcopy(union_features)
video_keys_for_merge = [k for k in features if features[k].get("dtype") == "video"]
merged_video_info = merge_video_feature_info_for_aggregate(all_metadata)
for vk in video_keys_for_merge:
if vk in merged_video_info:
features[vk] = merged_video_info[vk]
else:
features = merge_video_feature_info_for_aggregate(all_metadata)
video_keys = [key for key in features if features[key]["dtype"] == "video"]
dst_meta = LeRobotDatasetMetadata.create(
@@ -539,6 +558,31 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si
df = pd.read_parquet(src_path)
df = update_data_df(df, src_meta, dst_meta)
# Fill missing columns with default values (for lenient merge)
for col_name, feat_def in dst_meta.features.items():
if col_name in df.columns:
continue
if col_name in ("index", "episode_index", "task_index"):
continue
dtype = feat_def.get("dtype", "float32")
# Video/image features are stored as separate files, not in parquet
if dtype in ("video", "image"):
continue
n_rows = len(df)
if dtype == "language":
df[col_name] = [[] for _ in range(n_rows)]
elif dtype == "bool":
df[col_name] = False
elif dtype in ("float32", "float64"):
df[col_name] = 0.0
elif dtype in ("int32", "int64"):
df[col_name] = 0
elif dtype == "string":
df[col_name] = ""
else:
df[col_name] = 0.0
logging.info(f"Filled missing column '{col_name}' with default for {n_rows} rows")
# Write data and get the actual destination file it was written to
# This avoids duplicating the rotation logic here
data_idx, (dst_chunk, dst_file) = append_or_create_parquet_file(
+3
View File
@@ -274,6 +274,7 @@ def merge_datasets(
output_dir: str | Path | None = None,
concatenate_videos: bool = True,
concatenate_data: bool = True,
lenient: bool = False,
) -> LeRobotDataset:
"""Merge multiple LeRobotDatasets into a single dataset.
@@ -285,6 +286,7 @@ def merge_datasets(
output_dir: Root directory where the merged dataset will be stored. If not specified, defaults to $HF_LEROBOT_HOME/output_repo_id.
concatenate_videos: When False, keep one mp4 per source file instead of packing into shards.
concatenate_data: When False, keep one parquet per source file instead of packing into shards.
lenient: Allow merging datasets with different feature sets (union + fill defaults).
"""
if not datasets:
raise ValueError("No datasets to merge")
@@ -301,6 +303,7 @@ def merge_datasets(
aggr_root=output_dir,
concatenate_videos=concatenate_videos,
concatenate_data=concatenate_data,
lenient=lenient,
)
merged_dataset = LeRobotDataset(
+2 -2
View File
@@ -43,10 +43,10 @@ CORE_STYLES = {
# validation. Empty by default — populate from a downstream module that
# also extends ``PERSISTENT_STYLES`` or ``EVENT_ONLY_STYLES`` to declare
# the new style's column.
EXTENDED_STYLES: set[str] = set()
EXTENDED_STYLES: set[str] = {"advantage"}
STYLE_REGISTRY = CORE_STYLES | EXTENDED_STYLES
PERSISTENT_STYLES = {"subtask", "plan", "memory", "motion", "task_aug"}
PERSISTENT_STYLES = {"subtask", "plan", "memory", "motion", "task_aug", "advantage"}
EVENT_ONLY_STYLES = {"interjection", "vqa", "trace"}
# Styles whose ``content`` is grounded in a specific camera view. Rows of these
@@ -73,6 +73,19 @@ class MolmoAct2Config(PreTrainedConfig):
num_inference_steps: int | None = None
mask_action_dim_padding: bool = True
enable_inference_cuda_graph: bool = True
# Language conditioning (e.g. RECAP advantage). When set, RenderMessagesStep
# resolves language_persistent rows via the recipe YAML. Same mechanism as PI05.
recipe_path: str | None = None
# Inference-time advantage indicator (e.g. "Advantage: positive. ").
# Used during rollout when no language_persistent data is available.
# Placed after the user prompt, before action tokens.
advantage_prefix: str = ""
# Classifier-Free Guidance (CFG) scale for inference (RECAP Eq. 13).
# 1.0 = no guidance. >1.0 = dual-path: v = v_uncond + beta * (v_cond - v_uncond)
cfg_beta: float = 1.0
# MolmoAct2-local eval option. When enabled, stochastic continuous action
# generation uses a rollout-local generator derived from eval_seed.
per_episode_seed: bool = False
@@ -497,6 +497,56 @@ def _weighted_per_example(
return loss_sum * float(batch_size) / global_weight_sum
def _cat_action_contexts(cond_ctx, uncond_ctx):
"""Concatenate two ActionExpertContext objects along the batch dimension."""
from .molmoact2_hf_model.modeling_molmoact2 import ActionExpertContext
kv_contexts = []
for (k_c, v_c), (k_u, v_u) in zip(cond_ctx.kv_contexts, uncond_ctx.kv_contexts, strict=True):
kv_contexts.append((torch.cat([k_c, k_u], dim=0), torch.cat([v_c, v_u], dim=0)))
cross_mask = None
if cond_ctx.cross_mask is not None and uncond_ctx.cross_mask is not None:
cross_mask = torch.cat([cond_ctx.cross_mask, uncond_ctx.cross_mask], dim=0)
self_mask = None
if cond_ctx.self_mask is not None and uncond_ctx.self_mask is not None:
self_mask = torch.cat([cond_ctx.self_mask, uncond_ctx.self_mask], dim=0)
elif cond_ctx.self_mask is not None:
self_mask = cond_ctx.self_mask.repeat(2, *([1] * (cond_ctx.self_mask.ndim - 1)))
valid_action = None
if cond_ctx.valid_action is not None and uncond_ctx.valid_action is not None:
valid_action = torch.cat([cond_ctx.valid_action, uncond_ctx.valid_action], dim=0)
rope_cache = cond_ctx.rope_cache
return ActionExpertContext(
kv_contexts=kv_contexts,
cross_mask=cross_mask,
self_mask=self_mask,
valid_action=valid_action,
rope_cache=rope_cache,
)
def _clone_modulation_with_conditioning(modulation, batched_conditioning):
"""Create a modulation with doubled batch for batched CFG forward."""
from .molmoact2_hf_model.modeling_molmoact2 import ActionExpertStepModulation
batched_block_modulations = []
for block_mod in modulation.block_modulations:
batched_block_modulations.append(tuple(torch.cat([m, m], dim=0) for m in block_mod))
batched_final_modulation = tuple(torch.cat([m, m], dim=0) for m in modulation.final_modulation)
return ActionExpertStepModulation(
conditioning=batched_conditioning,
block_modulations=batched_block_modulations,
final_modulation=batched_final_modulation,
)
class MolmoAct2Policy(PreTrainedPolicy):
"""MolmoAct2 policy wrapping the vendored HF model for LeRobot.
@@ -1622,6 +1672,183 @@ class MolmoAct2Policy(PreTrainedPolicy):
metrics["loss"] = loss.detach().float().mean().item()
return loss, metrics
def _cfg_enabled_for_batch(self, batch: dict[str, Tensor]) -> bool:
"""Check if CFG should be used for this batch."""
return self.config.cfg_beta > 1.0 and batch.get("uncond_input_ids") is not None
def _uncond_model_inputs(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
"""Extract unconditional model inputs from the batch (prepared by processor)."""
compute_dtype = _torch_dtype(self.config.model_dtype)
uncond_inputs: dict[str, Tensor] = {}
for key in _MODEL_INPUT_KEYS:
uncond_key = f"uncond_{key}"
value = batch.get(uncond_key)
if value is not None:
uncond_inputs[key] = value.to(dtype=compute_dtype) if value.is_floating_point() else value
return uncond_inputs
def _generate_actions_with_cfg(
self,
*,
cond_model_inputs: dict[str, Tensor],
uncond_model_inputs: dict[str, Tensor],
action_dim_is_pad: Tensor | None,
num_steps: int | None,
generator: torch.Generator | None,
) -> Tensor:
"""CFG inference: dual VLM forward + batched flow denoising.
Caching strategy:
1. VLM backbone runs once per branch (cond + uncond) KV states cached.
2. Action expert context prepared once per branch from cached KV states.
3. Modulation cache (timestep embeddings) shared across branches.
4. Denoising loop: cond + uncond batched into a single action expert
forward per step (2x batch dim), then split and blended.
"""
backbone = self._backbone()
action_expert = self._action_expert()
# === VLM prefill (cached — runs once per branch) ===
cond_outputs = backbone(
**cond_model_inputs,
use_cache=True,
output_attentions=False,
output_hidden_states=False,
)
cond_encoder_kv_states = backbone._extract_kv_states(cond_outputs.past_key_values)
cond_encoder_attention_mask = self._encoder_attention_mask_for_action_expert(
input_ids=cond_model_inputs.get("input_ids"),
attention_mask=cond_model_inputs.get("attention_mask"),
)
cond_depth_gate, cond_depth_mask = backbone._depth_gate_from_condition(
input_ids=cond_model_inputs.get("input_ids"),
encoder_attention_mask=cond_encoder_attention_mask,
layer_kv_states=cond_encoder_kv_states,
)
cond_encoder_kv_states = backbone._apply_depth_gate_to_layer_kv_states(
cond_encoder_kv_states, cond_depth_mask, cond_depth_gate
)
uncond_outputs = backbone(
**uncond_model_inputs,
use_cache=True,
output_attentions=False,
output_hidden_states=False,
)
uncond_encoder_kv_states = backbone._extract_kv_states(uncond_outputs.past_key_values)
uncond_encoder_attention_mask = self._encoder_attention_mask_for_action_expert(
input_ids=uncond_model_inputs.get("input_ids"),
attention_mask=uncond_model_inputs.get("attention_mask"),
)
uncond_depth_gate, uncond_depth_mask = backbone._depth_gate_from_condition(
input_ids=uncond_model_inputs.get("input_ids"),
encoder_attention_mask=uncond_encoder_attention_mask,
layer_kv_states=uncond_encoder_kv_states,
)
uncond_encoder_kv_states = backbone._apply_depth_gate_to_layer_kv_states(
uncond_encoder_kv_states, uncond_depth_mask, uncond_depth_gate
)
# === Setup flow denoising ===
steps = int(num_steps or backbone.config.flow_matching_num_steps)
if steps <= 0:
raise ValueError(f"num_steps must be >= 1, got {steps}.")
source_tensor = cond_encoder_kv_states[0][0]
batch_size = int(source_tensor.shape[0])
device = source_tensor.device
trajectory = torch.randn(
batch_size,
self._generation_action_horizon(),
int(backbone.config.max_action_dim),
device=device,
dtype=torch.float32,
generator=generator,
)
if self.config.mask_action_dim_padding:
trajectory = _mask_action_dim_tensor(trajectory, action_dim_is_pad)
# === Prepare action contexts (cached — reused across all denoising steps) ===
cond_action_context = action_expert.prepare_context(
encoder_kv_states=cond_encoder_kv_states,
encoder_attention_mask=cond_encoder_attention_mask,
state_embeddings=None,
batch_size=batch_size,
seq_len=trajectory.shape[1],
device=device,
dtype=trajectory.dtype,
)
uncond_action_context = action_expert.prepare_context(
encoder_kv_states=uncond_encoder_kv_states,
encoder_attention_mask=uncond_encoder_attention_mask,
state_embeddings=None,
batch_size=batch_size,
seq_len=trajectory.shape[1],
device=device,
dtype=trajectory.dtype,
)
# Modulation cache shared between branches (timestep is prompt-independent)
flow_timesteps = [
torch.full((batch_size,), idx / steps, device=device, dtype=trajectory.dtype)
for idx in range(steps)
]
modulation_cache = action_expert.get_or_prepare_modulation_cache(
flow_timesteps,
cache_key=(steps, batch_size, device, trajectory.dtype),
)
# === Batched CFG denoising loop ===
# Instead of two sequential action expert forwards per step, we batch
# cond + uncond on the batch dimension for a single forward (2x batch).
# This maximizes GPU utilization (same pattern as PI05).
dt = 1.0 / steps
mask_enabled = self.config.mask_action_dim_padding
cfg_beta = self.config.cfg_beta
batched_action_dim_is_pad = (
torch.cat([action_dim_is_pad, action_dim_is_pad], dim=0)
if action_dim_is_pad is not None
else None
)
for idx in range(steps):
modulation = modulation_cache[idx]
# Duplicate trajectory and conditioning for batched forward
batched_trajectory = torch.cat([trajectory, trajectory], dim=0)
batched_conditioning = torch.cat([modulation.conditioning, modulation.conditioning], dim=0)
# Build batched context by concatenating cond + uncond contexts
batched_context = _cat_action_contexts(cond_action_context, uncond_action_context)
# Build batched modulation with doubled conditioning
batched_modulation = _clone_modulation_with_conditioning(modulation, batched_conditioning)
# Single batched forward through action expert
v_all = action_expert.forward_with_context(
batched_trajectory,
batched_conditioning,
context=batched_context,
modulation=batched_modulation,
)
if mask_enabled:
v_all = _mask_action_dim_tensor(v_all, batched_action_dim_is_pad)
# Split: first half = cond, second half = uncond
v_cond, v_uncond = v_all.chunk(2, dim=0)
# CFG interpolation: v = v_uncond + beta * (v_cond - v_uncond)
velocity = v_uncond + cfg_beta * (v_cond - v_uncond)
trajectory = trajectory + dt * velocity
if mask_enabled:
trajectory = _mask_action_dim_tensor(trajectory, action_dim_is_pad)
return trajectory
@torch.no_grad()
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
"""Generate an action chunk via continuous flow matching or discrete AR decoding."""
@@ -1657,6 +1884,15 @@ class MolmoAct2Policy(PreTrainedPolicy):
model_inputs=model_inputs,
action_dim=action_dim,
)
elif self._cfg_enabled_for_batch(batch):
uncond_model_inputs = self._uncond_model_inputs(batch)
actions = self._generate_actions_with_cfg(
cond_model_inputs=model_inputs,
uncond_model_inputs=uncond_model_inputs,
action_dim_is_pad=batch.get("action_dim_is_pad"),
num_steps=num_steps,
generator=generator,
)
elif self._rtc_enabled():
actions = self._generate_actions_from_inputs_with_rtc(
model_inputs=model_inputs,
@@ -359,6 +359,7 @@ def _build_robot_text(
add_setup_tokens: bool,
add_control_tokens: bool,
num_images: int,
advantage: str = "",
) -> str:
setup_text = _wrap_setup_text(setup_type, add_setup_tokens=add_setup_tokens)
control_text = _wrap_control_text(control_mode, add_control_tokens=add_control_tokens)
@@ -375,7 +376,10 @@ def _build_robot_text(
image_prefix = "<|image|>"
else:
image_prefix = "".join(f"Image {idx + 1}<|image|>" for idx in range(num_images))
return f"{image_prefix}<|im_start|>user\n{prompt}<|im_end|>\n<|im_start|>assistant\n{ACTION_OUTPUT_TOKEN}"
# Per RECAP paper (Section V-B): advantage indicator goes after context,
# before actions, so only action log-likelihoods are affected.
advantage_clause = f"Advantage: {advantage}. " if advantage else ""
return f"{image_prefix}<|im_start|>user\n{prompt}<|im_end|>\n<|im_start|>assistant\n{advantage_clause}{ACTION_OUTPUT_TOKEN}"
def _as_text_list(value: Any, batch_size: int) -> list[str]:
@@ -695,6 +699,39 @@ class MolmoAct2ClampNormalizedProcessorStep(ProcessorStep):
return features
@ProcessorStepRegistry.register(name="molmoact2_normalize_task")
@dataclass
class MolmoAct2NormalizeTaskStep(ProcessorStep):
"""Normalize the task text in complementary_data before recipe rendering.
Ensures ${task} in recipe templates gets the same normalized form that
MolmoAct2PackInputsProcessorStep would produce, so training prompts
match inference prompts.
"""
def __call__(self, transition: EnvTransition) -> EnvTransition:
complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA)
if not isinstance(complementary, dict):
return transition
task = complementary.get("task")
if task is None:
return transition
transition = transition.copy()
complementary = dict(complementary)
if isinstance(task, str):
complementary["task"] = _normalize_question_text(task)
elif isinstance(task, list):
complementary["task"] = [_normalize_question_text(t) for t in task]
transition[TransitionKey.COMPLEMENTARY_DATA] = complementary
return transition
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
@ProcessorStepRegistry.register(name="molmoact2_pack_inputs")
@dataclass
class MolmoAct2PackInputsProcessorStep(ProcessorStep):
@@ -715,6 +752,10 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
chunk_size: int = 30
max_action_dim: int = 32
env_action_dim: int | None = None
# RECAP: advantage indicator for inference (e.g. "Advantage: positive. ")
advantage_prefix: str = ""
# CFG scale for inference. >1.0 builds unconditional inputs for guidance.
cfg_beta: float = 1.0
def __post_init__(self) -> None:
require_package("transformers", extra="molmoact2")
@@ -757,6 +798,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
"chunk_size": self.chunk_size,
"max_action_dim": self.max_action_dim,
"env_action_dim": self.env_action_dim,
"advantage_prefix": self.advantage_prefix,
}
def _resolve_max_sequence_length(
@@ -919,8 +961,40 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
if task_source is None:
task_source = complementary.get("language_instruction")
tasks = _as_text_list(task_source, batch_size)
if self.normalize_language:
tasks = [_normalize_question_text(task) for task in tasks]
# Resolve the advantage indicator. Per RECAP paper (Section V-B), it goes
# after all context but before actions — handled by _build_robot_text.
# Source priority: recipe-rendered "advantage" key > config advantage_prefix.
advantages: list[str] = []
recipe_rendered = "base_task" in complementary
if recipe_rendered:
# Recipe rendered the task as "<task> Advantage: <value>".
# Extract the advantage value and restore the clean task.
clean_tasks: list[str] = []
for t in tasks:
if " Advantage: " in t:
split_idx = t.rindex(" Advantage: ")
clean_task = t[:split_idx]
adv = t[split_idx + len(" Advantage: ") :]
advantages.append(adv)
clean_tasks.append(clean_task)
else:
advantages.append("")
clean_tasks.append(t)
tasks = clean_tasks
else:
if self.normalize_language:
tasks = [_normalize_question_text(task) for task in tasks]
if self.advantage_prefix:
# Extract just the value from prefix like "Advantage: positive. "
prefix = self.advantage_prefix.strip()
if prefix.startswith("Advantage:"):
adv_val = prefix[len("Advantage:") :].strip().rstrip(".")
else:
adv_val = prefix
advantages = [adv_val] * batch_size
else:
advantages = [""] * batch_size
complementary["task"] = tasks
action_padded = None
@@ -953,6 +1027,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
add_setup_tokens=self.add_setup_tokens,
add_control_tokens=self.add_control_tokens,
num_images=len(images),
advantage=advantages[batch_idx],
)
prompt_texts.append(prompt)
if build_action_labels:
@@ -989,6 +1064,33 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
if build_action_labels:
inputs["labels"] = self._build_labels(inputs["input_ids"], inputs["attention_mask"])
# CFG: build unconditional inputs (no advantage) for inference-time guidance.
# Only produced when cfg_beta > 1.0 and we have advantage conditioning.
if self.cfg_beta > 1.0 and action is None and any(advantages):
uncond_prompt_texts: list[str] = []
for batch_idx in range(batch_size):
images = images_by_example[batch_idx]
discrete_state = _build_discrete_state_string(state_np[batch_idx], self.num_state_tokens)
uncond_prompt = _build_robot_text(
task=tasks[batch_idx],
discrete_state_string=discrete_state,
setup_type=self.setup_type,
control_mode=self.control_mode,
add_setup_tokens=self.add_setup_tokens,
add_control_tokens=self.add_control_tokens,
num_images=len(images),
advantage="",
)
uncond_prompt_texts.append(uncond_prompt)
uncond_inputs = self.processor(
text=uncond_prompt_texts, images=flat_images, return_tensors="pt", padding=True
)
complementary["uncond_input_ids"] = uncond_inputs["input_ids"]
complementary["uncond_attention_mask"] = uncond_inputs["attention_mask"]
for key in ("pixel_values", "image_token_pooling", "image_grids", "image_num_crops"):
if key in uncond_inputs:
complementary[f"uncond_{key}"] = uncond_inputs[key]
complementary.update(dict(inputs))
complementary["action_dim_is_pad"] = action_dim_is_pad
if action_horizon_is_pad is not None:
@@ -1164,28 +1266,49 @@ def make_molmoact2_pre_post_processors(
stats=masked_dataset_stats,
),
MolmoAct2ClampNormalizedProcessorStep(normalization_masks=normalization_masks),
MolmoAct2PackInputsProcessorStep(
checkpoint_path=config.checkpoint_path,
checkpoint_revision=config.checkpoint_revision,
checkpoint_force_download=config.checkpoint_force_download,
action_mode=config.action_mode,
discrete_action_tokenizer=config.discrete_action_tokenizer,
image_keys=image_keys,
allow_image_key_fallback=not bool(config.image_keys),
setup_type=setup_type,
control_mode=control_mode,
normalize_language=config.normalize_language,
add_setup_tokens=config.add_setup_tokens,
add_control_tokens=config.add_control_tokens,
num_state_tokens=config.num_state_tokens,
max_sequence_length=config.max_sequence_length,
chunk_size=chunk_size,
max_action_dim=config.expected_max_action_dim,
env_action_dim=env_action_dim,
),
DeviceProcessorStep(device=config.device),
]
# Insert language rendering steps when a recipe is configured (e.g. RECAP advantage)
if config.recipe_path is not None:
from lerobot.configs.recipe import load_recipe
from lerobot.processor.render_messages_processor import RenderMessagesStep
from lerobot.processor.rendered_messages_to_task import RenderedMessagesToTaskStep
recipe = load_recipe(config.recipe_path)
# Normalize task text before recipe uses ${task}, ensuring consistency
# between training (recipe-rendered) and inference (advantage_prefix).
if config.normalize_language:
input_steps.append(MolmoAct2NormalizeTaskStep())
input_steps.append(RenderMessagesStep(recipe=recipe))
input_steps.append(RenderedMessagesToTaskStep())
input_steps.extend(
[
MolmoAct2PackInputsProcessorStep(
checkpoint_path=config.checkpoint_path,
checkpoint_revision=config.checkpoint_revision,
checkpoint_force_download=config.checkpoint_force_download,
action_mode=config.action_mode,
discrete_action_tokenizer=config.discrete_action_tokenizer,
image_keys=image_keys,
allow_image_key_fallback=not bool(config.image_keys),
setup_type=setup_type,
control_mode=control_mode,
normalize_language=config.normalize_language,
add_setup_tokens=config.add_setup_tokens,
add_control_tokens=config.add_control_tokens,
num_state_tokens=config.num_state_tokens,
max_sequence_length=config.max_sequence_length,
chunk_size=chunk_size,
max_action_dim=config.expected_max_action_dim,
env_action_dim=env_action_dim,
advantage_prefix=config.advantage_prefix,
cfg_beta=config.cfg_beta,
),
DeviceProcessorStep(device=config.device),
]
)
output_steps: list[ProcessorStep] = [
MolmoAct2ClampActionProcessorStep(),
MolmoAct2MaskedUnnormalizerProcessorStep(
@@ -87,6 +87,17 @@ class PI05Config(PreTrainedConfig):
freeze_vision_encoder: bool = False # Freeze only the vision encoder
train_expert_only: bool = False # Freeze entire VLM, train only action expert and projections
# Language conditioning (e.g. RECAP advantage). When set, RenderMessagesStep
# is inserted into the preprocessor to resolve language_persistent rows via
# the recipe YAML before prompt construction.
recipe_path: str | None = None
# Classifier-Free Guidance (CFG) scale for inference (Eq. 13 in RECAP paper).
# 1.0 = no guidance (default). >1.0 enables dual-path denoising where:
# v = v_uncond + cfg_beta * (v_cond - v_uncond)
# VLM runs twice (cond + uncond prompts), action expert runs 2x per step.
cfg_beta: float = 1.0
# Optimizer settings: see openpi `AdamW`
optimizer_lr: float = 2.5e-5 # see openpi `CosineDecaySchedule: peak_lr`
optimizer_betas: tuple[float, float] = (0.9, 0.95)
+141 -2
View File
@@ -52,6 +52,8 @@ from lerobot.utils.constants import (
ACTION,
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
OBS_LANGUAGE_UNCOND_TOKENS,
OPENPI_ATTENTION_MASK_VALUE,
)
@@ -148,6 +150,20 @@ def clone_past_key_values(past_key_values):
)
def cat_past_key_values(kv_a, kv_b):
"""Concatenate two DynamicCaches along the batch dimension for batched CFG."""
return DynamicCache(
tuple(
(
torch.cat([ka, kb], dim=0),
torch.cat([va, vb], dim=0),
sw_a,
)
for (ka, va, sw_a), (kb, vb, _sw_b) in zip(kv_a, kv_b, strict=True)
)
)
def pad_vector(vector, new_dim):
"""Pad the last dimension of a vector to new_dim with zeros.
@@ -797,9 +813,17 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
masks,
noise=None,
num_steps=None,
uncond_tokens=None,
uncond_masks=None,
**kwargs: Unpack[ActionSelectKwargs],
) -> Tensor:
"""Do a full inference forward and compute the action."""
"""Do a full inference forward and compute the action.
When cfg_beta > 1.0 and uncond_tokens/uncond_masks are provided, performs
Classifier-Free Guidance: VLM runs twice (conditioned + unconditional), action
expert runs twice per denoising step, and velocities are interpolated via
v = v_uncond + cfg_beta * (v_cond - v_uncond).
"""
if num_steps is None:
num_steps = self.config.num_inference_steps
@@ -815,6 +839,9 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
) # Use config max_action_dim for internal processing
noise = self.sample_noise(actions_shape, device)
cfg_enabled = self.config.cfg_beta > 1.0 and uncond_tokens is not None and uncond_masks is not None
# Prefill VLM for conditioned prompt
prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(images, img_masks, tokens, masks)
prefix_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks)
prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1
@@ -830,6 +857,23 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
use_cache=True,
)
# Prefill VLM for unconditional prompt (CFG)
if cfg_enabled:
uncond_prefix_embs, uncond_prefix_pad_masks, uncond_prefix_att_masks = self.embed_prefix(
images, img_masks, uncond_tokens, uncond_masks
)
uncond_prefix_att_2d_masks = make_att_2d_masks(uncond_prefix_pad_masks, uncond_prefix_att_masks)
uncond_prefix_position_ids = torch.cumsum(uncond_prefix_pad_masks, dim=1) - 1
uncond_prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(uncond_prefix_att_2d_masks)
_, uncond_past_key_values = self.paligemma_with_expert.forward(
attention_mask=uncond_prefix_att_2d_masks_4d,
position_ids=uncond_prefix_position_ids,
past_key_values=None,
inputs_embeds=[uncond_prefix_embs, None],
use_cache=True,
)
dt = -1.0 / num_steps
x_t = noise
@@ -838,6 +882,15 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
if cfg_enabled:
return self.denoise_step_cfg_batched(
cond_prefix_pad_masks=prefix_pad_masks,
cond_past_key_values=past_key_values,
uncond_prefix_pad_masks=uncond_prefix_pad_masks,
uncond_past_key_values=uncond_past_key_values,
x_t=input_x_t,
timestep=current_timestep,
)
return self.denoise_step(
prefix_pad_masks=prefix_pad_masks,
past_key_values=past_key_values,
@@ -907,6 +960,80 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
suffix_out = suffix_out.to(dtype=torch.float32)
return self.action_out_proj(suffix_out)
def denoise_step_cfg_batched(
self,
cond_prefix_pad_masks,
cond_past_key_values,
uncond_prefix_pad_masks,
uncond_past_key_values,
x_t,
timestep,
):
"""Batched CFG denoising: runs cond + uncond in a single forward pass.
Concatenates cond and uncond inputs along the batch dimension, runs one
action expert forward (2x batch), then splits and applies CFG interpolation.
This is ~1.5x faster than two sequential denoise_step calls due to better
GPU utilization (inspired by Qwen2.5-Omni DiT / diffusers batched CFG).
"""
# Embed suffix once (same x_t and timestep for both branches)
suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, timestep)
bsize = cond_prefix_pad_masks.shape[0]
suffix_len = suffix_pad_masks.shape[1]
cond_prefix_len = cond_prefix_pad_masks.shape[1]
uncond_prefix_len = uncond_prefix_pad_masks.shape[1]
# Build attention masks for cond branch
cond_prefix_2d = cond_prefix_pad_masks[:, None, :].expand(bsize, suffix_len, cond_prefix_len)
cond_suffix_att_2d = make_att_2d_masks(suffix_pad_masks, suffix_att_masks)
cond_full_att = torch.cat([cond_prefix_2d, cond_suffix_att_2d], dim=2)
cond_prefix_offsets = torch.sum(cond_prefix_pad_masks, dim=-1)[:, None]
cond_position_ids = cond_prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
# Build attention masks for uncond branch
uncond_prefix_2d = uncond_prefix_pad_masks[:, None, :].expand(bsize, suffix_len, uncond_prefix_len)
uncond_suffix_att_2d = make_att_2d_masks(suffix_pad_masks, suffix_att_masks)
uncond_full_att = torch.cat([uncond_prefix_2d, uncond_suffix_att_2d], dim=2)
uncond_prefix_offsets = torch.sum(uncond_prefix_pad_masks, dim=-1)[:, None]
uncond_position_ids = uncond_prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
# Concatenate on batch dim: [cond_batch; uncond_batch]
batched_full_att = torch.cat([cond_full_att, uncond_full_att], dim=0)
batched_full_att_4d = self._prepare_attention_masks_4d(batched_full_att)
batched_position_ids = torch.cat([cond_position_ids, uncond_position_ids], dim=0)
batched_suffix_embs = torch.cat([suffix_embs, suffix_embs], dim=0)
batched_adarms_cond = torch.cat([adarms_cond, adarms_cond], dim=0)
# Concatenate KV caches on batch dim
batched_past_kv = cat_past_key_values(
clone_past_key_values(cond_past_key_values),
clone_past_key_values(uncond_past_key_values),
)
self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001
# Single forward pass for both branches
outputs_embeds, _ = self.paligemma_with_expert.forward(
attention_mask=batched_full_att_4d,
position_ids=batched_position_ids,
past_key_values=batched_past_kv,
inputs_embeds=[None, batched_suffix_embs],
use_cache=False,
adarms_cond=[None, batched_adarms_cond],
)
suffix_out = outputs_embeds[1]
suffix_out = suffix_out[:, -self.config.chunk_size :]
suffix_out = suffix_out.to(dtype=torch.float32)
v_all = self.action_out_proj(suffix_out)
# Split: first half = cond, second half = uncond
v_cond, v_uncond = v_all.chunk(2, dim=0)
# CFG interpolation: v = v_uncond + beta * (v_cond - v_uncond)
return v_uncond + self.config.cfg_beta * (v_cond - v_uncond)
class PI05Policy(PreTrainedPolicy):
"""PI05 Policy for LeRobot."""
@@ -1243,8 +1370,20 @@ class PI05Policy(PreTrainedPolicy):
images, img_masks = self._preprocess_images(batch)
tokens, masks = batch[f"{OBS_LANGUAGE_TOKENS}"], batch[f"{OBS_LANGUAGE_ATTENTION_MASK}"]
# CFG: pass unconditional tokens if available
uncond_tokens = batch.get(f"{OBS_LANGUAGE_UNCOND_TOKENS}")
uncond_masks = batch.get(f"{OBS_LANGUAGE_UNCOND_ATTENTION_MASK}")
# Sample actions using the model (pass through RTC kwargs, no separate state needed for PI05)
actions = self.model.sample_actions(images, img_masks, tokens, masks, **kwargs)
actions = self.model.sample_actions(
images,
img_masks,
tokens,
masks,
uncond_tokens=uncond_tokens,
uncond_masks=uncond_masks,
**kwargs,
)
# Unpad actions to actual action dimension
original_action_dim = self.config.output_features[ACTION].shape[0]
+69 -15
View File
@@ -40,6 +40,8 @@ from lerobot.processor import (
)
from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import (
OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
OBS_LANGUAGE_UNCOND_TOKENS,
OBS_STATE,
POLICY_POSTPROCESSOR_DEFAULT_NAME,
POLICY_PREPROCESSOR_DEFAULT_NAME,
@@ -57,6 +59,7 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
max_state_dim: int = 32
task_key: str = "task"
cfg_enabled: bool = False
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
@@ -84,8 +87,25 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
full_prompts.append(full_prompt)
transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts
# Normalize state to [-1, 1] range if needed (assuming it's already normalized by normalizer processor step!!)
# Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`)
# Build unconditional prompts for CFG (same state but original task without advantage)
if self.cfg_enabled:
base_tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get("base_task")
if base_tasks is None:
base_tasks = tasks
if isinstance(base_tasks, str):
base_tasks = [base_tasks] * len(tasks)
uncond_prompts = []
for i, base_task in enumerate(base_tasks):
cleaned_text = base_task.strip().replace("_", " ").replace("\n", " ")
state_str = " ".join(map(str, discretized_states[i]))
uncond_prompt = f"Task: {cleaned_text}, State: {state_str};\nAction: "
uncond_prompts.append(uncond_prompt)
transition[TransitionKey.COMPLEMENTARY_DATA]["uncond_task"] = uncond_prompts
return transition
def transform_features(
@@ -111,9 +131,10 @@ def make_pi05_pre_post_processors(
1. Renaming features to match pretrained configurations.
2. Normalizing input and output features based on dataset statistics.
3. Adding a batch dimension.
4. Appending a newline character to the task description for tokenizer compatibility.
5. Tokenizing the text prompt using the PaliGemma tokenizer.
6. Moving all data to the specified device.
4. (Optional) Rendering language annotations via recipe YAML.
5. (Optional) Flattening rendered messages into the task string.
6. Tokenizing the text prompt using the PaliGemma tokenizer.
7. Moving all data to the specified device.
The post-processing pipeline handles the model's output by:
1. Moving data to the CPU.
@@ -122,8 +143,6 @@ def make_pi05_pre_post_processors(
Args:
config: The configuration object for the PI0 policy.
dataset_stats: A dictionary of statistics for normalization.
preprocessor_kwargs: Additional arguments for the pre-processor pipeline.
postprocessor_kwargs: Additional arguments for the post-processor pipeline.
Returns:
A tuple containing the configured pre-processor and post-processor pipelines.
@@ -147,16 +166,51 @@ def make_pi05_pre_post_processors(
norm_map=config.normalization_mapping,
stats=dataset_stats,
),
Pi05PrepareStateTokenizerProcessorStep(max_state_dim=config.max_state_dim),
TokenizerProcessorStep(
tokenizer_name="google/paligemma-3b-pt-224",
max_length=config.tokenizer_max_length,
padding_side="right",
padding="max_length",
),
DeviceProcessorStep(device=config.device),
]
# Insert language rendering steps when a recipe is configured (e.g. RECAP advantage)
if config.recipe_path is not None:
from lerobot.configs.recipe import load_recipe
from lerobot.processor.render_messages_processor import RenderMessagesStep
from lerobot.processor.rendered_messages_to_task import RenderedMessagesToTaskStep
recipe = load_recipe(config.recipe_path)
input_steps.append(RenderMessagesStep(recipe=recipe))
input_steps.append(RenderedMessagesToTaskStep())
cfg_enabled = config.cfg_beta > 1.0
input_steps.extend(
[
Pi05PrepareStateTokenizerProcessorStep(
max_state_dim=config.max_state_dim,
cfg_enabled=cfg_enabled,
),
TokenizerProcessorStep(
tokenizer_name="google/paligemma-3b-pt-224",
max_length=config.tokenizer_max_length,
padding_side="right",
padding="max_length",
),
]
)
# Add unconditional prompt tokenizer for CFG inference
if cfg_enabled:
input_steps.append(
TokenizerProcessorStep(
tokenizer_name="google/paligemma-3b-pt-224",
max_length=config.tokenizer_max_length,
padding_side="right",
padding="max_length",
task_key="uncond_task",
output_tokens_key=OBS_LANGUAGE_UNCOND_TOKENS,
output_mask_key=OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
)
)
input_steps.append(DeviceProcessorStep(device=config.device))
output_steps: list[ProcessorStep] = [
UnnormalizerProcessorStep(
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
+4
View File
@@ -164,6 +164,10 @@ _COMPLEMENTARY_KEYS = (
"messages",
"message_streams",
"target_message_indices",
"mc_return",
"is_terminal",
"next.success",
"intervention",
)
+14 -2
View File
@@ -1054,8 +1054,20 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
try:
step_class = ProcessorStepRegistry.get(step_entry["registry_name"])
return step_class, step_entry["registry_name"]
except KeyError as e:
raise ImportError(f"Failed to load processor step from registry. {str(e)}") from e
except KeyError:
registry_name = step_entry["registry_name"]
module_path = f"lerobot.processor.{registry_name}"
try:
importlib.import_module(module_path)
step_class = ProcessorStepRegistry.get(registry_name)
return step_class, registry_name
except (ImportError, ModuleNotFoundError, KeyError):
raise ImportError(
f"Failed to load processor step from registry. "
f"Processor step '{registry_name}' not found in registry. "
f"Available steps: {list(ProcessorStepRegistry._registry.keys())}. "
f"Make sure the step is registered using @ProcessorStepRegistry.register()"
) from None
else:
# Fallback to dynamic import using the full class path
full_class_path = step_entry["class"]
@@ -40,11 +40,14 @@ class RenderMessagesStep(ProcessorStep):
``message_streams`` / ``target_message_indices`` keys.
"""
recipe: TrainingRecipe
recipe: TrainingRecipe | None = None
dataset_ctx: Any | None = None
def __call__(self, transition: EnvTransition) -> EnvTransition | None:
"""Render messages for a single transition; return ``None`` to drop it."""
if self.recipe is None:
return transition
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
persistent = complementary_data.get(LANGUAGE_PERSISTENT) or []
events = complementary_data.get(LANGUAGE_EVENTS) or []
@@ -0,0 +1,86 @@
#!/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.
# 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.
"""Adapter step that flattens rendered chat messages back into a task string.
Bridges RenderMessagesStep (which outputs structured messages) to policies
that expect a plain task string in complementary_data["task"] (e.g. PI05).
"""
from __future__ import annotations
from lerobot.configs import PipelineFeatureType, PolicyFeature
from .pipeline import ComplementaryDataProcessorStep, ProcessorStepRegistry
@ProcessorStepRegistry.register(name="rendered_messages_to_task")
class RenderedMessagesToTaskStep(ComplementaryDataProcessorStep):
"""Extract user-role message content from rendered messages into the task string.
After RenderMessagesStep renders a recipe into structured messages, this
step extracts content from all user-role messages, joins them, and writes
the result to complementary_data["task"]. This allows downstream steps
(like Pi05PrepareStateTokenizerProcessorStep) to consume the
advantage-conditioned prompt without modification.
No-ops when the "messages" key is absent (backward compatible with
pipelines that don't use language annotations).
"""
def complementary_data(self, complementary_data: dict) -> dict:
messages = complementary_data.get("messages")
if messages is None:
return complementary_data
user_parts = []
for msg in messages:
if msg.get("role") == "user":
content = msg.get("content", "")
if isinstance(content, str) and content:
user_parts.append(content)
elif isinstance(content, list):
# HF multimodal blocks: extract text blocks
for block in content:
if isinstance(block, dict) and block.get("type") == "text":
text = block.get("text", "")
if text:
user_parts.append(text)
new_complementary_data = dict(complementary_data)
if user_parts:
task = complementary_data.get("task")
# Preserve the original task for CFG unconditional prompt
new_complementary_data["base_task"] = task
# Wrap in list if the original task was a list (batched)
joined = "\n".join(user_parts)
if isinstance(task, list):
new_complementary_data["task"] = [joined] * len(task)
else:
new_complementary_data["task"] = joined
# Remove consumed rendering outputs
new_complementary_data.pop("messages", None)
new_complementary_data.pop("message_streams", None)
new_complementary_data.pop("target_message_indices", None)
return new_complementary_data
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
+8 -6
View File
@@ -81,6 +81,8 @@ class TokenizerProcessorStep(ObservationProcessorStep):
padding_side: str = "right"
padding: str = "max_length"
truncation: bool = True
output_tokens_key: str = OBS_LANGUAGE_TOKENS
output_mask_key: str = OBS_LANGUAGE_ATTENTION_MASK
# Internal tokenizer instance (not part of the config)
input_tokenizer: Any = field(default=None, init=False, repr=False)
@@ -201,8 +203,8 @@ class TokenizerProcessorStep(ObservationProcessorStep):
new_observation = dict(observation)
# Add tokenized data to the observation
new_observation[OBS_LANGUAGE_TOKENS] = tokenized_prompt["input_ids"]
new_observation[OBS_LANGUAGE_ATTENTION_MASK] = tokenized_prompt["attention_mask"].to(dtype=torch.bool)
new_observation[self.output_tokens_key] = tokenized_prompt["input_ids"]
new_observation[self.output_mask_key] = tokenized_prompt["attention_mask"].to(dtype=torch.bool)
# Tokenize subtask if available
subtask = self.get_subtask(self.transition)
@@ -309,14 +311,14 @@ class TokenizerProcessorStep(ObservationProcessorStep):
The updated dictionary of policy features.
"""
# Add a feature for the token IDs if it doesn't already exist
if OBS_LANGUAGE_TOKENS not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][OBS_LANGUAGE_TOKENS] = PolicyFeature(
if self.output_tokens_key not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][self.output_tokens_key] = PolicyFeature(
type=FeatureType.LANGUAGE, shape=(self.max_length,)
)
# Add a feature for the attention mask if it doesn't already exist
if OBS_LANGUAGE_ATTENTION_MASK not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][OBS_LANGUAGE_ATTENTION_MASK] = PolicyFeature(
if self.output_mask_key not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][self.output_mask_key] = PolicyFeature(
type=FeatureType.LANGUAGE, shape=(self.max_length,)
)
+4
View File
@@ -13,6 +13,9 @@
# limitations under the License.
from .classifier.configuration_classifier import RewardClassifierConfig as RewardClassifierConfig
from .distributional_value_function.configuration_distributional_value_function import (
DistributionalVFConfig as DistributionalVFConfig,
)
from .factory import (
get_reward_model_class as get_reward_model_class,
make_reward_model as make_reward_model,
@@ -26,6 +29,7 @@ from .topreward.configuration_topreward import TOPRewardConfig as TOPRewardConfi
__all__ = [
# Configuration classes
"DistributionalVFConfig",
"RewardClassifierConfig",
"RobometerConfig",
"SARMConfig",
@@ -0,0 +1,23 @@
# 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 .configuration_distributional_value_function import DistributionalVFConfig
from .modeling_distributional_value_function import DistributionalVFRewardModel
from .processor_distributional_value_function import make_distributional_vf_pre_post_processors
__all__ = [
"DistributionalVFConfig",
"DistributionalVFRewardModel",
"make_distributional_vf_pre_post_processors",
]
@@ -0,0 +1,112 @@
# 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.
"""Configuration for RECAP's distributional value function.
Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06
Distributional value function V^{pi_ref}(o_t, l) (Section IV-A).
Architecture (~670M params):
Vision: SigLIP2-so400m 27 layers, 1152-dim, 256 patches/image
LM: Gemma3-270M 18 layers, 640-dim
Proj: Linear(1152, 640) fresh init
Head: [CLS] Linear(640320) LN GELU Dropout Linear(320201)
Inputs: multi-camera images (3 x 256 patches) + ``"Task: {task}."`` prompt
Targets: MC returns in [-1, 0], cross-entropy on HL-Gauss (default) or Dirac delta
Init: SigLIP2 + Gemma3 from pretrained HF checkpoints; head normal_(std=0.02)
"""
from dataclasses import dataclass, field
from lerobot.configs import FeatureType, NormalizationMode
from lerobot.configs.rewards import RewardModelConfig
from lerobot.optim import AdamWConfig, CosineDecayWithWarmupSchedulerConfig
@RewardModelConfig.register_subclass("distributional_value_function")
@dataclass
class DistributionalVFConfig(RewardModelConfig):
"""Configuration for RECAP's distributional value function.
Predicts V^{pi_ref}(o_t, l) as a categorical distribution over B=201 bins in [-1, 0].
Trained with cross-entropy on HL-Gauss soft targets (default) or Dirac delta (C51),
with optional one-hot targets for terminal states.
Architecture: monolithic VLM SigLIP2-so400m (vision) + Gemma3-270M (language),
bidirectional prefix attention, one-way [CLS] readout, 2-layer MLP value head.
"""
# Backbone pretrained paths
siglip_path: str = "google/siglip2-so400m-patch14-224"
gemma3_path: str = "google/gemma-3-270m"
# Distributional head
num_value_bins: int = 201
value_support_min: float = -1.0
value_support_max: float = 0.0
hl_gauss_sigma_ratio: float = 5.0
# Target distribution method: "hl_gauss" (default, soft) or "dirac_delta" (C51, hard)
target_method: str = "hl_gauss"
# Whether to use one-hot targets for terminal states (exact return, no smoothing).
use_one_hot_terminal: bool = True
# Image
image_resolution: tuple[int, int] = (224, 224)
# Tokenizer (uses Gemma3's tokenizer)
tokenizer_max_length: int = 200
# Training controls
value_dropout: float = 0.0
freeze_vision_encoder: bool = False
freeze_language_model: bool = False
stop_gradient_to_vlm: bool = False
vision_encoder_lr_multiplier: float = 0.5
# Readout: "mean_pool" (average all tokens) or "last_token" (causal LM last position)
readout: str = "mean_pool"
# Normalization
normalization_mapping: dict[str, NormalizationMode] = field(
default_factory=lambda: {
"VISUAL": NormalizationMode.IDENTITY,
}
)
def get_optimizer_preset(self) -> AdamWConfig:
return AdamWConfig(
lr=5e-5,
weight_decay=1e-10,
grad_clip_norm=1.0,
)
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
return CosineDecayWithWarmupSchedulerConfig(
num_warmup_steps=500,
num_decay_steps=40000,
peak_lr=5e-5,
decay_lr=5e-5,
)
def validate_features(self) -> None:
if not self.input_features:
return
has_image = any(ft.type == FeatureType.VISUAL for ft in self.input_features.values())
if not has_image:
raise ValueError("DistributionalVFConfig requires at least one VISUAL input feature.")
@@ -0,0 +1,500 @@
# 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.
"""Modeling for RECAP's distributional value function.
Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06
Implements the distributional value function V^{pi_ref}(o_t, l) from Section IV-A.
Architecture: the paper uses a 670M-parameter Gemma 3 VLM (Figure 3)
SigLIP2-so400m (27 layers, 1152-dim) + Gemma3-270M (18 layers, 640-dim),
with a [CLS] token readout predicting a categorical distribution over
B=201 discrete value bins in [-1, 0]. This implementation uses a 2-layer
MLP value head (LinearLNGELUDropoutLinear) inspired by Robometer
(Chen et al., 2025).
"""
from __future__ import annotations
import math
from typing import TYPE_CHECKING, Any
import torch
import torch.nn.functional as F # noqa: N812
from torch import Tensor, nn
from lerobot.configs.types import FeatureType
from lerobot.rewards.pretrained import PreTrainedRewardModel
from lerobot.utils.constants import (
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
OPENPI_ATTENTION_MASK_VALUE as _ATTENTION_MASK_VALUE,
)
from lerobot.utils.import_utils import _transformers_available, require_package
from .configuration_distributional_value_function import DistributionalVFConfig
from .processor_distributional_value_function import IMAGE_MASK_SUFFIX
if TYPE_CHECKING or _transformers_available:
from transformers import Gemma3ForCausalLM, SiglipVisionModel
else:
Gemma3ForCausalLM = None # type: ignore[assignment]
SiglipVisionModel = None # type: ignore[assignment]
class ValueHead(nn.Module):
"""Categorical value projection: hidden state → bin logits.
2-layer MLP: Linear LayerNorm GELU Dropout Linear.
Also holds the ``bin_centers`` buffer used to compute E[V] = Σ p_i · c_i.
"""
def __init__(
self,
hidden_size: int,
num_bins: int,
v_min: float,
v_max: float,
dropout: float = 0.0,
):
super().__init__()
self.hidden_size = hidden_size
self.num_bins = num_bins
self.mlp = nn.Sequential(
nn.Linear(hidden_size, hidden_size // 2),
nn.LayerNorm(hidden_size // 2),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_size // 2, num_bins),
)
self.register_buffer("bin_centers", torch.linspace(v_min, v_max, num_bins), persistent=False)
def forward(self, hidden_states: Tensor) -> Tensor:
"""Project hidden state to value logits. Returns [B, num_bins]."""
hidden_states = hidden_states.to(self.mlp[0].weight.dtype)
return self.mlp(hidden_states)
class DistributionalVFRewardModel(PreTrainedRewardModel):
"""Distributional value function model for RECAP.
Predicts V^{pi_ref}(o_t, l) as a categorical distribution over B bins (default 201).
Trained with cross-entropy on HL-Gauss or Dirac delta targets centered on
per-task normalized Monte Carlo returns.
Architecture: SigLIP2-so400m + Linear(1152640) + Gemma3-270M.
Multi-camera images are encoded by SigLIP2 (256 patches each), projected to
Gemma3's hidden dim, concatenated with tokenized language, and processed by
all 18 Gemma3 transformer layers.
Mean-pooled last-layer hidden states are read out through a 2-layer MLP value head.
"""
name = "distributional_value_function"
config_class = DistributionalVFConfig
def __init__(self, config: DistributionalVFConfig, **kwargs) -> None:
require_package("transformers", extra="recap")
super().__init__(config)
self.config = config
self.vision_encoder = SiglipVisionModel.from_pretrained(config.siglip_path)
siglip_hidden = self.vision_encoder.config.hidden_size # 1152
self.gemma3 = Gemma3ForCausalLM.from_pretrained(config.gemma3_path)
self.gemma3_hidden = self.gemma3.config.hidden_size # 640
# Fresh image projection: SigLIP2 1152-dim → Gemma3 640-dim
self.image_proj = nn.Linear(siglip_hidden, self.gemma3_hidden, bias=True)
nn.init.normal_(self.image_proj.weight, std=0.02)
nn.init.zeros_(self.image_proj.bias)
# Value head: last-token hidden state → MLP → num_bins logits
self.value_head = ValueHead(
hidden_size=self.gemma3_hidden,
num_bins=config.num_value_bins,
v_min=config.value_support_min,
v_max=config.value_support_max,
dropout=config.value_dropout,
)
# HL-Gauss sigma for soft targets
bin_width = (config.value_support_max - config.value_support_min) / (config.num_value_bins - 1)
self.hl_gauss_sigma = float(config.hl_gauss_sigma_ratio * bin_width)
# Apply freezing
self._set_requires_grad()
def _set_requires_grad(self) -> None:
if self.config.freeze_vision_encoder:
for param in self.vision_encoder.parameters():
param.requires_grad = False
self.vision_encoder.eval()
if self.config.freeze_language_model:
for param in self.gemma3.parameters():
param.requires_grad = False
self.gemma3.eval()
def train(self, mode: bool = True):
super().train(mode)
if self.config.freeze_vision_encoder:
self.vision_encoder.eval()
if self.config.freeze_language_model:
self.gemma3.eval()
return self
def get_optim_params(self) -> list[dict]:
"""Optimizer param groups with per-component learning rates."""
vision_params = []
other_params = []
for name, param in self.named_parameters():
if not param.requires_grad:
continue
if name.startswith("vision_encoder"):
vision_params.append(param)
else:
other_params.append(param)
base_lr = self.config.get_optimizer_preset().lr
return [
{"params": other_params},
{"params": vision_params, "lr": base_lr * self.config.vision_encoder_lr_multiplier},
]
def embed_image(self, image: Tensor) -> Tensor:
"""Embed images: SigLIP2 → projection → [B, num_patches, gemma3_hidden].
Args:
image: [batch_size, channels, height, width] preprocessed image in [-1, 1].
Returns:
[B, 256, gemma3_hidden] projected image features.
"""
if image.dtype != torch.float32:
image = image.to(torch.float32)
feats = self.vision_encoder(pixel_values=image).last_hidden_state
return self.image_proj(feats)
def embed_text(self, token_ids: Tensor) -> Tensor:
"""Embed text using Gemma3's embedding table (includes sqrt(d) scaling).
Args:
token_ids: [B, seq_len] integer token IDs.
Returns:
[B, seq_len, gemma3_hidden] text embeddings.
"""
return self.gemma3.model.embed_tokens(token_ids)
def embed_prefix(
self,
images: list[Tensor],
img_masks: list[Tensor],
text_embeddings: Tensor,
text_padding_mask: Tensor,
) -> tuple[Tensor, Tensor]:
"""Build prefix: [img1_patches, img2_patches, ..., lang_tokens].
All prefix tokens use bidirectional attention (att_mask=0).
Returns:
embs: [B, total_prefix_len, hidden_dim]
pad_masks: [B, total_prefix_len] boolean
"""
embs: list[Tensor] = []
pad_masks: list[Tensor] = []
for img, img_mask in zip(images, img_masks, strict=True):
img_emb = self.embed_image(img)
bsize, num_patches = img_emb.shape[:2]
embs.append(img_emb)
pad_masks.append(img_mask[:, None].expand(bsize, num_patches))
embs.append(text_embeddings)
pad_masks.append(text_padding_mask)
return torch.cat(embs, dim=1), torch.cat(pad_masks, dim=1)
def hl_gauss_target(self, target_value: Tensor) -> Tensor:
"""HL-Gauss soft target distribution.
Places a Gaussian N(target, sigma^2) over the bin support and computes
per-bin probabilities as CDF differences at bin edges, normalized to sum to 1.
Reference: Farebrother et al. 2024, "Stop Regressing: Training Value
Functions via Classification for Scalable Deep RL", Section 3.1.
arXiv:2403.03950
Args:
target_value: [batch_size] or [batch_size, 1] target values.
Returns:
[batch_size, num_value_bins] target probability distribution.
"""
if target_value.ndim == 2:
target_value = target_value.squeeze(-1)
target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
# Bin edges: half a bin-width outside the first/last center
bin_width = (self.config.value_support_max - self.config.value_support_min) / (
self.config.num_value_bins - 1
)
support_edges = torch.linspace(
self.config.value_support_min - bin_width / 2,
self.config.value_support_max + bin_width / 2,
self.config.num_value_bins + 1,
device=target_value.device,
dtype=target_value.dtype,
)
# CDF of N(target, sigma^2) evaluated at each edge
cdf_at_edges = 0.5 * (
1.0
+ torch.erf(
(support_edges.unsqueeze(0) - target_value.unsqueeze(-1))
/ (self.hl_gauss_sigma * math.sqrt(2))
)
) # [batch_size, num_bins + 1]
# Normalize: z = cdf(max_edge) - cdf(min_edge)
normalization_constant = (cdf_at_edges[:, -1] - cdf_at_edges[:, 0]).unsqueeze(-1).clamp(min=1e-10)
# Bin probabilities = differences of consecutive CDF values, normalized
bin_probabilities = (cdf_at_edges[:, 1:] - cdf_at_edges[:, :-1]) / normalization_constant
return bin_probabilities
def dirac_delta_target(self, target_value: Tensor) -> Tensor:
"""Dirac delta (C51) projection: split probability between two nearest bins.
Standard distributional RL projection from Bellemare et al. 2017.
"A Distributional Perspective on Reinforcement Learning"
arXiv:1707.06887
Args:
target_value: [batch_size] or [batch_size, 1] target values.
Returns:
[batch_size, num_value_bins] target probability distribution.
"""
if target_value.ndim == 2:
target_value = target_value.squeeze(-1)
target_value = target_value.clamp(self.config.value_support_min, self.config.value_support_max)
target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
bin_width = self.value_head.bin_centers[1] - self.value_head.bin_centers[0]
normalized_position = (target_value - self.config.value_support_min) / bin_width
lower_bin_idx = normalized_position.floor().long().clamp(0, self.config.num_value_bins - 1)
upper_bin_idx = normalized_position.ceil().long().clamp(0, self.config.num_value_bins - 1)
weight_upper = normalized_position - lower_bin_idx.float()
weight_lower = upper_bin_idx.float() - normalized_position
same_bin = lower_bin_idx == upper_bin_idx
weight_upper = torch.where(same_bin, torch.zeros_like(weight_upper), weight_upper)
weight_lower = torch.where(same_bin, torch.ones_like(weight_lower), weight_lower)
batch_size = target_value.shape[0]
target_distribution = torch.zeros(batch_size, self.config.num_value_bins, device=target_value.device)
batch_indices = torch.arange(batch_size, device=target_value.device)
target_distribution[batch_indices, lower_bin_idx] += weight_lower
target_distribution[batch_indices, upper_bin_idx] += weight_upper
return target_distribution
def one_hot_target(self, target_value: Tensor) -> Tensor:
"""One-hot target for terminal states (exact return, no smoothing).
Args:
target_value: [batch_size] or [batch_size, 1] target values.
Returns:
[batch_size, num_value_bins] one-hot distribution at the nearest bin.
"""
if target_value.ndim == 2:
target_value = target_value.squeeze(-1)
target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
nearest_bin_idx = torch.argmin(
torch.abs(self.value_head.bin_centers.unsqueeze(0) - target_value.unsqueeze(-1)), dim=-1
)
return F.one_hot(nearest_bin_idx, num_classes=self.config.num_value_bins).to(
dtype=self.value_head.bin_centers.dtype
)
def compute_target_distribution(
self,
target_value: Tensor,
is_terminal: Tensor,
method: str = "hl_gauss",
use_one_hot_terminal: bool = True,
) -> Tensor:
"""Compute target distribution using configured method.
Args:
target_value: [batch_size] scalar return targets
is_terminal: [batch_size] boolean terminal flags
method: "hl_gauss" or "dirac_delta"
use_one_hot_terminal: if True, terminal states get one-hot targets
(exact return, no smoothing). If False, all states use the same method.
Returns:
[batch_size, num_value_bins] target probability distribution
"""
if method == "hl_gauss":
base_distribution = self.hl_gauss_target(target_value)
elif method == "dirac_delta":
base_distribution = self.dirac_delta_target(target_value)
else:
raise ValueError(f"Unknown target method: {method}. Use 'hl_gauss' or 'dirac_delta'.")
if not use_one_hot_terminal:
return base_distribution
terminal_distribution = self.one_hot_target(target_value)
return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution)
def _vlm_forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, Tensor]:
"""Shared VLM forward: images + text → Gemma3 → last-token hidden → logits.
Returns:
(value_logits [B, num_bins], predicted_value [B, 1])
"""
images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
text_embs = self.embed_text(token_ids)
prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
if self.config.stop_gradient_to_vlm:
prefix_embs = prefix_embs.detach()
device = prefix_embs.device
model_dtype = next(self.gemma3.parameters()).dtype
# Bidirectional attention: every valid token attends to every valid token
att_2d = prefix_pad_masks[:, None, :] * prefix_pad_masks[:, :, None]
att_4d = torch.where(
att_2d[:, None, :, :],
torch.tensor(0.0, dtype=model_dtype, device=device),
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
)
position_ids = torch.cumsum(prefix_pad_masks.long(), dim=1) - 1
if prefix_embs.dtype != model_dtype:
prefix_embs = prefix_embs.to(model_dtype)
outputs = self.gemma3.model(
inputs_embeds=prefix_embs,
attention_mask=att_4d,
position_ids=position_ids,
)
# Readout from last hidden layer
if self.config.readout == "mean_pool":
hidden = outputs.last_hidden_state
mask = prefix_pad_masks.unsqueeze(-1).to(dtype=hidden.dtype)
readout = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1)
else:
readout = outputs.last_hidden_state[:, -1, :]
value_logits = self.value_head(readout)
value_probs = F.softmax(value_logits, dim=-1)
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
dim=-1, keepdim=True
)
return value_logits, predicted_value
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
"""Training forward pass — cross-entropy loss on MC return targets."""
mc_return = batch["mc_return"]
is_terminal = batch["is_terminal"]
value_logits, predicted_value = self._vlm_forward(batch)
# Compute target distribution from MC returns
target_dist = self.compute_target_distribution(
mc_return,
is_terminal,
method=self.config.target_method,
use_one_hot_terminal=self.config.use_one_hot_terminal,
)
# Cross-entropy loss between predicted and target distributions (Eq. 1 in pi*0.6 paper)
log_probs = F.log_softmax(value_logits, dim=-1)
loss = -(target_dist * log_probs).sum(dim=-1).mean()
# Diagnostic metrics
clamped_return = (
mc_return.float().view(-1).clamp(self.config.value_support_min, self.config.value_support_max)
)
bin_width = self.value_head.bin_centers[1] - self.value_head.bin_centers[0]
normalized_position = (clamped_return - self.config.value_support_min) / bin_width
lower_bin_idx = normalized_position.floor().long().clamp(0, self.config.num_value_bins - 1)
upper_bin_idx = normalized_position.ceil().long().clamp(0, self.config.num_value_bins - 1)
dist_to_lower = normalized_position - lower_bin_idx.float()
dist_to_upper = upper_bin_idx.float() - normalized_position
same_bin = lower_bin_idx == upper_bin_idx
dist_to_lower = torch.where(same_bin, torch.zeros_like(dist_to_lower), dist_to_lower)
dist_to_upper = torch.where(same_bin, torch.ones_like(dist_to_upper), dist_to_upper)
pred_bin = value_logits.argmax(dim=-1)
best_target_bin = torch.where(dist_to_upper >= dist_to_lower, lower_bin_idx, upper_bin_idx)
acc_best = (pred_bin == best_target_bin).float().mean().item()
acc_neighbor = ((pred_bin == lower_bin_idx) | (pred_bin == upper_bin_idx)).float().mean().item()
min_bin_dist = torch.min((pred_bin - lower_bin_idx).abs(), (pred_bin - upper_bin_idx).abs()).float()
mae = (min_bin_dist * bin_width).mean().item()
output_dict: dict[str, Any] = {
"loss": loss.item(),
"predicted_value_mean": predicted_value.mean().item(),
"mc_return_mean": mc_return.mean().item(),
"acc_best": acc_best,
"acc_neighbor": acc_neighbor,
"mae": mae,
}
return loss, output_dict
def _get_model_inputs(
self, batch: dict[str, Tensor]
) -> tuple[list[Tensor], list[Tensor], Tensor, Tensor]:
"""Extract images, masks, token_ids, text_pad_mask from a preprocessed batch."""
image_keys = [k for k, v in self.config.input_features.items() if v.type == FeatureType.VISUAL]
images = [batch[k] for k in image_keys]
img_masks = [batch[k + IMAGE_MASK_SUFFIX].bool() for k in image_keys]
token_ids = batch[OBS_LANGUAGE_TOKENS]
text_pad_mask = batch[OBS_LANGUAGE_ATTENTION_MASK].bool()
return images, img_masks, token_ids, text_pad_mask
def compute_reward(self, batch: dict[str, Tensor]) -> Tensor:
"""Compute V(s) for a batch of observations. Used for advantage scoring.
Args:
batch: preprocessed batch with images, masks, and tokenized text.
Returns:
[batch_size] tensor of predicted values V(s).
"""
_, predicted_value = self._vlm_forward(batch)
return predicted_value.squeeze(-1)
@@ -0,0 +1,283 @@
# 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.
"""Processor for RECAP's distributional value function.
Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06
Prepares inputs for V^{pi_ref}(o_t, l):
1. Resize multi-camera images to 224x224 (with aspect-preserving padding)
2. Normalize images from [0,1] [-1,1] (SigLIP standard)
3. Handle missing cameras (placeholder + mask)
4. Format task prompt: ``"Task: {task}."``
5. Tokenize with Gemma3 tokenizer
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
import torch.nn.functional as F # noqa: N812
from torch import Tensor
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
from lerobot.processor import (
AddBatchDimensionProcessorStep,
DeviceProcessorStep,
NormalizerProcessorStep,
PolicyAction,
PolicyProcessorPipeline,
ProcessorStep,
ProcessorStepRegistry,
RenameObservationsProcessorStep,
TokenizerProcessorStep,
batch_to_transition,
policy_action_to_transition,
transition_to_batch,
)
from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import (
POLICY_POSTPROCESSOR_DEFAULT_NAME,
POLICY_PREPROCESSOR_DEFAULT_NAME,
)
from .configuration_distributional_value_function import DistributionalVFConfig
# Keys used by the image processor to store per-camera validity masks.
IMAGE_MASK_SUFFIX = ".mask"
def resize_with_pad_torch(
images: Tensor,
height: int,
width: int,
mode: str = "bilinear",
) -> Tensor:
"""Resize images preserving aspect ratio, padding with black.
Matches ``resize_with_pad_torch`` in PI0/PI05/PI0-FAST.
Args:
images: [*b, h, w, c] or [*b, c, h, w] tensor.
height: Target height.
width: Target width.
mode: Interpolation mode.
Returns:
Resized and padded tensor with same shape format as input.
"""
if images.shape[-1] <= 4:
channels_last = True
if images.dim() == 3:
images = images.unsqueeze(0)
images = images.permute(0, 3, 1, 2)
else:
channels_last = False
if images.dim() == 3:
images = images.unsqueeze(0)
batch_size, channels, cur_height, cur_width = images.shape
ratio = max(cur_width / width, cur_height / height)
resized_height = int(cur_height / ratio)
resized_width = int(cur_width / ratio)
resized_images = F.interpolate(
images,
size=(resized_height, resized_width),
mode=mode,
align_corners=False if mode == "bilinear" else None,
)
if images.dtype == torch.uint8:
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
elif images.dtype == torch.float32:
resized_images = resized_images.clamp(-1.0, 1.0)
pad_h0, remainder_h = divmod(height - resized_height, 2)
pad_h1 = pad_h0 + remainder_h
pad_w0, remainder_w = divmod(width - resized_width, 2)
pad_w1 = pad_w0 + remainder_w
constant_value = 0 if images.dtype == torch.uint8 else -1.0
padded_images = F.pad(
resized_images,
(pad_w0, pad_w1, pad_h0, pad_h1),
mode="constant",
value=constant_value,
)
if channels_last:
padded_images = padded_images.permute(0, 2, 3, 1)
return padded_images
@ProcessorStepRegistry.register(name="distributional_vf_image_preprocessor")
@dataclass
class DistributionalVFImagePreprocessorStep(ProcessorStep):
"""Resize and normalize multi-camera images for the VF.
Produces [B, 3, H, W] tensors in [-1, 1] for each camera, plus boolean
masks indicating which cameras are present. Missing cameras get a black
placeholder image and mask=False.
"""
image_resolution: tuple[int, int] = (224, 224)
image_keys: tuple[str, ...] = ()
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
observation = dict(transition.get(TransitionKey.OBSERVATION, {}))
for key in self.image_keys:
if key in observation:
img = observation[key]
if img.dtype != torch.float32:
img = img.to(torch.float32)
is_channels_first = img.shape[1] == 3
if is_channels_first:
img = img.permute(0, 2, 3, 1) # BCHW → BHWC
if img.shape[1:3] != self.image_resolution:
img = resize_with_pad_torch(img, *self.image_resolution)
if img.min() >= 0.0 and img.max() <= 1.0:
img = img * 2.0 - 1.0
observation[key] = img.permute(0, 3, 1, 2) # BHWC → BCHW
observation[key + IMAGE_MASK_SUFFIX] = torch.ones(
img.shape[0], dtype=torch.bool, device=img.device
)
else:
bsize = self._infer_batch_size(observation)
h, w = self.image_resolution
observation[key] = torch.full((bsize, 3, h, w), -1.0)
observation[key + IMAGE_MASK_SUFFIX] = torch.zeros(bsize, dtype=torch.bool)
transition[TransitionKey.OBSERVATION] = observation
return transition
def _infer_batch_size(self, observation: dict) -> int:
for v in observation.values():
if isinstance(v, Tensor) and v.ndim >= 2:
return v.shape[0]
return 1
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
def get_config(self) -> dict[str, Any]:
return {
"image_resolution": self.image_resolution,
"image_keys": self.image_keys,
}
@ProcessorStepRegistry.register(name="distributional_vf_prepare_task_prompt")
@dataclass
class DistributionalVFPrepareTaskPromptStep(ProcessorStep):
"""Format the task string: ``"Task: {task}."``"""
task_key: str = "task"
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key)
if tasks is None:
raise ValueError("No task found in complementary data")
if isinstance(tasks, str):
tasks = [tasks]
full_prompts = []
for task in tasks:
cleaned_text = task.strip().replace("_", " ").replace("\n", " ")
full_prompts.append(f"Task: {cleaned_text}.")
new_complementary_data = dict(transition.get(TransitionKey.COMPLEMENTARY_DATA, {}))
new_complementary_data[self.task_key] = full_prompts
transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
return transition
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
def get_config(self) -> dict[str, Any]:
return {"task_key": self.task_key}
def make_distributional_vf_pre_post_processors(
config: DistributionalVFConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction],
]:
"""Create pre/post processors for the distributional value function.
Preprocessor steps:
1. Rename observations (no-op by default)
2. Add a batch dimension
3. Normalize features (identity for images)
4. Resize + normalize images [B, 3, 224, 224] in [-1, 1]
5. Format task prompt: ``"Task: {task}."``
6. Tokenize with Gemma3 tokenizer
7. Move tensors to the configured device
Training targets (mc_return, is_terminal) are not processed here.
The postprocessor is a no-op (value function does not produce actions).
"""
image_keys = tuple(k for k, v in config.input_features.items() if v.type == FeatureType.VISUAL)
preprocessor = PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
steps=[
RenameObservationsProcessorStep(rename_map={}),
AddBatchDimensionProcessorStep(),
NormalizerProcessorStep(
features={**config.input_features, **config.output_features},
norm_map=config.normalization_mapping,
stats=dataset_stats,
),
DistributionalVFImagePreprocessorStep(
image_resolution=config.image_resolution,
image_keys=image_keys,
),
DistributionalVFPrepareTaskPromptStep(),
TokenizerProcessorStep(
tokenizer_name=config.gemma3_path,
max_length=config.tokenizer_max_length,
padding_side="right",
padding="max_length",
),
DeviceProcessorStep(device=config.device or "cpu"),
],
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
to_transition=batch_to_transition,
to_output=transition_to_batch,
)
postprocessor = PolicyProcessorPipeline(
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
to_transition=policy_action_to_transition,
)
return preprocessor, postprocessor
+19
View File
@@ -24,6 +24,7 @@ from lerobot.configs.rewards import RewardModelConfig
from lerobot.processor import PolicyAction, PolicyProcessorPipeline
from .classifier.configuration_classifier import RewardClassifierConfig
from .distributional_value_function.configuration_distributional_value_function import DistributionalVFConfig
from .pretrained import PreTrainedRewardModel
from .robometer.configuration_robometer import RobometerConfig
from .sarm.configuration_sarm import SARMConfig
@@ -63,6 +64,12 @@ def get_reward_model_class(name: str) -> type[PreTrainedRewardModel]:
from lerobot.rewards.topreward.modeling_topreward import TOPRewardModel
return TOPRewardModel
elif name == "distributional_value_function":
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
return DistributionalVFRewardModel
else:
try:
return _get_reward_model_cls_from_name(name=name)
@@ -96,6 +103,8 @@ def make_reward_model_config(reward_type: str, **kwargs) -> RewardModelConfig:
return RobometerConfig(**kwargs)
elif reward_type == "topreward":
return TOPRewardConfig(**kwargs)
elif reward_type == "distributional_value_function":
return DistributionalVFConfig(**kwargs)
else:
try:
config_cls = RewardModelConfig.get_choice_class(reward_type)
@@ -192,6 +201,16 @@ def make_reward_pre_post_processors(
dataset_stats=kwargs.get("dataset_stats"),
)
elif isinstance(reward_cfg, DistributionalVFConfig):
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
make_distributional_vf_pre_post_processors,
)
return make_distributional_vf_pre_post_processors(
config=reward_cfg,
dataset_stats=kwargs.get("dataset_stats"),
)
else:
try:
processors = _make_processors_from_reward_model_config(
@@ -182,7 +182,7 @@ class SOFollower(Robot):
def get_observation(self) -> RobotObservation:
# Read arm position
start = time.perf_counter()
obs_dict = self.bus.sync_read("Present_Position")
obs_dict = self.bus.sync_read("Present_Position", num_retry=3)
obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read state: {dt_ms:.1f}ms")
@@ -223,12 +223,12 @@ class SOFollower(Robot):
# Cap goal position when too far away from present position.
# /!\ Slower fps expected due to reading from the follower.
if self.config.max_relative_target is not None:
present_pos = self.bus.sync_read("Present_Position")
present_pos = self.bus.sync_read("Present_Position", num_retry=3)
goal_present_pos = {key: (g_pos, present_pos[key]) for key, g_pos in goal_pos.items()}
goal_pos = ensure_safe_goal_position(goal_present_pos, self.config.max_relative_target)
# Send goal position to the arm
self.bus.sync_write("Goal_Position", goal_pos)
self.bus.sync_write("Goal_Position", goal_pos, num_retry=3)
return {f"{motor}.pos": val for motor, val in goal_pos.items()}
@check_if_not_connected
+8
View File
@@ -106,6 +106,8 @@ class DAggerKeyboardConfig:
pause_resume: str = "space"
correction: str = "tab"
upload: str = "enter"
success: str = "s"
failure: str = "f"
@dataclass
@@ -119,6 +121,8 @@ class DAggerPedalConfig:
pause_resume: str = "KEY_A"
correction: str = "KEY_B"
upload: str = "KEY_C"
success: str = "KEY_D"
failure: str = "KEY_E"
@RolloutStrategyConfig.register_subclass("episodic")
@@ -165,6 +169,10 @@ class DAggerStrategyConfig(RolloutStrategyConfig):
2. **correction** toggle human correction recording.
3. **upload** push dataset to hub on demand (corrections-only mode).
Episode success labeling:
4. **success** mark current episode as successful.
5. **failure** mark current episode as failed.
When ``record_autonomous=False`` (default) only human-correction windows
are recorded each correction becomes its own episode. Set to ``True``
to record both autonomous and correction frames with size-based episode
+5
View File
@@ -350,6 +350,11 @@ def build_rollout_context(
"shape": (1,),
"names": None,
}
dataset_features["next.success"] = {
"dtype": "bool",
"shape": (1,),
"names": None,
}
repo_name = cfg.dataset.repo_id.split("/", 1)[-1]
if not repo_name.startswith("rollout_"):
+141 -16
View File
@@ -112,6 +112,14 @@ class DAggerEvents:
# Session-level flags
self.stop_recording = Event()
self.upload_requested = Event()
# Set when operator presses success/failure key to end the current episode.
self.save_episode_requested = Event()
# Episode success labeling
self._episode_success: bool | None = None
# Episode success labeling
self._episode_success: bool | None = None
# -- Thread-safe phase access ------------------------------------------
@@ -155,7 +163,43 @@ class DAggerEvents:
with self._lock:
self._phase = DAggerPhase.AUTONOMOUS
self._pending_transition = None
self._episode_success = None
self.upload_requested.clear()
self.save_episode_requested.clear()
def mark_success(self) -> None:
"""Mark the current episode as successful (called from input threads)."""
with self._lock:
self._episode_success = True
def mark_failure(self) -> None:
"""Mark the current episode as failed (called from input threads)."""
with self._lock:
self._episode_success = False
def consume_episode_success(self) -> bool | None:
"""Consume and reset the episode success label. Returns None if unlabeled."""
with self._lock:
result = self._episode_success
self._episode_success = None
return result
def mark_success(self) -> None:
"""Mark the current episode as successful (called from input threads)."""
with self._lock:
self._episode_success = True
def mark_failure(self) -> None:
"""Mark the current episode as failed (called from input threads)."""
with self._lock:
self._episode_success = False
def consume_episode_success(self) -> bool | None:
"""Consume and reset the episode success label. Returns None if unlabeled."""
with self._lock:
result = self._episode_success
self._episode_success = None
return result
# ---------------------------------------------------------------------------
@@ -186,12 +230,20 @@ def _init_dagger_keyboard(events: DAggerEvents, cfg: DAggerKeyboardConfig):
events.request_transition(key_to_event[name])
if name == cfg.upload:
events.upload_requested.set()
if name == cfg.success:
events.mark_success()
events.save_episode_requested.set()
logger.info("Episode marked as SUCCESS — saving")
if name == cfg.failure:
events.mark_failure()
events.save_episode_requested.set()
logger.info("Episode marked as FAILURE — saving")
return create_key_listener(
dispatch,
controls_help=(
f"pause_resume='{cfg.pause_resume}', correction='{cfg.correction}', "
f"upload='{cfg.upload}', ESC=stop"
f"upload='{cfg.upload}', success='{cfg.success}', failure='{cfg.failure}', ESC=stop"
),
)
@@ -211,6 +263,12 @@ def _init_dagger_pedal(events: DAggerEvents, cfg: DAggerPedalConfig):
events.request_transition(code_to_event[code])
if code == cfg.upload:
events.upload_requested.set()
if code == cfg.success:
events.mark_success()
logger.info("Episode marked as SUCCESS (pedal)")
if code == cfg.failure:
events.mark_failure()
logger.info("Episode marked as FAILURE (pedal)")
logger.info("Initializing DAgger foot pedal listener (device=%s)", cfg.device_path)
return start_pedal_listener(on_press, device_path=cfg.device_path)
@@ -313,6 +371,31 @@ class DAggerStrategy(RolloutStrategy):
)
logger.info("DAgger strategy teardown complete")
# ------------------------------------------------------------------
# Episode success labeling
# ------------------------------------------------------------------
def _stamp_episode_success(self, dataset) -> None:
"""Set next.success on the terminal frame based on operator label.
Called just before save_episode(). If the operator pressed the success
key during this episode, the last frame's next.success is set to True.
Otherwise all frames remain False (unlabeled = assumed failure).
"""
buf = dataset.writer.episode_buffer
if buf is None:
return
success_buf = buf.get("next.success")
if not success_buf:
return
label = self._events.consume_episode_success()
if label:
success_buf[-1] = np.array([True], dtype=bool)
logger.info("Terminal frame stamped next.success=True")
# ------------------------------------------------------------------
# Continuous recording mode (record_autonomous=True)
# ------------------------------------------------------------------
@@ -350,7 +433,12 @@ class DAggerStrategy(RolloutStrategy):
episode_start = time.perf_counter()
episodes_since_push = 0
episode_duration_s = self._episode_duration_s
logger.info("DAgger continuous recording started (episode_duration=%.0fs)", episode_duration_s)
num_episodes = self.config.num_episodes
logger.info(
"DAgger continuous recording started (episode_duration=%.0fs, target=%s eps)",
episode_duration_s,
num_episodes if num_episodes is not None else "",
)
with VideoEncodingManager(dataset):
try:
@@ -399,6 +487,7 @@ class DAggerStrategy(RolloutStrategy):
**action_frame,
"task": task_str,
"intervention": np.array([True], dtype=bool),
"next.success": np.array([False], dtype=bool),
}
dataset.add_frame(frame)
record_tick += 1
@@ -427,23 +516,32 @@ class DAggerStrategy(RolloutStrategy):
**action_frame,
"task": task_str,
"intervention": np.array([False], dtype=bool),
"next.success": np.array([False], dtype=bool),
}
dataset.add_frame(frame)
record_tick += 1
# Episode rotation derived from the video file-size target.
# Saving is deferred while a correction is ongoing so the
# episode boundary lands on a clean autonomous frame.
# Episode rotation: either the operator pressed success/failure,
# or the video file-size target was reached.
# Defer the save while a correction is ongoing so the episode
# boundary lands on a clean autonomous frame. The event stays
# set until we actually save, so it won't be lost.
manual_save = events.save_episode_requested.is_set()
elapsed = time.perf_counter() - episode_start
if elapsed >= episode_duration_s and phase != DAggerPhase.CORRECTING:
if (manual_save or elapsed >= episode_duration_s) and phase != DAggerPhase.CORRECTING:
if manual_save:
events.save_episode_requested.clear()
with self._episode_lock:
self._stamp_episode_success(dataset)
dataset.save_episode()
episodes_since_push += 1
self._needs_push.set()
save_reason = "manual save" if manual_save else f"elapsed {elapsed:.1f}s"
logger.info(
"Episode saved (total: %d, elapsed: %.1fs)",
"Episode saved (%s, total: %d)",
save_reason,
dataset.num_episodes,
elapsed,
)
log_say(f"Episode {dataset.num_episodes} saved", play_sounds)
@@ -451,6 +549,25 @@ class DAggerStrategy(RolloutStrategy):
self._background_push(dataset, cfg)
episodes_since_push = 0
if num_episodes is not None and dataset.num_episodes >= num_episodes:
logger.info("Target episode count reached (%d), stopping session", num_episodes)
log_say(f"All {num_episodes} episodes collected", play_sounds)
events.stop_recording.set()
break
# Pause after manual save: stop the policy, return robot to
# initial position, and wait for the operator to reset the
# environment and press SPACE.
if manual_save:
engine.pause()
events.phase = DAggerPhase.PAUSED
self._return_to_initial_position(ctx.hardware)
last_action = None
logger.info(
"Episode saved — paused for environment reset. Press SPACE to start next episode."
)
log_say("Reset the environment, then press space", play_sounds)
episode_start = time.perf_counter()
dt = time.perf_counter() - loop_start
@@ -465,10 +582,13 @@ class DAggerStrategy(RolloutStrategy):
logger.info("DAgger continuous control loop ended — pausing engine")
engine.pause()
with contextlib.suppress(Exception):
with self._episode_lock:
dataset.save_episode()
self._needs_push.set()
logger.info("Final in-progress episode saved")
buf = dataset.writer.episode_buffer
if buf and any(len(v) > 0 for v in buf.values() if isinstance(v, list)):
with self._episode_lock:
self._stamp_episode_success(dataset)
dataset.save_episode()
self._needs_push.set()
logger.info("Final in-progress episode saved")
# ------------------------------------------------------------------
# Corrections-only mode (record_autonomous=False)
@@ -540,6 +660,7 @@ class DAggerStrategy(RolloutStrategy):
# Correction ended -> save episode (blocking if not streaming)
if old_phase == DAggerPhase.CORRECTING and new_phase == DAggerPhase.PAUSED:
with self._episode_lock:
self._stamp_episode_success(dataset)
dataset.save_episode()
recorded += 1
self._needs_push.set()
@@ -581,6 +702,7 @@ class DAggerStrategy(RolloutStrategy):
**action_frame,
"task": task_str,
"intervention": np.array([True], dtype=bool),
"next.success": np.array([False], dtype=bool),
}
)
record_tick += 1
@@ -614,10 +736,13 @@ class DAggerStrategy(RolloutStrategy):
logger.info("DAgger corrections-only loop ended — pausing engine")
engine.pause()
with contextlib.suppress(Exception):
with self._episode_lock:
dataset.save_episode()
self._needs_push.set()
logger.info("Final in-progress episode saved")
buf = dataset.writer.episode_buffer
if buf and any(len(v) > 0 for v in buf.values() if isinstance(v, list)):
with self._episode_lock:
self._stamp_episode_success(dataset)
dataset.save_episode()
self._needs_push.set()
logger.info("Final in-progress episode saved")
# ------------------------------------------------------------------
# State-machine transition side-effects
+39 -13
View File
@@ -34,6 +34,7 @@ from lerobot.annotations.steerable_pipeline.config import AnnotationPipelineConf
from lerobot.annotations.steerable_pipeline.executor import Executor
from lerobot.annotations.steerable_pipeline.frames import make_frame_provider
from lerobot.annotations.steerable_pipeline.modules import (
AdvantageModule,
GeneralVqaModule,
InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule,
@@ -47,6 +48,7 @@ logger = logging.getLogger(__name__)
def _resolve_root(cfg: AnnotationPipelineConfig) -> Path:
"""Resolve the dataset root by downloading the full snapshot if needed."""
if cfg.root is not None:
return Path(cfg.root)
if cfg.repo_id is not None:
@@ -63,17 +65,24 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
root = _resolve_root(cfg)
logger.info("annotate: root=%s", root)
vlm = make_vlm_client(cfg.vlm)
frame_provider = make_frame_provider(root, camera_key=cfg.vlm.camera_key, video_backend=cfg.video_backend)
needs_vlm = cfg.plan.enabled or cfg.interjections.enabled or cfg.vqa.enabled
needs_video = needs_vlm or cfg.advantage.enabled
vlm = make_vlm_client(cfg.vlm) if needs_vlm else None
frame_provider = (
make_frame_provider(root, camera_key=cfg.vlm.camera_key, video_backend=cfg.video_backend)
if needs_video
else None
)
# Surface the resolved cameras up front so a silent vqa-module no-op
# is obvious in job output rather than discovered post-hoc by counting
# parquet rows.
cam_keys = list(getattr(frame_provider, "camera_keys", []) or [])
logger.info(
"annotate: frame_provider default camera=%r, all cameras=%s",
getattr(frame_provider, "camera_key", None),
cam_keys,
)
cam_keys = list(getattr(frame_provider, "camera_keys", []) or []) if frame_provider else []
if frame_provider:
logger.info(
"annotate: frame_provider default camera=%r, all cameras=%s",
getattr(frame_provider, "camera_key", None),
cam_keys,
)
if cfg.vqa.enabled and not cam_keys:
logger.warning(
"annotate: the vqa module is enabled but no cameras were "
@@ -81,14 +90,30 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
"meta/info.json for observation.images.* features, or pass "
"--vlm.camera_key=<key> to seed the cameras list."
)
plan = PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider)
interjections = InterjectionsAndSpeechModule(
vlm=vlm, config=cfg.interjections, seed=cfg.seed, frame_provider=frame_provider
plan = (
PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider)
if needs_vlm
else None
)
interjections = (
InterjectionsAndSpeechModule(
vlm=vlm, config=cfg.interjections, seed=cfg.seed, frame_provider=frame_provider
)
if needs_vlm
else None
)
vqa = (
GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed, frame_provider=frame_provider)
if needs_vlm
else None
)
advantage = AdvantageModule(
config=cfg.advantage,
**({"frame_provider": frame_provider} if frame_provider is not None else {}),
)
vqa = GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed, frame_provider=frame_provider)
writer = LanguageColumnsWriter()
validator = StagingValidator(
dataset_camera_keys=tuple(getattr(frame_provider, "camera_keys", []) or []) or None,
dataset_camera_keys=tuple(cam_keys) or None,
)
executor = Executor(
@@ -96,6 +121,7 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
plan=plan,
interjections=interjections,
vqa=vqa,
advantage=advantage,
writer=writer,
validator=validator,
)
@@ -0,0 +1,409 @@
#!/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.
# 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.
"""Compute per-frame ``is_terminal`` and ``mc_return`` for a LeRobot dataset.
Implements the sparse reward function from pi*0.6 / RECAP (Eq. 5):
r_t = -1 for non-terminal steps
r_T = 0 for terminal success
r_T = -C_fail for terminal failure
Returns are normalized by ``H - 1 + C_fail`` so mc_return [-1, 0], where
``H`` is the longest episode (or ``--max-episode-length``).
The columns are written directly into the dataset's parquet data shards as
flat per-frame scalars. These serve as training targets for the distributional
value function.
Usage:
# Compute returns using the default "next.success" column (from lerobot-eval/rollout)
lerobot-compute-returns \\
--dataset-repo-id lerobot/aloha_sim_insertion_human_image
# Override: treat all episodes as successful (demo-only datasets)
lerobot-compute-returns \\
--dataset-repo-id lerobot/aloha_sim_insertion_human_image \\
--default-success true
# Custom success key, failure penalty, and discount
lerobot-compute-returns \\
--dataset-repo-id my_org/my_dataset \\
--success-key episode_success \\
--c-fail 100 \\
--gamma 0.99
"""
from __future__ import annotations
import argparse
import json
import logging
from dataclasses import dataclass, field
from pathlib import Path
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
from tqdm import tqdm
logger = logging.getLogger(__name__)
IS_TERMINAL_COL = "is_terminal"
MC_RETURN_COL = "mc_return"
@dataclass
class ComputeReturnsConfig:
"""Configuration for the returns computation script."""
dataset_repo_id: str = ""
root: str | None = None
success_key: str = "next.success"
default_success: bool | None = None
max_episode_length: int | None = None
c_fail: float = 50.0
gamma: float = 1.0
episodes: list[int] = field(default_factory=list)
force: bool = False
push_to_hub: bool = False
def _get_episode_success(
episode_table: pa.Table,
success_key: str,
default_success: bool | None,
) -> bool:
"""Determine whether an episode was successful.
Priority:
1. If ``default_success`` is set, use it unconditionally.
2. Look for ``success_key`` in the parquet columns and reduce with any().
3. Fall back to True (assume success for demo datasets).
"""
if default_success is not None:
return default_success
if success_key in episode_table.column_names:
col = episode_table.column(success_key)
for val in col:
py_val = val.as_py()
if isinstance(py_val, bool) and py_val:
return True
if isinstance(py_val, (int, float)) and py_val:
return True
return False
return True
def compute_episode_returns(
num_frames: int,
success: bool,
c_fail: float,
gamma: float,
max_episode_length: int,
) -> tuple[np.ndarray, np.ndarray]:
"""Compute is_terminal and mc_return arrays for a single episode.
Rewards (RECAP Eq. 5): r_t = -1, r_T = 0 (success) or -C_fail (failure),
normalized by ``H - 1 + C_fail`` so mc_return [-1, 0].
Args:
num_frames: Number of frames in the episode.
success: Whether the episode ended successfully.
c_fail: Failure penalty constant.
gamma: Discount factor (1.0 = undiscounted).
max_episode_length: Normalization horizon H.
Returns:
Tuple of (is_terminal, mc_return) arrays, each of length num_frames.
"""
normalizer = max_episode_length - 1 + c_fail
rewards = np.full(num_frames, -1.0 / normalizer, dtype=np.float64)
if success:
rewards[-1] = 0.0
else:
rewards[-1] = -c_fail / normalizer
is_terminal = np.zeros(num_frames, dtype=bool)
is_terminal[-1] = True
if gamma == 1.0:
mc_return = np.cumsum(rewards[::-1])[::-1].astype(np.float32)
else:
mc_return = np.zeros(num_frames, dtype=np.float64)
mc_return[-1] = rewards[-1]
for t in range(num_frames - 2, -1, -1):
mc_return[t] = rewards[t] + gamma * mc_return[t + 1]
mc_return = mc_return.astype(np.float32)
return is_terminal, mc_return
def compute_returns(config: ComputeReturnsConfig) -> Path:
"""Compute returns and write them into parquet shards."""
from lerobot.datasets import LeRobotDataset
logger.info(f"Loading dataset: {config.dataset_repo_id}")
kwargs = {"repo_id": config.dataset_repo_id, "download_videos": False}
if config.root:
kwargs["root"] = config.root
dataset = LeRobotDataset(**kwargs)
meta = dataset.meta
root = Path(meta.root)
logger.info(f"Dataset root: {root}")
logger.info(f"Episodes: {meta.total_episodes}, Frames: {meta.total_frames}")
episode_indices = config.episodes if config.episodes else list(range(meta.total_episodes))
if config.max_episode_length is not None:
max_ep_len = config.max_episode_length
else:
max_ep_len = max(int(meta.episodes[i]["length"]) for i in episode_indices)
normalizer = max_ep_len - 1 + config.c_fail
logger.info(
f"H={max_ep_len}, normalizer={normalizer:.1f}, "
f"success=[{-(max_ep_len - 1) / normalizer:.3f}, 0.0], "
f"failure_terminal={-config.c_fail / normalizer:.3f}"
)
parquet_files_to_rewrite: dict[Path, list[int]] = {}
for ep_idx in episode_indices:
rel_path = meta.get_data_file_path(ep_idx)
abs_path = root / rel_path
parquet_files_to_rewrite.setdefault(abs_path, []).append(ep_idx)
logger.info(f"Parquet shards to rewrite: {len(parquet_files_to_rewrite)}")
for parquet_path, ep_indices_in_file in tqdm(parquet_files_to_rewrite.items(), desc="Processing shards"):
table = pq.read_table(parquet_path)
if not config.force and IS_TERMINAL_COL in table.column_names:
logger.info(f"Skipping {parquet_path.name} (already has {IS_TERMINAL_COL})")
continue
all_is_terminal = np.zeros(len(table), dtype=bool)
all_mc_return = np.zeros(len(table), dtype=np.float32)
episode_col = table.column("episode_index").to_pylist()
for ep_idx in ep_indices_in_file:
ep_info = meta.episodes[ep_idx]
ep_from = int(ep_info["dataset_from_index"])
ep_to = int(ep_info["dataset_to_index"])
ep_len = ep_to - ep_from
mask = np.array([v == ep_idx for v in episode_col], dtype=bool)
local_indices = np.where(mask)[0]
if len(local_indices) != ep_len:
logger.warning(
f"Episode {ep_idx}: expected {ep_len} frames in shard, "
f"found {len(local_indices)}. Using found count."
)
ep_len = len(local_indices)
if ep_len == 0:
continue
ep_subtable = table.filter(mask)
success = _get_episode_success(ep_subtable, config.success_key, config.default_success)
is_terminal, mc_return = compute_episode_returns(
num_frames=ep_len,
success=success,
c_fail=config.c_fail,
gamma=config.gamma,
max_episode_length=max_ep_len,
)
all_is_terminal[local_indices] = is_terminal
all_mc_return[local_indices] = mc_return
if IS_TERMINAL_COL in table.column_names:
table = table.drop(IS_TERMINAL_COL)
if MC_RETURN_COL in table.column_names:
table = table.drop(MC_RETURN_COL)
table = table.append_column(IS_TERMINAL_COL, pa.array(all_is_terminal))
table = table.append_column(MC_RETURN_COL, pa.array(all_mc_return))
pq.write_table(table, parquet_path)
_update_info_json(root, meta)
logger.info("Done. Columns written: is_terminal, mc_return")
if config.push_to_hub:
from huggingface_hub import HfApi
api = HfApi()
logger.info(f"Pushing updated dataset to Hub: {config.dataset_repo_id}")
api.upload_folder(
folder_path=str(root),
repo_id=config.dataset_repo_id,
repo_type="dataset",
)
logger.info("Push to Hub complete.")
return root
def _update_info_json(root: Path, meta) -> None:
"""Add is_terminal and mc_return to the dataset's info.json features."""
info_path = root / "meta" / "info.json"
if not info_path.exists():
logger.warning(f"info.json not found at {info_path}, skipping metadata update.")
return
info = json.loads(info_path.read_text())
features = info.get("features", {})
changed = False
if IS_TERMINAL_COL not in features:
features[IS_TERMINAL_COL] = {
"dtype": "bool",
"shape": [1],
"names": None,
}
changed = True
if MC_RETURN_COL not in features:
features[MC_RETURN_COL] = {
"dtype": "float32",
"shape": [1],
"names": None,
}
changed = True
if changed:
info["features"] = features
info_path.write_text(json.dumps(info, indent=2) + "\n")
logger.info("Updated meta/info.json with is_terminal and mc_return features.")
def main():
parser = argparse.ArgumentParser(
description="Compute per-frame is_terminal and mc_return for a LeRobot dataset.",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Use the 'success' column from the dataset
lerobot-compute-returns --dataset-repo-id lerobot/aloha_sim_insertion_human_image
# Override all episodes as successful (demo-only data)
lerobot-compute-returns --dataset-repo-id my_org/my_dataset --default-success true
# Custom failure penalty
lerobot-compute-returns --dataset-repo-id my_org/my_dataset --c-fail 100
""",
)
parser.add_argument(
"--dataset-repo-id",
type=str,
required=True,
help="HuggingFace dataset repo id or local path.",
)
parser.add_argument(
"--root",
type=str,
default=None,
help="Local root directory override for the dataset.",
)
parser.add_argument(
"--success-key",
type=str,
default="next.success",
help="Column name in parquet that indicates episode success (default: 'next.success').",
)
parser.add_argument(
"--default-success",
type=str,
default=None,
choices=["true", "false"],
help="Override success for all episodes ('true' or 'false').",
)
parser.add_argument(
"--max-episode-length",
type=int,
default=None,
help="Normalization horizon H. If not provided, inferred from the dataset as the longest episode.",
)
parser.add_argument(
"--c-fail",
type=float,
default=900.0,
help="Failure penalty constant (default: 900.0). Larger values increase separation "
"between success and failure returns.",
)
parser.add_argument(
"--gamma",
type=float,
default=1.0,
help="Discount factor (default: 1.0, undiscounted).",
)
parser.add_argument(
"--episodes",
type=int,
nargs="+",
default=None,
help="Process only these episode indices (default: all).",
)
parser.add_argument(
"--force",
action="store_true",
help="Overwrite existing is_terminal/mc_return columns.",
)
parser.add_argument(
"--push-to-hub",
action="store_true",
help="Push the updated dataset to the Hugging Face Hub after computing returns.",
)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
default_success = None
if args.default_success is not None:
default_success = args.default_success.lower() == "true"
config = ComputeReturnsConfig(
dataset_repo_id=args.dataset_repo_id,
root=args.root,
success_key=args.success_key,
default_success=default_success,
max_episode_length=args.max_episode_length,
c_fail=args.c_fail,
gamma=args.gamma,
episodes=args.episodes or [],
force=args.force,
push_to_hub=args.push_to_hub,
)
root = compute_returns(config)
logger.info(f"Returns computed and written to: {root}")
logger.info(f" Columns added: {IS_TERMINAL_COL}, {MC_RETURN_COL}")
logger.info("To train the distributional value function, these columns")
logger.info("will be read as flat batch keys during training.")
if __name__ == "__main__":
main()
@@ -290,6 +290,8 @@ class MergeConfig(OperationConfig):
# When False, keep one file per source file instead of packing into shards.
concatenate_videos: bool = True
concatenate_data: bool = True
# Allow merging datasets with different feature sets (union + fill defaults).
lenient: bool = False
@OperationConfig.register_subclass("remove_feature")
@@ -498,6 +500,7 @@ def handle_merge(cfg: EditDatasetConfig) -> None:
output_dir=output_dir,
concatenate_videos=cfg.operation.concatenate_videos,
concatenate_data=cfg.operation.concatenate_data,
lenient=cfg.operation.lenient,
)
logging.info(f"Merged dataset saved to {output_dir}")
+2
View File
@@ -741,6 +741,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
# PEFT only applies when training a policy — reward models use the plain path.
if not cfg.is_reward_model_training and cfg.policy.use_peft:
unwrapped_model.push_model_to_hub(cfg, peft_model=unwrapped_model, dataset_meta=dataset.meta)
elif cfg.is_reward_model_training:
unwrapped_model.push_model_to_hub(cfg)
else:
unwrapped_model.push_model_to_hub(cfg, state_dict=model_state_dict, dataset_meta=dataset.meta)
preprocessor.push_to_hub(active_cfg.repo_id)
@@ -145,7 +145,7 @@ class SOLeader(Teleoperator):
@check_if_not_connected
def get_action(self) -> dict[str, float]:
start = time.perf_counter()
action = self.bus.sync_read("Present_Position")
action = self.bus.sync_read("Present_Position", num_retry=3)
action = {f"{motor}.pos": val for motor, val in action.items()}
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read action: {dt_ms:.1f}ms")
@@ -155,7 +155,7 @@ class SOLeader(Teleoperator):
def send_feedback(self, feedback: dict[str, float]) -> None:
goals = {k.removesuffix(".pos"): v for k, v in feedback.items() if k.endswith(".pos")}
if goals:
self.bus.sync_write("Goal_Position", goals)
self.bus.sync_write("Goal_Position", goals, num_retry=3)
@check_if_not_connected
def disconnect(self) -> None:
+3
View File
@@ -26,6 +26,9 @@ OBS_IMAGES = OBS_IMAGE + "s"
OBS_LANGUAGE = OBS_STR + ".language"
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
OBS_LANGUAGE_UNCOND = OBS_STR + ".language_uncond"
OBS_LANGUAGE_UNCOND_TOKENS = OBS_LANGUAGE_UNCOND + ".tokens"
OBS_LANGUAGE_UNCOND_ATTENTION_MASK = OBS_LANGUAGE_UNCOND + ".attention_mask"
OBS_LANGUAGE_SUBTASK = OBS_STR + ".subtask"
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
+3 -1
View File
@@ -28,9 +28,10 @@ import sys
import tempfile
from pathlib import Path
from lerobot.annotations.steerable_pipeline.config import AnnotationPipelineConfig
from lerobot.annotations.steerable_pipeline.config import AdvantageConfig, AnnotationPipelineConfig
from lerobot.annotations.steerable_pipeline.executor import Executor
from lerobot.annotations.steerable_pipeline.modules import (
AdvantageModule,
GeneralVqaModule,
InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule,
@@ -85,6 +86,7 @@ def main() -> int:
plan=PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan),
interjections=InterjectionsAndSpeechModule(vlm=vlm, config=cfg.interjections, seed=cfg.seed),
vqa=GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed),
advantage=AdvantageModule(config=AdvantageConfig(enabled=False)),
writer=LanguageColumnsWriter(),
validator=StagingValidator(),
)
+295
View File
@@ -0,0 +1,295 @@
#!/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.
# 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.
"""Tests for the advantage scoring annotation module."""
from __future__ import annotations
from pathlib import Path
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
from lerobot.annotations.steerable_pipeline.config import AdvantageConfig
from lerobot.annotations.steerable_pipeline.modules.advantage import AdvantageModule
from lerobot.annotations.steerable_pipeline.reader import EpisodeRecord
from lerobot.annotations.steerable_pipeline.staging import EpisodeStaging
def _make_record(
episode_index: int = 0,
num_frames: int = 20,
task: str = "pick up the cup",
mc_returns: np.ndarray | None = None,
intervention_mask: np.ndarray | None = None,
fps: float = 10.0,
) -> EpisodeRecord:
"""Build a minimal EpisodeRecord with a mocked frames_df."""
import pandas as pd
timestamps = tuple(round(i / fps, 6) for i in range(num_frames))
frame_indices = tuple(range(num_frames))
if mc_returns is None:
mc_returns = np.linspace(-0.9, -0.1, num_frames).astype(np.float32)
data = {
"episode_index": [episode_index] * num_frames,
"frame_index": list(range(num_frames)),
"timestamp": list(timestamps),
"mc_return": mc_returns,
}
if intervention_mask is not None:
data["intervention"] = intervention_mask.astype(bool)
df = pd.DataFrame(data)
record = EpisodeRecord(
episode_index=episode_index,
episode_task=task,
frame_timestamps=timestamps,
frame_indices=frame_indices,
data_path=Path("/fake/data.parquet"),
row_offset=0,
row_count=num_frames,
)
record._frames_df_cache = df
return record
@pytest.fixture
def staging(tmp_path: Path) -> EpisodeStaging:
return EpisodeStaging(tmp_path, episode_index=0)
def test_advantage_module_disabled():
"""Disabled module has enabled=False."""
cfg = AdvantageConfig(enabled=False)
module = AdvantageModule(config=cfg)
assert not module.enabled
def test_advantage_module_enabled_by_default():
"""Module is enabled by default."""
cfg = AdvantageConfig()
module = AdvantageModule(config=cfg)
assert module.enabled
def test_run_episode_skips_without_value_function_path(staging: EpisodeStaging):
"""Module gracefully returns when no value_function_path is configured."""
cfg = AdvantageConfig(value_function_path="")
module = AdvantageModule(config=cfg)
record = _make_record()
module.run_episode(record, staging)
rows = staging.read("advantage")
assert rows == []
def test_binarization_with_mock_values(staging: EpisodeStaging):
"""Advantage binarization produces positive/negative labels based on threshold."""
num_frames = 10
mc_returns = np.array([-0.5, -0.4, -0.3, -0.2, -0.1, -0.5, -0.6, -0.7, -0.8, -0.9], dtype=np.float32)
mock_values = np.array([-0.4, -0.4, -0.4, -0.4, -0.4, -0.4, -0.4, -0.4, -0.4, -0.4], dtype=np.float32)
cfg = AdvantageConfig(
value_function_path="/fake/vf",
threshold_percentile=0.5,
)
module = AdvantageModule(config=cfg)
record = _make_record(num_frames=num_frames, mc_returns=mc_returns)
with (
patch.object(module, "_ensure_model_loaded"),
patch.object(module, "_compute_values", return_value=mock_values),
):
module.run_episode(record, staging)
rows = staging.read("advantage")
assert len(rows) == num_frames
# A_t = mc_returns - values
# advantages = [-0.1, 0.0, 0.1, 0.2, 0.3, -0.1, -0.2, -0.3, -0.4, -0.5]
# Median (50th pctile) = -0.1
# positive: advantage > -0.1 → indices 1,2,3,4
# negative: advantage <= -0.1 → indices 0,5,6,7,8,9
positives = [r for r in rows if r["content"] == "positive"]
negatives = [r for r in rows if r["content"] == "negative"]
assert len(positives) == 4
assert len(negatives) == 6
def test_intervention_frames_forced_positive(staging: EpisodeStaging):
"""Intervention frames are always scored as positive regardless of advantage value."""
num_frames = 5
mc_returns = np.array([-0.9, -0.9, -0.9, -0.9, -0.9], dtype=np.float32)
mock_values = np.array([-0.1, -0.1, -0.1, -0.1, -0.1], dtype=np.float32)
intervention = np.array([False, False, True, False, False])
cfg = AdvantageConfig(
value_function_path="/fake/vf",
force_positive_on_intervention=True,
)
module = AdvantageModule(config=cfg)
record = _make_record(num_frames=num_frames, mc_returns=mc_returns, intervention_mask=intervention)
with (
patch.object(module, "_ensure_model_loaded"),
patch.object(module, "_compute_values", return_value=mock_values),
):
module.run_episode(record, staging)
rows = staging.read("advantage")
# Frame 2 (intervention) should be positive despite negative advantage
assert rows[2]["content"] == "positive"
def test_all_frames_labeled(staging: EpisodeStaging):
"""Every frame gets an advantage label (no annotation-level dropout)."""
num_frames = 100
mc_returns = np.linspace(-0.9, -0.1, num_frames).astype(np.float32)
mock_values = np.full(num_frames, -0.5, dtype=np.float32)
cfg = AdvantageConfig(value_function_path="/fake/vf")
module = AdvantageModule(config=cfg)
record = _make_record(num_frames=num_frames, mc_returns=mc_returns)
with (
patch.object(module, "_ensure_model_loaded"),
patch.object(module, "_compute_values", return_value=mock_values),
):
module.run_episode(record, staging)
rows = staging.read("advantage")
assert len(rows) == num_frames
def test_staged_row_format(staging: EpisodeStaging):
"""Staged rows have the correct schema for language_persistent."""
num_frames = 5
mc_returns = np.array([-0.5, -0.4, -0.3, -0.2, -0.1], dtype=np.float32)
mock_values = np.full(5, -0.3, dtype=np.float32)
cfg = AdvantageConfig(value_function_path="/fake/vf")
module = AdvantageModule(config=cfg)
record = _make_record(num_frames=num_frames, mc_returns=mc_returns)
with (
patch.object(module, "_ensure_model_loaded"),
patch.object(module, "_compute_values", return_value=mock_values),
):
module.run_episode(record, staging)
rows = staging.read("advantage")
for row in rows:
assert row["role"] == "user"
assert row["content"] in ("positive", "negative")
assert row["style"] == "advantage"
assert isinstance(row["timestamp"], float)
assert row["camera"] is None
assert row["tool_calls"] is None
def test_n_step_advantage():
"""N-step advantage uses partial returns + bootstrapped value."""
num_frames = 10
mc_returns = np.linspace(-0.9, 0.0, num_frames).astype(np.float32)
mock_values = np.full(num_frames, -0.45, dtype=np.float32)
cfg = AdvantageConfig(
value_function_path="/fake/vf",
n_step=3,
)
module = AdvantageModule(config=cfg)
record = _make_record(num_frames=num_frames, mc_returns=mc_returns)
with patch.object(module, "_ensure_model_loaded"):
advantages, _ = (
module.compute_advantages_for_episode.__wrapped__(module, record)
if hasattr(module.compute_advantages_for_episode, "__wrapped__")
else (None, None)
)
# Just verify computation works - use the internal method directly
module._model = MagicMock()
module._preprocessor = MagicMock()
with patch.object(module, "_compute_values", return_value=mock_values):
advantages, _ = module.compute_advantages_for_episode(record)
# For t where t+n < num_frames: A = mc_return[t] - mc_return[t+n] + values[t+n] - values[t]
# Since values are constant: A = mc_return[t] - mc_return[t+n]
# For t where t+n >= num_frames: A = mc_return[t] - values[t]
for t in range(num_frames):
if t + 3 < num_frames:
expected = mc_returns[t] - mc_returns[t + 3] + mock_values[t + 3] - mock_values[t]
else:
expected = mc_returns[t] - mock_values[t]
np.testing.assert_almost_equal(advantages[t], expected, decimal=5)
def test_compute_threshold():
"""Threshold is computed as configured percentile of non-intervention advantages."""
cfg = AdvantageConfig(threshold_percentile=0.3)
module = AdvantageModule(config=cfg)
advantages = np.array([-1.0, -0.5, 0.0, 0.5, 1.0], dtype=np.float32)
intervention_mask = np.array([False, False, False, False, False])
threshold = module._compute_threshold(advantages, intervention_mask)
expected = float(np.percentile(advantages, 30))
assert abs(threshold - expected) < 1e-6
def test_compute_threshold_excludes_intervention():
"""Threshold computation excludes intervention frames."""
cfg = AdvantageConfig(threshold_percentile=0.5)
module = AdvantageModule(config=cfg)
advantages = np.array([100.0, -1.0, 0.0, 1.0, 100.0], dtype=np.float32)
intervention_mask = np.array([True, False, False, False, True])
threshold = module._compute_threshold(advantages, intervention_mask)
# Only non-intervention: [-1.0, 0.0, 1.0], median = 0.0
expected = float(np.percentile([-1.0, 0.0, 1.0], 50))
assert abs(threshold - expected) < 1e-6
def test_missing_mc_return_raises():
"""Module raises if mc_return column is missing from dataset."""
import pandas as pd
cfg = AdvantageConfig(value_function_path="/fake/vf")
module = AdvantageModule(config=cfg)
module._model = MagicMock()
module._preprocessor = MagicMock()
record = EpisodeRecord(
episode_index=0,
episode_task="test",
frame_timestamps=(0.0, 0.1),
frame_indices=(0, 1),
data_path=Path("/fake/data.parquet"),
row_offset=0,
row_count=2,
)
record._frames_df_cache = pd.DataFrame({"episode_index": [0, 0], "frame_index": [0, 1]})
with pytest.raises(KeyError, match="mc_return"):
module.compute_advantages_for_episode(record)
@@ -30,6 +30,7 @@ pytest.importorskip("pandas", reason="pandas is required (install lerobot[datase
import pyarrow.parquet as pq # noqa: E402
from lerobot.annotations.steerable_pipeline.config import ( # noqa: E402
AdvantageConfig,
AnnotationPipelineConfig,
InterjectionsConfig,
PlanConfig,
@@ -37,6 +38,7 @@ from lerobot.annotations.steerable_pipeline.config import ( # noqa: E402
)
from lerobot.annotations.steerable_pipeline.executor import Executor # noqa: E402
from lerobot.annotations.steerable_pipeline.modules import ( # noqa: E402
AdvantageModule,
GeneralVqaModule,
InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule,
@@ -132,6 +134,7 @@ def _build_executor() -> Executor:
plan=PlanSubtasksMemoryModule(vlm=vlm, config=config.plan),
interjections=InterjectionsAndSpeechModule(vlm=vlm, config=config.interjections, seed=config.seed),
vqa=GeneralVqaModule(vlm=vlm, config=config.vqa, seed=config.seed),
advantage=AdvantageModule(config=AdvantageConfig(enabled=False)),
writer=LanguageColumnsWriter(),
validator=StagingValidator(),
)
+145
View File
@@ -0,0 +1,145 @@
#!/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.
# 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.
"""Tests for RECAP advantage conditioning recipes."""
from __future__ import annotations
from pathlib import Path
from lerobot.configs.recipe import load_recipe
from lerobot.datasets.language_render import render_sample
RECIPES_DIR = Path(__file__).resolve().parents[2] / "src" / "lerobot" / "configs" / "recipes"
def _persistent_rows(advantage: str | None = None):
"""Build minimal persistent rows with optional advantage."""
rows = [
{
"role": "user",
"content": "pick up the cup",
"style": "task_aug",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
},
{
"role": "assistant",
"content": "reaching for the cup",
"style": "subtask",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
},
]
if advantage is not None:
rows.append(
{
"role": "user",
"content": advantage,
"style": "advantage",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
}
)
return rows
def test_recap_advantage_recipe_loads():
"""The recap_advantage.yaml recipe loads without errors."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage.yaml")
assert recipe.messages is not None
assert len(recipe.messages) == 3
assert recipe.bindings == {"advantage": "active_at(t, style=advantage)"}
def test_advantage_present_renders_indicator():
"""When advantage annotation exists, the prompt includes 'Advantage: positive'."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage.yaml")
result = render_sample(
recipe=recipe,
persistent=_persistent_rows(advantage="positive"),
events=[],
t=0.5,
sample_idx=0,
task="pick up the cup",
)
assert result is not None
messages = result["messages"]
assert len(messages) == 3
assert messages[1]["content"] == "Advantage: positive"
def test_advantage_negative_renders_indicator():
"""Negative advantage also appears in the prompt."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage.yaml")
result = render_sample(
recipe=recipe,
persistent=_persistent_rows(advantage="negative"),
events=[],
t=0.5,
sample_idx=0,
task="pick up the cup",
)
assert result is not None
messages = result["messages"]
assert messages[1]["content"] == "Advantage: negative"
def test_advantage_absent_skips_turn():
"""When no advantage annotation exists (dropout), the advantage turn is skipped."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage.yaml")
result = render_sample(
recipe=recipe,
persistent=_persistent_rows(advantage=None),
events=[],
t=0.5,
sample_idx=0,
task="pick up the cup",
)
assert result is not None
messages = result["messages"]
# Only task + subtask, no advantage turn
assert len(messages) == 2
assert messages[0]["content"] == "pick up the cup"
assert messages[1]["content"] == "reaching for the cup"
def test_advantage_absent_still_has_target():
"""Even without advantage, the target message (subtask) is preserved."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage.yaml")
result = render_sample(
recipe=recipe,
persistent=_persistent_rows(advantage=None),
events=[],
t=0.5,
sample_idx=0,
task="pick up the cup",
)
assert result is not None
assert result["target_message_indices"] == [1]
def test_blend_recipe_loads():
"""The blend recipe has two components with correct weights."""
recipe = load_recipe(RECIPES_DIR / "recap_advantage_blend.yaml")
assert recipe.blend is not None
assert "advantage_conditioned" in recipe.blend
assert "unconditional" in recipe.blend
assert recipe.blend["advantage_conditioned"].weight == 0.7
assert recipe.blend["unconditional"].weight == 0.3
+224
View File
@@ -0,0 +1,224 @@
#!/usr/bin/env python
"""Tests for PI05 Classifier-Free Guidance (CFG) inference."""
import pytest
pytest.importorskip("transformers", reason="transformers is required for PI05")
import torch # noqa: E402
from lerobot.configs.types import FeatureType, PolicyFeature # noqa: E402
from lerobot.policies.pi05 import PI05Config, make_pi05_pre_post_processors # noqa: E402
from lerobot.processor.converters import create_transition # noqa: E402
from lerobot.processor.rendered_messages_to_task import RenderedMessagesToTaskStep # noqa: E402
from lerobot.types import TransitionKey # noqa: E402
from lerobot.utils.constants import ( # noqa: E402
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
OBS_LANGUAGE_UNCOND_TOKENS,
)
class TestRenderedMessagesToTaskBaseTaskPreservation:
"""Tests that RenderedMessagesToTaskStep preserves base_task for CFG."""
def test_preserves_string_base_task(self):
transition = create_transition(
complementary_data={
"task": "pick up the cup",
"messages": [
{"role": "user", "content": "pick up the cup, Advantage: positive"},
],
}
)
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["base_task"] == "pick up the cup"
assert data["task"] == "pick up the cup, Advantage: positive"
def test_preserves_list_base_task(self):
transition = create_transition(
complementary_data={
"task": ["task1", "task2"],
"messages": [
{"role": "user", "content": "rendered with advantage"},
],
}
)
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["base_task"] == ["task1", "task2"]
def test_no_base_task_when_messages_absent(self):
transition = create_transition(complementary_data={"task": "pick up the cup"})
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert "base_task" not in data
class TestPi05PrepareStateTokenizerCfg:
"""Tests for Pi05PrepareStateTokenizerProcessorStep with cfg_enabled."""
def _make_transition(self, task, base_task=None):
complementary_data = {"task": task}
if base_task is not None:
complementary_data["base_task"] = base_task
return create_transition(
observation={"observation.state": torch.zeros(1, 14)},
complementary_data=complementary_data,
)
def test_cfg_disabled_no_uncond_task(self):
from lerobot.policies.pi05.processor_pi05 import Pi05PrepareStateTokenizerProcessorStep
step = Pi05PrepareStateTokenizerProcessorStep(max_state_dim=14, cfg_enabled=False)
transition = self._make_transition(task=["pick up the cup, Advantage: positive"])
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert "uncond_task" not in data
def test_cfg_enabled_produces_uncond_task_from_base(self):
from lerobot.policies.pi05.processor_pi05 import Pi05PrepareStateTokenizerProcessorStep
step = Pi05PrepareStateTokenizerProcessorStep(max_state_dim=14, cfg_enabled=True)
transition = self._make_transition(
task=["pick up the cup, Advantage: positive"],
base_task=["pick up the cup"],
)
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert "uncond_task" in data
assert len(data["uncond_task"]) == 1
# Unconditional prompt uses base_task (no advantage)
assert "Advantage" not in data["uncond_task"][0]
assert "pick up the cup" in data["uncond_task"][0]
assert "State:" in data["uncond_task"][0]
def test_cfg_enabled_falls_back_to_task_when_no_base(self):
from lerobot.policies.pi05.processor_pi05 import Pi05PrepareStateTokenizerProcessorStep
step = Pi05PrepareStateTokenizerProcessorStep(max_state_dim=14, cfg_enabled=True)
transition = self._make_transition(task=["pick up the cup"])
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
# Falls back to using task itself as unconditional
assert "uncond_task" in data
assert "pick up the cup" in data["uncond_task"][0]
class TestCfgPipelineConstruction:
"""Tests that the processor pipeline is constructed correctly for CFG."""
def _make_config(self, cfg_beta=1.0, recipe_path=None):
config = PI05Config(
max_action_dim=7,
max_state_dim=14,
cfg_beta=cfg_beta,
recipe_path=recipe_path,
device="cpu",
)
config.input_features = {
"observation.state": PolicyFeature(type=FeatureType.STATE, shape=(14,)),
"observation.images.base_0_rgb": PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
}
config.output_features = {
"action": PolicyFeature(type=FeatureType.ACTION, shape=(7,)),
}
return config
def _make_dataset_stats(self):
return {
"observation.state": {
"mean": torch.zeros(14),
"std": torch.ones(14),
"min": torch.zeros(14),
"max": torch.ones(14),
"q01": torch.zeros(14),
"q99": torch.ones(14),
},
"action": {
"mean": torch.zeros(7),
"std": torch.ones(7),
"min": torch.zeros(7),
"max": torch.ones(7),
"q01": torch.zeros(7),
"q99": torch.ones(7),
},
"observation.images.base_0_rgb": {
"mean": torch.zeros(3, 224, 224),
"std": torch.ones(3, 224, 224),
"q01": torch.zeros(3, 224, 224),
"q99": torch.ones(3, 224, 224),
},
}
def test_no_uncond_tokenizer_when_cfg_disabled(self):
from lerobot.processor import TokenizerProcessorStep
config = self._make_config(cfg_beta=1.0)
preprocessor, _ = make_pi05_pre_post_processors(config, self._make_dataset_stats())
tokenizer_steps = [s for s in preprocessor.steps if isinstance(s, TokenizerProcessorStep)]
assert len(tokenizer_steps) == 1
def test_uncond_tokenizer_added_when_cfg_enabled(self):
from lerobot.processor import TokenizerProcessorStep
config = self._make_config(cfg_beta=2.0)
preprocessor, _ = make_pi05_pre_post_processors(config, self._make_dataset_stats())
tokenizer_steps = [s for s in preprocessor.steps if isinstance(s, TokenizerProcessorStep)]
assert len(tokenizer_steps) == 2
uncond_tokenizer = tokenizer_steps[1]
assert uncond_tokenizer.task_key == "uncond_task"
assert uncond_tokenizer.output_tokens_key == OBS_LANGUAGE_UNCOND_TOKENS
assert uncond_tokenizer.output_mask_key == OBS_LANGUAGE_UNCOND_ATTENTION_MASK
def test_cfg_pipeline_produces_both_token_sets(self):
config = self._make_config(cfg_beta=2.0)
preprocessor, _ = make_pi05_pre_post_processors(config, self._make_dataset_stats())
batch = {
"observation.state": torch.randn(14),
"observation.images.base_0_rgb": torch.rand(3, 224, 224),
"task": "pick up the cup",
}
processed = preprocessor(batch)
assert OBS_LANGUAGE_TOKENS in processed
assert OBS_LANGUAGE_ATTENTION_MASK in processed
assert OBS_LANGUAGE_UNCOND_TOKENS in processed
assert OBS_LANGUAGE_UNCOND_ATTENTION_MASK in processed
# Both should be tensors with the same shape
assert processed[OBS_LANGUAGE_TOKENS].shape == processed[OBS_LANGUAGE_UNCOND_TOKENS].shape
assert (
processed[OBS_LANGUAGE_ATTENTION_MASK].shape
== processed[OBS_LANGUAGE_UNCOND_ATTENTION_MASK].shape
)
def test_cfg_beta_1_no_uncond_tokens_in_output(self):
config = self._make_config(cfg_beta=1.0)
preprocessor, _ = make_pi05_pre_post_processors(config, self._make_dataset_stats())
batch = {
"observation.state": torch.randn(14),
"observation.images.base_0_rgb": torch.rand(3, 224, 224),
"task": "pick up the cup",
}
processed = preprocessor(batch)
assert OBS_LANGUAGE_TOKENS in processed
assert OBS_LANGUAGE_UNCOND_TOKENS not in processed
@@ -0,0 +1,186 @@
#!/usr/bin/env python
"""Tests for RenderedMessagesToTaskStep and PI05 pipeline integration with advantage."""
import pytest
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
import torch # noqa: E402
from lerobot.configs.recipe import MessageTurn, TrainingRecipe # noqa: E402
from lerobot.processor.converters import create_transition # noqa: E402
from lerobot.processor.render_messages_processor import RenderMessagesStep # noqa: E402
from lerobot.processor.rendered_messages_to_task import RenderedMessagesToTaskStep # noqa: E402
from lerobot.types import TransitionKey # noqa: E402
def test_rendered_messages_to_task_noops_without_messages():
"""Without messages key, the step is a no-op."""
transition = create_transition(complementary_data={"task": "pick up the cup"})
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["task"] == "pick up the cup"
def test_rendered_messages_to_task_extracts_user_content():
"""Extracts user-role message content and joins with newline."""
transition = create_transition(
complementary_data={
"task": "original task",
"messages": [
{"role": "user", "content": "pick up the cup"},
{"role": "user", "content": "Advantage: positive"},
{"role": "assistant", "content": "reach for cup"},
],
"message_streams": ["high_level", "high_level", "low_level"],
"target_message_indices": [2],
}
)
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["task"] == "pick up the cup\nAdvantage: positive"
assert "messages" not in data
assert "message_streams" not in data
assert "target_message_indices" not in data
def test_rendered_messages_to_task_handles_multimodal_blocks():
"""Extracts text from HF multimodal content blocks."""
transition = create_transition(
complementary_data={
"task": "original",
"messages": [
{
"role": "user",
"content": [
{"type": "image", "image": "placeholder"},
{"type": "text", "text": "describe this"},
],
},
{"role": "assistant", "content": "a cup on a table"},
],
"message_streams": ["high_level", "low_level"],
"target_message_indices": [1],
}
)
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["task"] == "describe this"
def test_rendered_messages_to_task_preserves_list_task_format():
"""When original task is a list (batched), output is also a list."""
transition = create_transition(
complementary_data={
"task": ["task1", "task2"],
"messages": [
{"role": "user", "content": "rendered task"},
{"role": "assistant", "content": "do it", "target": True},
],
"message_streams": ["high_level", "low_level"],
"target_message_indices": [1],
}
)
step = RenderedMessagesToTaskStep()
out = step(transition)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["task"] == ["rendered task", "rendered task"]
def test_full_render_then_flatten_pipeline():
"""RenderMessagesStep + RenderedMessagesToTaskStep produces correct task string."""
recipe = TrainingRecipe(
messages=[
MessageTurn(role="user", content="${task}", stream="high_level"),
MessageTurn(
role="user",
content="Advantage: ${advantage}",
stream="high_level",
if_present="advantage",
),
MessageTurn(role="assistant", content="${subtask}", stream="low_level", target=True),
]
)
transition = create_transition(
complementary_data={
"task": "pick up the cup",
"timestamp": torch.tensor(0.5),
"index": torch.tensor(0),
"language_persistent": [
{
"role": "assistant",
"content": "reach for the cup",
"style": "subtask",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
},
{
"role": "user",
"content": "positive",
"style": "advantage",
"timestamp": 0.1,
"camera": None,
"tool_calls": None,
},
],
"language_events": [],
}
)
# Step 1: Render recipe
rendered = RenderMessagesStep(recipe=recipe)(transition)
# Step 2: Flatten to task string
out = RenderedMessagesToTaskStep()(rendered)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert "pick up the cup" in data["task"]
assert "Advantage: positive" in data["task"]
def test_full_render_advantage_absent_skips_turn():
"""When advantage row is absent, the advantage turn is skipped via if_present."""
recipe = TrainingRecipe(
messages=[
MessageTurn(role="user", content="${task}", stream="high_level"),
MessageTurn(
role="user",
content="Advantage: ${advantage}",
stream="high_level",
if_present="advantage",
),
MessageTurn(role="assistant", content="${subtask}", stream="low_level", target=True),
]
)
transition = create_transition(
complementary_data={
"task": "pick up the cup",
"timestamp": torch.tensor(0.5),
"index": torch.tensor(0),
"language_persistent": [
{
"role": "assistant",
"content": "reach for the cup",
"style": "subtask",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
},
],
"language_events": [],
}
)
rendered = RenderMessagesStep(recipe=recipe)(transition)
out = RenderedMessagesToTaskStep()(rendered)
data = out[TransitionKey.COMPLEMENTARY_DATA]
assert data["task"] == "pick up the cup"
assert "Advantage" not in data["task"]
@@ -0,0 +1,717 @@
# 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.
"""Tests for RECAP's distributional value function."""
from __future__ import annotations
import pytest
import torch
from lerobot.configs.rewards import RewardModelConfig
from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature
from lerobot.rewards.distributional_value_function.configuration_distributional_value_function import (
DistributionalVFConfig,
)
from lerobot.types import TransitionKey
from lerobot.utils.constants import OBS_IMAGES
from tests.utils import skip_if_package_missing
BATCH_SIZE = 4
NUM_BINS = 201
IMAGE_KEY = f"{OBS_IMAGES}.top"
IMAGE_KEY_WRIST_LEFT = f"{OBS_IMAGES}.wrist_left"
IMAGE_KEY_WRIST_RIGHT = f"{OBS_IMAGES}.wrist_right"
def _make_config(**overrides) -> DistributionalVFConfig:
defaults = {
"device": "cpu",
"image_resolution": (224, 224),
}
defaults.update(overrides)
config = DistributionalVFConfig(**defaults)
config.input_features = {
IMAGE_KEY: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
IMAGE_KEY_WRIST_LEFT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
IMAGE_KEY_WRIST_RIGHT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
}
config.output_features = {}
config.normalization_mapping = {
"VISUAL": NormalizationMode.IDENTITY,
}
return config
def _make_model():
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
return DistributionalVFRewardModel(_make_config())
def _make_batch(batch_size: int = BATCH_SIZE, device: str = "cpu") -> dict[str, torch.Tensor]:
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
)
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
return {
IMAGE_KEY: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
IMAGE_KEY + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
IMAGE_KEY_WRIST_LEFT: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
IMAGE_KEY_WRIST_RIGHT: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
IMAGE_KEY_WRIST_RIGHT + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
OBS_LANGUAGE_TOKENS: torch.randint(0, 1000, (batch_size, 200), device=device),
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(batch_size, 200, dtype=torch.bool, device=device),
"mc_return": torch.rand(batch_size, device=device) * -1.0,
"is_terminal": torch.zeros(batch_size, dtype=torch.bool, device=device),
}
# ------------------------------------------------------------------
# Config / registry tests
# ------------------------------------------------------------------
def test_config_registered_in_reward_model_registry():
"""DistributionalVFConfig is discoverable via RewardModelConfig registry."""
known = RewardModelConfig.get_known_choices()
assert "distributional_value_function" in known
def test_factory_returns_correct_class():
"""get_reward_model_class returns DistributionalVFRewardModel."""
from lerobot.rewards.factory import get_reward_model_class
cls = get_reward_model_class("distributional_value_function")
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
assert cls is DistributionalVFRewardModel
def test_make_reward_model_config_factory():
"""make_reward_model_config creates DistributionalVFConfig with overrides."""
from lerobot.rewards.factory import make_reward_model_config
config = make_reward_model_config("distributional_value_function", num_value_bins=101)
assert isinstance(config, DistributionalVFConfig)
assert config.num_value_bins == 101
# ------------------------------------------------------------------
# Target distribution tests (HL-Gauss, Dirac delta, one-hot)
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_hl_gauss_sums_to_one():
"""HL-Gauss target distribution sums to 1 for each sample."""
model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9, -0.0])
dist = model.hl_gauss_target(targets)
assert dist.shape == (4, NUM_BINS)
torch.testing.assert_close(dist.sum(dim=-1), torch.ones(4), atol=1e-5, rtol=0)
@skip_if_package_missing("transformers")
def test_hl_gauss_non_negative():
"""HL-Gauss target probabilities are all non-negative."""
model = _make_model()
targets = torch.linspace(-1.0, 0.0, 10)
dist = model.hl_gauss_target(targets)
assert (dist >= 0).all()
@skip_if_package_missing("transformers")
def test_hl_gauss_expected_value_matches():
"""E[V] under HL-Gauss distribution matches the target value."""
model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9])
dist = model.hl_gauss_target(targets)
expected = (dist * model.value_head.bin_centers).sum(dim=-1)
torch.testing.assert_close(expected, targets, atol=1e-4, rtol=0)
@skip_if_package_missing("transformers")
def test_hl_gauss_handles_2d_input():
"""HL-Gauss handles [batch_size, 1] shaped inputs correctly."""
model = _make_model()
targets = torch.tensor([-0.5, -0.3]).unsqueeze(-1)
dist = model.hl_gauss_target(targets)
assert dist.shape == (2, NUM_BINS)
torch.testing.assert_close(dist.sum(dim=-1), torch.ones(2), atol=1e-5, rtol=0)
@skip_if_package_missing("transformers")
def test_dirac_delta_sums_to_one():
"""Dirac delta target distribution sums to 1 for each sample."""
model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9, -1.0, 0.0])
dist = model.dirac_delta_target(targets)
assert dist.shape == (5, NUM_BINS)
torch.testing.assert_close(dist.sum(dim=-1), torch.ones(5), atol=1e-6, rtol=0)
@skip_if_package_missing("transformers")
def test_dirac_delta_at_most_two_nonzero():
"""Dirac delta places probability on at most two adjacent bins."""
model = _make_model()
targets = torch.tensor([-0.7523, -0.0013])
dist = model.dirac_delta_target(targets)
for i in range(2):
assert (dist[i] > 0).sum() <= 2
@skip_if_package_missing("transformers")
def test_dirac_delta_expected_value_matches():
"""E[V] under Dirac delta distribution matches the target value."""
model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9])
dist = model.dirac_delta_target(targets)
expected = (dist * model.value_head.bin_centers).sum(dim=-1)
torch.testing.assert_close(expected, targets, atol=1e-5, rtol=0)
@skip_if_package_missing("transformers")
def test_dirac_delta_boundary_values_clamped():
"""Values outside support are clamped to boundary bins."""
model = _make_model()
targets = torch.tensor([-1.5, 0.5])
dist = model.dirac_delta_target(targets)
torch.testing.assert_close(dist.sum(dim=-1), torch.ones(2), atol=1e-6, rtol=0)
assert dist[0, 0] == 1.0
assert dist[1, -1] == 1.0
@skip_if_package_missing("transformers")
def test_one_hot_single_nonzero():
"""One-hot target has exactly one non-zero bin per sample."""
model = _make_model()
targets = torch.tensor([-0.5, -0.1, -1.0, 0.0])
dist = model.one_hot_target(targets)
assert dist.shape == (4, NUM_BINS)
for i in range(4):
assert (dist[i] > 0).sum() == 1
assert dist[i].sum() == 1.0
@skip_if_package_missing("transformers")
def test_one_hot_nearest_bin():
"""One-hot target activates the bin closest to the target value."""
model = _make_model()
targets = torch.tensor([-0.5])
dist = model.one_hot_target(targets)
hot_idx = dist[0].argmax()
assert model.value_head.bin_centers[hot_idx].item() == pytest.approx(-0.5, abs=0.003)
@skip_if_package_missing("transformers")
def test_terminal_gets_one_hot():
"""Terminal states receive one-hot targets; non-terminal get HL-Gauss."""
model = _make_model()
targets = torch.tensor([-0.5, -0.3, -0.7, -0.9])
is_terminal = torch.tensor([False, True, False, True])
dist = model.compute_target_distribution(
targets, is_terminal, method="hl_gauss", use_one_hot_terminal=True
)
for i in range(4):
assert dist[i].sum().item() == pytest.approx(1.0, abs=1e-5)
assert (dist[1] > 0).sum() == 1
assert (dist[3] > 0).sum() == 1
assert (dist[0] > 0).sum() > 2
assert (dist[2] > 0).sum() > 2
@skip_if_package_missing("transformers")
def test_no_terminal_override_when_disabled():
"""When use_one_hot_terminal=False, terminal states use the base method."""
model = _make_model()
targets = torch.tensor([-0.5, -0.3])
is_terminal = torch.tensor([False, True])
dist = model.compute_target_distribution(
targets, is_terminal, method="hl_gauss", use_one_hot_terminal=False
)
assert (dist[1] > 0).sum() > 2
# ------------------------------------------------------------------
# Architecture / component tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_model_has_expected_components():
"""Model scaffold contains the SigLIP2+Gemma3+ValueHead components."""
model = _make_model()
assert hasattr(model, "vision_encoder")
assert hasattr(model, "gemma3")
assert hasattr(model, "image_proj")
assert hasattr(model, "value_head")
assert hasattr(model, "cls_embedding")
assert hasattr(model.value_head, "mlp")
assert hasattr(model.value_head, "bin_centers")
@skip_if_package_missing("transformers")
def test_model_bin_centers_shape():
"""Value head bin_centers buffer has shape (num_value_bins,)."""
model = _make_model()
assert model.value_head.bin_centers.shape == (NUM_BINS,)
@skip_if_package_missing("transformers")
def test_value_head_output_dim():
"""Value head linear projection outputs num_value_bins logits."""
model = _make_model()
assert model.value_head.mlp[-1].out_features == NUM_BINS
@skip_if_package_missing("transformers")
def test_cls_embedding_is_nn_embedding():
"""CLS is nn.Embedding (FSDP-safe) with correct shape."""
model = _make_model()
from torch import nn
assert isinstance(model.cls_embedding, nn.Embedding)
assert model.cls_embedding.num_embeddings == 1
assert model.cls_embedding.embedding_dim == model.gemma3_hidden
@skip_if_package_missing("transformers")
def test_image_proj_dimensions():
"""Image projection maps SigLIP2 hidden to Gemma3 hidden."""
model = _make_model()
siglip_hidden = model.vision_encoder.config.hidden_size
assert model.image_proj.in_features == siglip_hidden
assert model.image_proj.out_features == model.gemma3_hidden
# ------------------------------------------------------------------
# Forward / inference tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_forward_returns_loss_and_dict():
"""Forward pass returns a finite scalar loss and output dict with expected keys."""
model = _make_model()
batch = _make_batch()
loss, output_dict = model.forward(batch)
assert loss.shape == ()
assert torch.isfinite(loss)
assert "loss" in output_dict
assert "predicted_value_mean" in output_dict
assert "mc_return_mean" in output_dict
assert "acc_best" in output_dict
assert "acc_neighbor" in output_dict
assert "mae" in output_dict
@skip_if_package_missing("transformers")
def test_forward_loss_is_positive():
"""Cross-entropy loss is strictly positive for random weights."""
model = _make_model()
batch = _make_batch()
loss, _ = model.forward(batch)
assert loss.item() > 0
@skip_if_package_missing("transformers")
def test_compute_reward_returns_correct_shape():
"""compute_reward returns [batch_size] tensor of finite float32 values."""
model = _make_model()
model.eval()
batch = _make_batch(batch_size=3)
with torch.no_grad():
values = model.compute_reward(batch)
assert values.shape == (3,)
assert values.dtype == torch.float32
assert torch.isfinite(values).all()
@skip_if_package_missing("transformers")
def test_compute_reward_values_in_support_range():
"""Predicted values lie within [value_support_min, value_support_max]."""
model = _make_model()
model.eval()
batch = _make_batch(batch_size=8)
with torch.no_grad():
values = model.compute_reward(batch)
assert (values >= -1.0 - 0.01).all()
assert (values <= 0.0 + 0.01).all()
# ------------------------------------------------------------------
# Gradient flow tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_gradient_flows_through_value_head():
"""Backprop produces non-zero gradients on the value head projection."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.value_head.mlp[-1].weight.grad is not None
assert not torch.all(model.value_head.mlp[-1].weight.grad == 0)
@skip_if_package_missing("transformers")
def test_gradient_flows_through_cls_embedding():
"""Backprop produces non-zero gradients on the learned [CLS] embedding."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.cls_embedding.weight.grad is not None
assert not torch.all(model.cls_embedding.weight.grad == 0)
@skip_if_package_missing("transformers")
def test_gradient_flows_through_image_proj():
"""Backprop produces non-zero gradients on the image projection."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.image_proj.weight.grad is not None
assert not torch.all(model.image_proj.weight.grad == 0)
# ------------------------------------------------------------------
# Freeze / training infrastructure tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_freeze_vision_encoder():
"""freeze_vision_encoder disables requires_grad on SigLIP2."""
model = _make_model()
model.config.freeze_vision_encoder = True
model._set_requires_grad()
for p in model.vision_encoder.parameters():
assert not p.requires_grad
for p in model.value_head.parameters():
assert p.requires_grad
@skip_if_package_missing("transformers")
def test_freeze_language_model():
"""freeze_language_model disables requires_grad on Gemma3."""
model = _make_model()
model.config.freeze_language_model = True
model._set_requires_grad()
for p in model.gemma3.parameters():
assert not p.requires_grad
for p in model.value_head.parameters():
assert p.requires_grad
@skip_if_package_missing("transformers")
def test_stop_gradient_to_vlm_preserves_cls_grad():
"""With stop_gradient_to_vlm, CLS embedding still gets gradients."""
config = _make_config(stop_gradient_to_vlm=True)
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = DistributionalVFRewardModel(config)
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.cls_embedding.weight.grad is not None
assert not torch.all(model.cls_embedding.weight.grad == 0)
# ------------------------------------------------------------------
# Config validation tests
# ------------------------------------------------------------------
def test_config_requires_visual_feature():
"""validate_features raises if no VISUAL feature is present."""
config = DistributionalVFConfig()
config.input_features = {
"observation.state": PolicyFeature(type=FeatureType.STATE, shape=(14,)),
}
with pytest.raises(ValueError, match="VISUAL"):
config.validate_features()
def test_config_passes_with_visual_feature():
"""validate_features succeeds when a VISUAL feature is present."""
config = _make_config()
config.validate_features()
# ------------------------------------------------------------------
# Processor tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_processor_pipeline_produces_expected_keys():
"""Full preprocessor pipeline produces tokenized text, preprocessed images, and masks."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
make_distributional_vf_pre_post_processors,
)
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
config = _make_config()
preprocessor, _ = make_distributional_vf_pre_post_processors(config)
raw_batch = {
IMAGE_KEY: torch.rand(3, 224, 224),
IMAGE_KEY_WRIST_LEFT: torch.rand(3, 224, 224),
IMAGE_KEY_WRIST_RIGHT: torch.rand(3, 224, 224),
"task": "pick up the cup",
}
processed = preprocessor(raw_batch)
assert OBS_LANGUAGE_TOKENS in processed
assert OBS_LANGUAGE_ATTENTION_MASK in processed
assert IMAGE_KEY in processed
assert IMAGE_KEY + IMAGE_MASK_SUFFIX in processed
assert IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX in processed
assert IMAGE_KEY_WRIST_RIGHT + IMAGE_MASK_SUFFIX in processed
img = processed[IMAGE_KEY]
assert img.shape == (1, 3, 224, 224)
assert img.min() >= -1.0 - 1e-5
assert img.max() <= 1.0 + 1e-5
def test_task_prompt_formats_correctly():
"""Task prompt step builds 'Task: {task}.' format."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFPrepareTaskPromptStep,
)
step = DistributionalVFPrepareTaskPromptStep()
transition = {
TransitionKey.COMPLEMENTARY_DATA: {"task": ["pick_up_the_cup"]},
}
result = step(transition)
prompt = result[TransitionKey.COMPLEMENTARY_DATA]["task"][0]
assert prompt == "Task: pick up the cup."
def test_task_prompt_handles_string_input():
"""Task prompt step accepts a plain string task."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFPrepareTaskPromptStep,
)
step = DistributionalVFPrepareTaskPromptStep()
transition = {
TransitionKey.COMPLEMENTARY_DATA: {"task": "open_drawer"},
}
result = step(transition)
prompt = result[TransitionKey.COMPLEMENTARY_DATA]["task"][0]
assert prompt == "Task: open drawer."
def test_task_prompt_raises_on_missing_task():
"""Task prompt step raises ValueError when task key is absent."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFPrepareTaskPromptStep,
)
step = DistributionalVFPrepareTaskPromptStep()
transition = {
TransitionKey.COMPLEMENTARY_DATA: {},
}
with pytest.raises(ValueError, match="No task found"):
step(transition)
def test_image_preprocessor_resize_and_normalize():
"""Image preprocessor resizes, normalizes to [-1,1], and adds masks."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
DistributionalVFImagePreprocessorStep,
)
step = DistributionalVFImagePreprocessorStep(
image_resolution=(224, 224),
image_keys=(IMAGE_KEY,),
)
transition = {
TransitionKey.OBSERVATION: {
IMAGE_KEY: torch.rand(2, 3, 320, 240), # non-square, [0, 1]
}
}
result = step(transition)
obs = result[TransitionKey.OBSERVATION]
assert obs[IMAGE_KEY].shape == (2, 3, 224, 224)
assert obs[IMAGE_KEY].min() >= -1.0 - 1e-5
assert obs[IMAGE_KEY].max() <= 1.0 + 1e-5
assert obs[IMAGE_KEY + IMAGE_MASK_SUFFIX].all()
def test_image_preprocessor_missing_camera_gets_placeholder():
"""Missing cameras get black placeholder and mask=False."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
DistributionalVFImagePreprocessorStep,
)
step = DistributionalVFImagePreprocessorStep(
image_resolution=(224, 224),
image_keys=(IMAGE_KEY, IMAGE_KEY_WRIST_LEFT),
)
transition = {
TransitionKey.OBSERVATION: {
IMAGE_KEY: torch.rand(2, 3, 224, 224),
}
}
result = step(transition)
obs = result[TransitionKey.OBSERVATION]
assert obs[IMAGE_KEY + IMAGE_MASK_SUFFIX].all()
assert not obs[IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX].any()
assert obs[IMAGE_KEY_WRIST_LEFT].shape == (2, 3, 224, 224)
# ------------------------------------------------------------------
# Save / load roundtrip
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_save_load_pretrained_roundtrip(tmp_path):
"""Saved model can be loaded back with identical weights."""
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = _make_model()
model._save_pretrained(tmp_path)
loaded = DistributionalVFRewardModel.from_pretrained(str(tmp_path))
orig_sd = model.state_dict()
loaded_sd = loaded.state_dict()
assert set(orig_sd.keys()) == set(loaded_sd.keys())
for key in orig_sd:
torch.testing.assert_close(orig_sd[key], loaded_sd[key], msg=f"Mismatch in {key}")
# ------------------------------------------------------------------
# Attention mask utility test
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_make_att_2d_masks():
"""Verify attention mask construction for prefix + CLS."""
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
make_att_2d_masks,
)
pad = torch.ones(1, 4, dtype=torch.bool)
att = torch.tensor([[0, 0, 0, 1]])
mask = make_att_2d_masks(pad, att)[0]
assert mask[0, 0] # prefix sees prefix
assert mask[1, 2] # prefix sees prefix
assert not mask[0, 3] # prefix does NOT see CLS
assert mask[3, 0] # CLS sees prefix
assert mask[3, 3] # CLS sees itself
# ------------------------------------------------------------------
# Categorical metrics test
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_categorical_metrics_perfect_prediction():
"""Metrics return acc_best=1 when logits peak at the correct bin."""
model = _make_model()
bin_centers = model.value_head.bin_centers
target = bin_centers[100].unsqueeze(0) # exact bin center
batch = _make_batch(batch_size=1)
batch["mc_return"] = target
batch["is_terminal"] = torch.zeros(1, dtype=torch.bool)
with torch.no_grad():
_, output_dict = model.forward(batch)
assert "acc_best" in output_dict
assert "acc_neighbor" in output_dict
assert "mae" in output_dict
assert isinstance(output_dict["acc_best"], float)
assert isinstance(output_dict["mae"], float)
+514
View File
@@ -0,0 +1,514 @@
#!/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.
# 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.
"""Tests for lerobot-compute-returns script."""
import json
from pathlib import Path
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
from lerobot.scripts.lerobot_compute_returns import (
IS_TERMINAL_COL,
MC_RETURN_COL,
ComputeReturnsConfig,
_get_episode_success,
compute_episode_returns,
)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def parquet_dataset(tmp_path):
"""Build a minimal parquet shard + info.json for testing I/O logic.
Mirrors the lerobot-rollout DAgger convention: ``next.success`` is False
on all frames except the terminal frame of successful episodes.
Even episodes are successful, odd episodes are failures.
"""
num_episodes = 3
frames_per_ep = 10
root = tmp_path / "test_dataset"
data_dir = root / "data" / "chunk-000"
meta_dir = root / "meta"
data_dir.mkdir(parents=True)
meta_dir.mkdir(parents=True)
all_rows = []
episodes_meta = []
global_idx = 0
for ep in range(num_episodes):
ep_from = global_idx
is_successful = ep % 2 == 0
for frame in range(frames_per_ep):
is_last_frame = frame == frames_per_ep - 1
all_rows.append(
{
"episode_index": ep,
"frame_index": frame,
"index": global_idx,
"next.success": is_successful and is_last_frame,
}
)
global_idx += 1
ep_to = global_idx
episodes_meta.append(
{
"episode_index": ep,
"length": frames_per_ep,
"dataset_from_index": ep_from,
"dataset_to_index": ep_to,
}
)
table = pa.table(
{
"episode_index": [r["episode_index"] for r in all_rows],
"frame_index": [r["frame_index"] for r in all_rows],
"index": [r["index"] for r in all_rows],
"next.success": [r["next.success"] for r in all_rows],
}
)
parquet_path = data_dir / "episode_000000.parquet"
pq.write_table(table, parquet_path)
info = {
"codebase_version": "v3.0",
"total_episodes": num_episodes,
"total_frames": global_idx,
"fps": 30,
"features": {
"episode_index": {"dtype": "int64", "shape": [1], "names": None},
"frame_index": {"dtype": "int64", "shape": [1], "names": None},
"index": {"dtype": "int64", "shape": [1], "names": None},
"next.success": {"dtype": "bool", "shape": [1], "names": None},
},
}
(meta_dir / "info.json").write_text(json.dumps(info, indent=2))
return root, parquet_path, episodes_meta
def _rewrite_shard(parquet_path: Path, episodes_meta: list[dict], config: ComputeReturnsConfig):
"""Rewrite a single parquet shard using the core logic from compute_returns."""
table = pq.read_table(parquet_path)
if not config.force and IS_TERMINAL_COL in table.column_names:
return
all_is_terminal = np.zeros(len(table), dtype=bool)
all_mc_return = np.zeros(len(table), dtype=np.float32)
episode_col = table.column("episode_index").to_pylist()
for ep_info in episodes_meta:
ep_idx = ep_info["episode_index"]
ep_len = ep_info["length"]
mask = np.array([v == ep_idx for v in episode_col], dtype=bool)
local_indices = np.where(mask)[0]
ep_subtable = table.filter(mask)
success = _get_episode_success(ep_subtable, config.success_key, config.default_success)
is_terminal, mc_return = compute_episode_returns(
num_frames=ep_len,
success=success,
c_fail=config.c_fail,
gamma=config.gamma,
max_episode_length=config.max_episode_length or ep_len,
)
all_is_terminal[local_indices] = is_terminal
all_mc_return[local_indices] = mc_return
if IS_TERMINAL_COL in table.column_names:
table = table.drop(IS_TERMINAL_COL)
if MC_RETURN_COL in table.column_names:
table = table.drop(MC_RETURN_COL)
table = table.append_column(IS_TERMINAL_COL, pa.array(all_is_terminal))
table = table.append_column(MC_RETURN_COL, pa.array(all_mc_return))
pq.write_table(table, parquet_path)
# ---------------------------------------------------------------------------
# Tests: compute_episode_returns (pure math, no I/O)
# ---------------------------------------------------------------------------
def test_successful_episode_terminal_reward_is_zero():
"""Terminal MC return for a successful episode should be 0."""
_, mc_return = compute_episode_returns(
num_frames=10, success=True, c_fail=50.0, gamma=1.0, max_episode_length=10
)
assert mc_return[-1] == pytest.approx(0.0, abs=1e-6)
def test_failed_episode_terminal_reward_reflects_cfail():
"""Terminal MC return for a failed episode should be -C_fail / H."""
horizon = 100
c_fail = 50.0
_, mc_return = compute_episode_returns(
num_frames=10, success=False, c_fail=c_fail, gamma=1.0, max_episode_length=horizon
)
assert mc_return[-1] == pytest.approx(-c_fail / horizon, abs=1e-5)
def test_is_terminal_only_last_frame():
"""Only the last frame of an episode should be marked terminal."""
is_terminal, _ = compute_episode_returns(
num_frames=20, success=True, c_fail=50.0, gamma=1.0, max_episode_length=20
)
assert is_terminal[-1] == True # noqa: E712
assert not any(is_terminal[:-1])
def test_mc_return_monotonically_increases_for_success():
"""For a successful undiscounted episode, returns should increase toward 0."""
_, mc_return = compute_episode_returns(
num_frames=50, success=True, c_fail=50.0, gamma=1.0, max_episode_length=50
)
for i in range(len(mc_return) - 1):
assert mc_return[i] <= mc_return[i + 1]
def test_mc_return_bounded_negative_to_zero():
"""MC returns for successful episodes should be in (-1, 0]."""
_, mc_return = compute_episode_returns(
num_frames=100, success=True, c_fail=50.0, gamma=1.0, max_episode_length=100
)
assert mc_return[-1] == pytest.approx(0.0, abs=1e-6)
assert all(v <= 0.0 for v in mc_return)
assert all(v >= -1.0 - 1e-6 for v in mc_return)
def test_first_frame_return_success():
"""First frame return for successful episode equals -(N-1)/H."""
num_frames = 10
horizon = 10
_, mc_return = compute_episode_returns(
num_frames=num_frames, success=True, c_fail=50.0, gamma=1.0, max_episode_length=horizon
)
expected = -(num_frames - 1) / horizon
assert mc_return[0] == pytest.approx(expected, abs=1e-5)
def test_first_frame_return_failure():
"""First frame return for failed episode includes the failure penalty."""
num_frames = 10
horizon = 100
c_fail = 50.0
_, mc_return = compute_episode_returns(
num_frames=num_frames, success=False, c_fail=c_fail, gamma=1.0, max_episode_length=horizon
)
expected = (-(num_frames - 1) / horizon) + (-c_fail / horizon)
assert mc_return[0] == pytest.approx(expected, abs=1e-5)
def test_discount_factor_less_than_one():
"""Discount factor < 1 should make earlier frames have smaller magnitude."""
_, mc_undiscounted = compute_episode_returns(
num_frames=20, success=True, c_fail=50.0, gamma=1.0, max_episode_length=20
)
_, mc_discounted = compute_episode_returns(
num_frames=20, success=True, c_fail=50.0, gamma=0.99, max_episode_length=20
)
assert abs(mc_discounted[0]) < abs(mc_undiscounted[0])
def test_single_frame_episode_success():
"""Single-frame successful episode: return should be 0."""
is_terminal, mc_return = compute_episode_returns(
num_frames=1, success=True, c_fail=50.0, gamma=1.0, max_episode_length=1
)
assert mc_return[0] == pytest.approx(0.0, abs=1e-6)
assert is_terminal[0] == True # noqa: E712
def test_single_frame_episode_failure():
"""Single-frame failed episode: return should be -C_fail/H."""
horizon = 100
c_fail = 50.0
is_terminal, mc_return = compute_episode_returns(
num_frames=1, success=False, c_fail=c_fail, gamma=1.0, max_episode_length=horizon
)
assert mc_return[0] == pytest.approx(-c_fail / horizon, abs=1e-5)
assert is_terminal[0] == True # noqa: E712
def test_horizon_normalization_scales_returns():
"""Larger horizon should scale down the per-step penalty."""
_, mc_small_h = compute_episode_returns(
num_frames=10, success=True, c_fail=50.0, gamma=1.0, max_episode_length=10
)
_, mc_large_h = compute_episode_returns(
num_frames=10, success=True, c_fail=50.0, gamma=1.0, max_episode_length=100
)
assert abs(mc_large_h[0]) < abs(mc_small_h[0])
# ---------------------------------------------------------------------------
# Tests: _get_episode_success (in-memory PyArrow tables)
# ---------------------------------------------------------------------------
def test_default_success_overrides_column():
"""default_success should override any column value."""
table = pa.table({"next.success": [True, True, True]})
assert _get_episode_success(table, "next.success", default_success=False) is False
def test_reads_bool_column():
"""Should detect success via any() reduction over the column."""
table_success = pa.table({"next.success": [False, False, True]})
table_fail = pa.table({"next.success": [False, False, False]})
assert _get_episode_success(table_success, "next.success", None) is True
assert _get_episode_success(table_fail, "next.success", None) is False
def test_reads_int_column():
"""Should interpret integer success column (0/1) as bool via any()."""
table = pa.table({"task_success": [0, 0, 1]})
assert _get_episode_success(table, "task_success", None) is True
def test_all_zeros_means_failure():
"""An episode with all-zero success values is a failure."""
table = pa.table({"next.success": [0, 0, 0]})
assert _get_episode_success(table, "next.success", None) is False
def test_missing_column_defaults_to_true():
"""When success column is missing, assume success (demo data)."""
table = pa.table({"frame_index": [0, 1, 2]})
assert _get_episode_success(table, "next.success", None) is True
# ---------------------------------------------------------------------------
# Tests: parquet rewriting (integration, writes to disk)
# ---------------------------------------------------------------------------
def test_writes_columns_to_parquet(parquet_dataset):
"""The rewrite logic should add is_terminal and mc_return columns."""
root, parquet_path, episodes_meta = parquet_dataset
table_before = pq.read_table(parquet_path)
assert IS_TERMINAL_COL not in table_before.column_names
assert MC_RETURN_COL not in table_before.column_names
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=10, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table_after = pq.read_table(parquet_path)
assert IS_TERMINAL_COL in table_after.column_names
assert MC_RETURN_COL in table_after.column_names
def test_terminal_frames_correct(parquet_dataset):
"""Only the last frame of each episode should be terminal."""
root, parquet_path, episodes_meta = parquet_dataset
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=10, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table = pq.read_table(parquet_path)
is_terminal = table.column(IS_TERMINAL_COL).to_pylist()
terminal_indices = [i for i, v in enumerate(is_terminal) if v]
assert terminal_indices == [9, 19, 29]
def test_success_episodes_return_zero_at_terminal(tmp_path):
"""Successful episodes (ep 0) should have mc_return=0 at terminal."""
num_episodes = 2
frames_per_ep = 5
root = tmp_path / "test_dataset"
data_dir = root / "data" / "chunk-000"
meta_dir = root / "meta"
data_dir.mkdir(parents=True)
meta_dir.mkdir(parents=True)
all_rows = []
episodes_meta = []
global_idx = 0
for ep in range(num_episodes):
ep_from = global_idx
is_successful = ep % 2 == 0
for frame in range(frames_per_ep):
is_last_frame = frame == frames_per_ep - 1
all_rows.append(
{
"episode_index": ep,
"frame_index": frame,
"index": global_idx,
"next.success": is_successful and is_last_frame,
}
)
global_idx += 1
episodes_meta.append(
{
"episode_index": ep,
"length": frames_per_ep,
"dataset_from_index": ep_from,
"dataset_to_index": global_idx,
}
)
table = pa.table(
{
"episode_index": [r["episode_index"] for r in all_rows],
"frame_index": [r["frame_index"] for r in all_rows],
"index": [r["index"] for r in all_rows],
"next.success": [r["next.success"] for r in all_rows],
}
)
parquet_path = data_dir / "episode_000000.parquet"
pq.write_table(table, parquet_path)
info = {
"codebase_version": "v3.0",
"total_episodes": num_episodes,
"total_frames": global_idx,
"fps": 30,
"features": {
"episode_index": {"dtype": "int64", "shape": [1], "names": None},
"frame_index": {"dtype": "int64", "shape": [1], "names": None},
"index": {"dtype": "int64", "shape": [1], "names": None},
"next.success": {"dtype": "bool", "shape": [1], "names": None},
},
}
(meta_dir / "info.json").write_text(json.dumps(info, indent=2))
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=5, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table = pq.read_table(parquet_path)
mc_return = table.column(MC_RETURN_COL).to_pylist()
assert mc_return[4] == pytest.approx(0.0, abs=1e-5)
def test_failed_episodes_have_negative_terminal(tmp_path):
"""Failed episodes (ep 1) should have mc_return < 0 at terminal."""
num_episodes = 2
frames_per_ep = 5
root = tmp_path / "test_dataset"
data_dir = root / "data" / "chunk-000"
meta_dir = root / "meta"
data_dir.mkdir(parents=True)
meta_dir.mkdir(parents=True)
all_rows = []
episodes_meta = []
global_idx = 0
for ep in range(num_episodes):
ep_from = global_idx
is_successful = ep % 2 == 0
for frame in range(frames_per_ep):
is_last_frame = frame == frames_per_ep - 1
all_rows.append(
{
"episode_index": ep,
"frame_index": frame,
"index": global_idx,
"next.success": is_successful and is_last_frame,
}
)
global_idx += 1
episodes_meta.append(
{
"episode_index": ep,
"length": frames_per_ep,
"dataset_from_index": ep_from,
"dataset_to_index": global_idx,
}
)
table = pa.table(
{
"episode_index": [r["episode_index"] for r in all_rows],
"frame_index": [r["frame_index"] for r in all_rows],
"index": [r["index"] for r in all_rows],
"next.success": [r["next.success"] for r in all_rows],
}
)
parquet_path = data_dir / "episode_000000.parquet"
pq.write_table(table, parquet_path)
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=5, c_fail=50.0, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table = pq.read_table(parquet_path)
mc_return = table.column(MC_RETURN_COL).to_pylist()
assert mc_return[9] < 0.0
def test_idempotent_with_force_flag(parquet_dataset):
"""Running twice with force should produce identical results."""
root, parquet_path, episodes_meta = parquet_dataset
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=10, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table1 = pq.read_table(parquet_path)
mc1 = table1.column(MC_RETURN_COL).to_pylist()
_rewrite_shard(parquet_path, episodes_meta, config)
table2 = pq.read_table(parquet_path)
mc2 = table2.column(MC_RETURN_COL).to_pylist()
assert mc1 == mc2
def test_skips_if_columns_exist_without_force(parquet_dataset):
"""Without force, existing columns should not be overwritten."""
root, parquet_path, episodes_meta = parquet_dataset
config = ComputeReturnsConfig(success_key="next.success", max_episode_length=10, force=True)
_rewrite_shard(parquet_path, episodes_meta, config)
table = pq.read_table(parquet_path)
original_mc = table.column(MC_RETURN_COL).to_pylist()
config_no_force = ComputeReturnsConfig(success_key="next.success", max_episode_length=20, force=False)
_rewrite_shard(parquet_path, episodes_meta, config_no_force)
table2 = pq.read_table(parquet_path)
assert table2.column(MC_RETURN_COL).to_pylist() == original_mc
def test_updates_info_json(parquet_dataset):
"""info.json should be updated with is_terminal and mc_return features."""
from lerobot.scripts.lerobot_compute_returns import _update_info_json
root, parquet_path, episodes_meta = parquet_dataset
_update_info_json(root, None)
info_path = root / "meta" / "info.json"
info = json.loads(info_path.read_text())
assert IS_TERMINAL_COL in info["features"]
assert MC_RETURN_COL in info["features"]
assert info["features"][IS_TERMINAL_COL]["dtype"] == "bool"
assert info["features"][MC_RETURN_COL]["dtype"] == "float32"
+97
View File
@@ -338,6 +338,103 @@ def test_dagger_events_reset():
assert not events.upload_requested.is_set()
def test_dagger_mark_success():
"""mark_success sets the episode label to True."""
from lerobot.rollout.strategies import DAggerEvents
events = DAggerEvents()
assert events.consume_episode_success() is None
events.mark_success()
assert events.consume_episode_success() is True
# Consuming clears the label
assert events.consume_episode_success() is None
def test_dagger_mark_failure():
"""mark_failure sets the episode label to False."""
from lerobot.rollout.strategies import DAggerEvents
events = DAggerEvents()
events.mark_failure()
assert events.consume_episode_success() is False
def test_dagger_success_overrides_failure():
"""Last label wins — success after failure overrides."""
from lerobot.rollout.strategies import DAggerEvents
events = DAggerEvents()
events.mark_failure()
events.mark_success()
assert events.consume_episode_success() is True
def test_dagger_reset_clears_success_label():
"""reset() clears any pending episode success label."""
from lerobot.rollout.strategies import DAggerEvents
events = DAggerEvents()
events.mark_success()
events.reset()
assert events.consume_episode_success() is None
def test_stamp_episode_success_labels_terminal_frame():
"""_stamp_episode_success sets last frame's next.success to True."""
import numpy as np
from lerobot.rollout.strategies.dagger import DAggerStrategy
strategy = DAggerStrategy.__new__(DAggerStrategy)
strategy.config = MagicMock()
from lerobot.rollout.strategies import DAggerEvents
strategy._events = DAggerEvents()
strategy._events.mark_success()
dataset = MagicMock()
dataset.writer.episode_buffer = {
"next.success": [
np.array([False], dtype=bool),
np.array([False], dtype=bool),
np.array([False], dtype=bool),
],
}
strategy._stamp_episode_success(dataset)
assert dataset.writer.episode_buffer["next.success"][-1].item() is True
assert dataset.writer.episode_buffer["next.success"][0].item() is False
def test_stamp_episode_success_no_label_stays_false():
"""Without a label, all frames remain False."""
import numpy as np
from lerobot.rollout.strategies.dagger import DAggerStrategy
strategy = DAggerStrategy.__new__(DAggerStrategy)
strategy.config = MagicMock()
from lerobot.rollout.strategies import DAggerEvents
strategy._events = DAggerEvents()
dataset = MagicMock()
dataset.writer.episode_buffer = {
"next.success": [
np.array([False], dtype=bool),
np.array([False], dtype=bool),
],
}
strategy._stamp_episode_success(dataset)
assert all(v.item() is False for v in dataset.writer.episode_buffer["next.success"])
# ---------------------------------------------------------------------------
# Context dataclass
# ---------------------------------------------------------------------------
Generated
+1243 -1483
View File
File diff suppressed because it is too large Load Diff