Compare commits

..

46 Commits

Author SHA1 Message Date
Khalil Meftah 9947afa0db fix ux 2026-07-29 15:39:02 +02:00
Khalil Meftah 2e95c57548 fix video ui 2026-07-29 15:34:13 +02:00
Khalil Meftah 4e423dcf2e fix video encoding/opencv 2026-07-29 15:23:40 +02:00
Khalil Meftah d2173eed14 create video overlay for advantage 2026-07-29 15:17:47 +02:00
Khalil Meftah 8a779f1e9a eval(rewards): evaluate value function 2026-07-26 20:09:53 +02:00
Khalil Meftah 64613d8f03 fix(rewards): harden distributional value target projection
- Clamp HL-Gauss targets
- Handle terminal mask shapes safely
- Align the default smoothing ratio with Stop Regressing
2026-07-25 17:17:35 +02:00
Khalil Meftah b162e977c0 fix: six frames history 2026-07-24 09:31:44 +02:00
Khalil Meftah 0a342ad4d0 fix: exclude _is_pad mask from images 2026-07-24 09:28:04 +02:00
Khalil Meftah 39cdc4dea3 fix: nan in eval 2026-07-24 09:24:39 +02:00
Khalil Meftah 02f67a9f54 fix: nanoVLM processing 2026-07-23 22:50:29 +02:00
Khalil Meftah 08953c3a9e fix: add attention_mask 2026-07-23 22:29:22 +02:00
Khalil Meftah 39284f43c7 fix: populate config from dataset metadata 2026-07-23 22:17:39 +02:00
Khalil Meftah d015474483 fix temporal siglip value function 2026-07-23 21:55:24 +02:00
Khalil Meftah 8ff4a96b5e feat(rewards): add temporal SigLIP2 and nanoVLM value functions
- add isolated configs, models, processors, and tests
- add shared distributional value utilities
- add VF overfit comparison harness
- add standalone Gemma3 VLM alignment workflow
- keep the committed RECAP baseline unchanged
2026-07-23 21:24:52 +02:00
Khalil Meftah 2aa7f601cd feat(recap): align value function with pi0.6 Gemma3 VLM
- use 448px SigLIP2 inputs with 2x2 pooling
- adopt native Gemma3 projector and attention masking
- support aligned VLM checkpoints
- update preprocessing, targets, and tests
2026-07-23 13:38:07 +02:00
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
69 changed files with 9285 additions and 1572 deletions
@@ -0,0 +1,119 @@
# RECAP value-function experiments
All variants use the same `mc_return`, `is_terminal`, 201-bin support, Dirac/HL-Gauss
targets, metrics, and LeRobot training pipeline. Only the representation backbone changes.
## Current RECAP Gemma3 baseline
```bash
lerobot-train \
--reward_model.type=distributional_value_function \
--reward_model.target_method=dirac_delta \
--reward_model.device=cuda \
--dataset.repo_id=<dataset_repo_id> \
--output_dir=outputs/vf_recap_gemma3 \
--steps=40000 \
--batch_size=8
```
This initializes SigLIP2 and Gemma3-270M from unimodal checkpoints and creates a
fresh Gemma3 multimodal connector.
## Temporal SigLIP2
```bash
lerobot-train \
--reward_model.type=temporal_siglip_value_function \
--reward_model.history_steps=6 \
--reward_model.frame_gap=30 \
--reward_model.target_method=dirac_delta \
--reward_model.device=cuda \
--dataset.repo_id=<dataset_repo_id> \
--output_dir=outputs/vf_temporal_siglip2 \
--steps=40000 \
--batch_size=16
```
The dataset factory supplies six past-only frames for every observation key.
The model requires `observation.state` and all configured camera streams.
## nanoVLM-460M
Run a frozen-backbone probe first:
```bash
lerobot-train \
--reward_model.type=nanovlm_value_function \
--reward_model.nanovlm_pretrained_path=lusxvr/nanoVLM-460M-8k \
--reward_model.freeze_vision_encoder=true \
--reward_model.freeze_multimodal_projector=true \
--reward_model.freeze_language_model=true \
--reward_model.device=cuda \
--dataset.repo_id=<dataset_repo_id> \
--output_dir=outputs/vf_nanovlm_probe \
--steps=5000 \
--batch_size=1
```
The released checkpoint's native preprocessing resizes the long image side to
2048 and creates 512px global/split views. A 480x640 camera therefore produces
13 vision inputs and roughly 832 image placeholders; use batch size 1 initially
for a three-camera setup.
Then load the probe checkpoint and selectively fine-tune the projector/decoder
at a lower learning rate.
## Standalone Gemma3 VLM alignment
The VLM trainer is intentionally separate from LeRobot:
```bash
cd third_party/nanoVLM
# Projector warmup
torchrun --standalone --nproc_per_node=4 train_recap_gemma3.py \
--output_dir=checkpoints/recap_vlm_warmup \
--steps=8000 \
--freeze_language_model
# Full multimodal alignment
torchrun --standalone --nproc_per_node=4 train_recap_gemma3.py \
--resume_from_checkpoint=checkpoints/recap_vlm_warmup/final \
--output_dir=checkpoints/recap_vlm_aligned \
--steps=50000
```
The output is a standard Hugging Face `Gemma3ForConditionalGeneration`
checkpoint.
## Aligned RECAP value function (run last)
```bash
lerobot-train \
--reward_model.type=distributional_value_function \
--reward_model.vlm_pretrained_path=third_party/nanoVLM/checkpoints/recap_vlm_aligned/final \
--reward_model.freeze_vision_encoder=true \
--reward_model.device=cuda \
--dataset.repo_id=<dataset_repo_id> \
--output_dir=outputs/vf_recap_aligned \
--steps=40000 \
--batch_size=8
```
## Small-batch verification
Before each full run:
```bash
uv run python scripts/overfit_vf_variant.py \
--dataset_repo_id=<dataset_repo_id> \
--reward_type=<distributional_value_function|temporal_siglip_value_function|nanovlm_value_function> \
--num_samples=16 \
--steps=500
```
For `nanovlm_value_function`, start with `--num_samples=2` because all overfit
samples are held in one batch and native image tiling is memory intensive.
Compare runs using held-out episode NLL/MAE, per-episode return rank correlation,
terminal success/failure separation, and the matched-versus-shuffled image loss gap.
+5
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]"] 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]"] robometer = ["lerobot[transformers-dep]", "lerobot[qwen-vl-utils-dep]", "lerobot[peft-dep]"]
topreward = ["lerobot[transformers-dep]"] topreward = ["lerobot[transformers-dep]"]
recap = ["lerobot[transformers-dep]"]
xvla = ["lerobot[transformers-dep]"] xvla = ["lerobot[transformers-dep]"]
eo1 = ["lerobot[transformers-dep]", "lerobot[qwen-vl-utils-dep]"] eo1 = ["lerobot[transformers-dep]", "lerobot[qwen-vl-utils-dep]"]
fastwam = [ fastwam = [
@@ -332,6 +333,7 @@ all = [
"lerobot[sarm]", "lerobot[sarm]",
"lerobot[robometer]", "lerobot[robometer]",
"lerobot[topreward]", "lerobot[topreward]",
"lerobot[recap]",
"lerobot[peft]", "lerobot[peft]",
# "lerobot[unitree_g1]", TODO: Unitree requires specific installation instructions for unitree_sdk2 # "lerobot[unitree_g1]", TODO: Unitree requires specific installation instructions for unitree_sdk2
] ]
@@ -355,6 +357,9 @@ lerobot-edit-dataset="lerobot.scripts.lerobot_edit_dataset:main"
lerobot-setup-can="lerobot.scripts.lerobot_setup_can:main" lerobot-setup-can="lerobot.scripts.lerobot_setup_can:main"
lerobot-annotate="lerobot.scripts.lerobot_annotate:main" lerobot-annotate="lerobot.scripts.lerobot_annotate:main"
lerobot-rollout="lerobot.scripts.lerobot_rollout:main" lerobot-rollout="lerobot.scripts.lerobot_rollout:main"
lerobot-compute-returns="lerobot.scripts.lerobot_compute_returns:main"
lerobot-eval-reward-model="lerobot.scripts.lerobot_eval_reward_model:main"
lerobot-create-advantage-video="lerobot.scripts.lerobot_create_advantage_video:main"
# ---------------- Tool Configurations ---------------- # ---------------- Tool Configurations ----------------
+199
View File
@@ -0,0 +1,199 @@
"""Overfit any distributional VF architecture on a small real-data batch."""
from __future__ import annotations
import argparse
import numpy as np
import torch
from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature
from lerobot.datasets import LeRobotDataset, LeRobotDatasetMetadata
from lerobot.datasets.factory import resolve_delta_timestamps
from lerobot.rewards.distributional_value_function.configuration_distributional_value_function import (
DistributionalVFConfig,
)
from lerobot.rewards.factory import make_reward_model, make_reward_pre_post_processors
from lerobot.rewards.nanovlm_value_function.configuration_nanovlm_value_function import (
NanoVLMVFConfig,
)
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
NANOVLM_IMAGES,
)
from lerobot.rewards.temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
TemporalSiglipVFConfig,
)
from lerobot.utils.collate import lerobot_collate_fn
from lerobot.utils.constants import OBS_STATE
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--dataset_repo_id", required=True)
parser.add_argument("--root", default=None)
parser.add_argument(
"--reward_type",
choices=(
"distributional_value_function",
"temporal_siglip_value_function",
"nanovlm_value_function",
),
required=True,
)
parser.add_argument("--vlm_pretrained_path", default=None)
parser.add_argument("--nanovlm_pretrained_path", default="lusxvr/nanoVLM-460M-8k")
parser.add_argument("--num_samples", type=int, default=16)
parser.add_argument("--steps", type=int, default=500)
parser.add_argument("--lr_head", type=float, default=1e-3)
parser.add_argument("--lr_backbone", type=float, default=1e-5)
parser.add_argument("--history_steps", type=int, default=6)
parser.add_argument("--history_frame_gap", type=int, default=30)
parser.add_argument("--log_every", type=int, default=25)
args = parser.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
metadata = LeRobotDatasetMetadata(args.dataset_repo_id, root=args.root)
input_features = {
key: PolicyFeature(type=FeatureType.VISUAL, shape=tuple(metadata.features[key]["shape"]))
for key in metadata.camera_keys
}
if OBS_STATE in metadata.features:
input_features[OBS_STATE] = PolicyFeature(
type=FeatureType.STATE,
shape=tuple(metadata.features[OBS_STATE]["shape"]),
)
common = {"input_features": input_features, "device": str(device), "target_method": "dirac_delta"}
if args.reward_type == "distributional_value_function":
config = DistributionalVFConfig(
**common,
vlm_pretrained_path=args.vlm_pretrained_path,
freeze_vision_encoder=True,
)
elif args.reward_type == "temporal_siglip_value_function":
config = TemporalSiglipVFConfig(
**common,
history_steps=args.history_steps,
frame_gap=args.history_frame_gap,
)
config.normalization_mapping = {
"VISUAL": NormalizationMode.IDENTITY,
"STATE": NormalizationMode.MEAN_STD,
}
else:
config = NanoVLMVFConfig(
**common,
nanovlm_pretrained_path=args.nanovlm_pretrained_path,
)
delta_timestamps = resolve_delta_timestamps(config, metadata)
dataset = LeRobotDataset(
args.dataset_repo_id,
root=args.root,
delta_timestamps=delta_timestamps,
video_backend="pyav",
)
indices = np.linspace(0, len(dataset) - 1, args.num_samples, dtype=int).tolist()
preprocessor, _ = make_reward_pre_post_processors(
config,
dataset_stats=metadata.stats,
)
samples = []
returns = []
terminals = []
for index in indices:
sample = dataset[index]
returns.append(torch.as_tensor(sample["mc_return"]).reshape(-1)[0])
terminals.append(torch.as_tensor(sample["is_terminal"]).reshape(-1)[0])
samples.append(sample)
raw_batch = lerobot_collate_fn(samples)
if raw_batch is None:
raise ValueError("The selected overfit samples produced an empty batch")
batch = preprocessor(raw_batch)
batch["mc_return"] = torch.stack(returns).to(device)
batch["is_terminal"] = torch.stack(terminals).bool().to(device)
model = make_reward_model(config).to(device)
_run_stage(
model,
batch,
steps=args.steps // 2,
learning_rate=args.lr_head,
head_only=True,
log_every=args.log_every,
label="head probe",
)
_run_stage(
model,
batch,
steps=args.steps - args.steps // 2,
learning_rate=args.lr_backbone,
head_only=False,
log_every=args.log_every,
label="fine-tune",
)
_image_shuffle_diagnostic(model, batch, metadata.camera_keys)
def _set_trainable(model, *, head_only: bool):
for param in model.parameters():
param.requires_grad = False
for param in model.value_head.parameters():
param.requires_grad = True
if hasattr(model, "value_query"):
for param in model.value_query.parameters():
param.requires_grad = True
if head_only:
return
if hasattr(model, "multi_modal_projector"):
model.multi_modal_projector.requires_grad_(True)
model.language_model.requires_grad_(True)
elif hasattr(model, "temporal_transformer"):
for name, param in model.named_parameters():
if not name.startswith("siglip."):
param.requires_grad = True
else:
model.nanovlm.MP.requires_grad_(True)
model.nanovlm.decoder.requires_grad_(True)
def _run_stage(model, batch, *, steps, learning_rate, head_only, log_every, label):
_set_trainable(model, head_only=head_only)
model.train()
params = [param for param in model.parameters() if param.requires_grad]
optimizer = torch.optim.AdamW(params, lr=learning_rate)
print(f"\n{label}: {sum(param.numel() for param in params):,} trainable parameters")
for step in range(steps + 1):
optimizer.zero_grad(set_to_none=True)
loss, metrics = model(batch)
loss.backward()
optimizer.step()
if step % log_every == 0 or step == steps:
print(
f"step={step:04d} loss={metrics['loss']:.4f} "
f"mae={metrics['mae']:.4f} acc={metrics['acc_neighbor']:.3f}"
)
@torch.no_grad()
def _image_shuffle_diagnostic(model, batch, camera_keys):
model.eval()
matched_loss, _ = model(batch)
shuffled = dict(batch)
permutation = torch.roll(torch.arange(batch["mc_return"].shape[0], device=matched_loss.device), 1)
if NANOVLM_IMAGES in batch:
shuffled[NANOVLM_IMAGES] = [batch[NANOVLM_IMAGES][index] for index in permutation.cpu().tolist()]
for key in camera_keys:
if key in batch:
shuffled[key] = batch[key][permutation]
shuffled_loss, _ = model(shuffled)
print(
f"\nvisual dependence: matched_loss={matched_loss.item():.4f} "
f"shuffled_loss={shuffled_loss.item():.4f} "
f"gap={shuffled_loss.item() - matched_loss.item():+.4f}"
)
if __name__ == "__main__":
main()
@@ -169,6 +169,56 @@ class ExecutorConfig:
episode_parallelism: int = 16 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 = ""
# Optional CSV from ``lerobot-eval-reward-model``. When set, annotation
# consumes its per-frame predictions instead of rerunning VF inference.
# This is the recommended path for temporal and multi-camera value models.
predictions_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
# Percentile for binarization threshold ε_. Appendix F uses a threshold
# yielding about 40% positive rollout actions during post-training, i.e.
# the 60th percentile. Pre-training uses 0.7 (~30% positive), while the
# slow-but-successful laundry setting uses 0.9 (~10% positive).
threshold_percentile: float = 0.6
# 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 @dataclass
class AnnotationPipelineConfig: class AnnotationPipelineConfig:
"""Top-level config for ``lerobot-annotate`` (rewrites data shards in place).""" """Top-level config for ``lerobot-annotate`` (rewrites data shards in place)."""
@@ -190,6 +240,7 @@ class AnnotationPipelineConfig:
plan: PlanConfig = field(default_factory=PlanConfig) plan: PlanConfig = field(default_factory=PlanConfig)
interjections: InterjectionsConfig = field(default_factory=InterjectionsConfig) interjections: InterjectionsConfig = field(default_factory=InterjectionsConfig)
vqa: VqaConfig = field(default_factory=VqaConfig) vqa: VqaConfig = field(default_factory=VqaConfig)
advantage: AdvantageConfig = field(default_factory=AdvantageConfig)
vlm: VlmConfig = field(default_factory=VlmConfig) vlm: VlmConfig = field(default_factory=VlmConfig)
executor: ExecutorConfig = field(default_factory=ExecutorConfig) executor: ExecutorConfig = field(default_factory=ExecutorConfig)
@@ -15,20 +15,24 @@
# limitations under the License. # limitations under the License.
"""In-process executor that runs the annotation phases. """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 1: ``plan`` module (plan + subtasks + memory)
phase 2: ``interjections`` module (interjections + speech) phase 2: ``interjections`` module (interjections + speech)
phase 3: ``plan`` plan-update pass — re-runs plan emission at every phase 3: ``plan`` plan-update pass — re-runs plan emission at every
interjection timestamp produced by phase 2 interjection timestamp produced by phase 2
phase 4: ``vqa`` module (VQA) phase 4: ``vqa`` module (VQA)
phase 5: validator phase 5: ``advantage`` module (advantage scoring via frozen VF)
phase 6: writer phase 6: validator
phase 7: writer
Phase 3 is why the ``plan`` module must be re-entered after the Phase 3 is why the ``plan`` module must be re-entered after the
``interjections`` module — to refresh ``plan`` rows at interjection ``interjections`` module — to refresh ``plan`` rows at interjection
timestamps. 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 Distributed execution is provided by Hugging Face Jobs (see
``examples/annotations/run_hf_job.py``); the runner inside the job ``examples/annotations/run_hf_job.py``); the runner inside the job
invokes ``lerobot-annotate`` which uses this in-process executor. invokes ``lerobot-annotate`` which uses this in-process executor.
@@ -74,7 +78,7 @@ class PipelineRunSummary:
@dataclass @dataclass
class Executor: 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`` Episode-level concurrency comes from ``ExecutorConfig.episode_parallelism``
(a thread pool); cluster-level concurrency comes from running this (a thread pool); cluster-level concurrency comes from running this
@@ -86,6 +90,7 @@ class Executor:
plan: Any # PlanSubtasksMemoryModule plan: Any # PlanSubtasksMemoryModule
interjections: Any # InterjectionsAndSpeechModule interjections: Any # InterjectionsAndSpeechModule
vqa: Any # GeneralVqaModule vqa: Any # GeneralVqaModule
advantage: Any # AdvantageModule
writer: LanguageColumnsWriter writer: LanguageColumnsWriter
validator: StagingValidator validator: StagingValidator
@@ -112,6 +117,12 @@ class Executor:
phases.append(self._run_plan_update_phase(records, staging_dir)) phases.append(self._run_plan_update_phase(records, staging_dir))
# Phase 4: ``vqa`` module (VQA) # Phase 4: ``vqa`` module (VQA)
phases.append(self._run_module_phase("vqa", records, staging_dir, self.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) print("[annotate] running validator...", flush=True)
report = self.validator.validate(records, staging_dir) report = self.validator.validate(records, staging_dir)
@@ -179,7 +190,7 @@ class Executor:
staging_dir: Path, staging_dir: Path,
module: Any, module: Any,
) -> PhaseResult: ) -> PhaseResult:
if not module.enabled: if module is None or not module.enabled:
print(f"[annotate] phase={name} skipped (module disabled)", flush=True) print(f"[annotate] phase={name} skipped (module disabled)", flush=True)
return PhaseResult(name=name, episodes_processed=0, episodes_skipped=len(records)) return PhaseResult(name=name, episodes_processed=0, episodes_skipped=len(records))
n = len(records) n = len(records)
@@ -231,7 +242,7 @@ class Executor:
``plan`` module with the interjection timestamps so its existing ``plan`` module with the interjection timestamps so its existing
prompt path is reused. 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)) return PhaseResult(name="plan_update", episodes_processed=0, episodes_skipped=len(records))
processed = 0 processed = 0
for record in records: for record in records:
@@ -14,11 +14,13 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from .advantage import AdvantageModule
from .general_vqa import GeneralVqaModule from .general_vqa import GeneralVqaModule
from .interjections_and_speech import InterjectionsAndSpeechModule from .interjections_and_speech import InterjectionsAndSpeechModule
from .plan_subtasks_memory import PlanSubtasksMemoryModule from .plan_subtasks_memory import PlanSubtasksMemoryModule
__all__ = [ __all__ = [
"AdvantageModule",
"GeneralVqaModule", "GeneralVqaModule",
"InterjectionsAndSpeechModule", "InterjectionsAndSpeechModule",
"PlanSubtasksMemoryModule", "PlanSubtasksMemoryModule",
@@ -0,0 +1,429 @@
#!/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 csv
import logging
from dataclasses import dataclass, field
from pathlib import Path
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)
_prediction_lookup: dict[tuple[int, int], float] | None = field(default=None, 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.
"""
if self.config.predictions_path:
self._ensure_predictions_loaded()
else:
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.
"""
if self.config.predictions_path:
self._ensure_predictions_loaded()
assert self._prediction_lookup is not None
missing = [
frame_index
for frame_index in record.frame_indices
if (record.episode_index, frame_index) not in self._prediction_lookup
]
if missing:
raise KeyError(
f"Predictions CSV is missing {len(missing)} frame(s) from episode "
f"{record.episode_index}; first missing frame_index={missing[0]}"
)
return np.asarray(
[
self._prediction_lookup[(record.episode_index, frame_index)]
for frame_index in record.frame_indices
],
dtype=np.float32,
)
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 _ensure_predictions_loaded(self) -> None:
if self._prediction_lookup is not None:
return
path = Path(self.config.predictions_path)
if not path.is_file():
raise FileNotFoundError(f"Advantage predictions CSV not found: {path}")
lookup: dict[tuple[int, int], float] = {}
with path.open(newline="") as handle:
reader = csv.DictReader(handle)
required = {"episode_index", "frame_index", "predicted_value"}
missing_columns = required.difference(reader.fieldnames or ())
if missing_columns:
raise ValueError(f"Predictions CSV is missing required columns: {sorted(missing_columns)}")
for row in reader:
key = (int(row["episode_index"]), int(row["frame_index"]))
if key in lookup:
raise ValueError(f"Duplicate prediction for episode/frame {key}")
lookup[key] = float(row["predicted_value"])
self._prediction_lookup = lookup
logger.info("Loaded %d per-frame value predictions from %s", len(lookup), path)
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 and not self.config.predictions_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 and not self.config.predictions_path:
logger.warning(
"No value_function_path, predictions_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 from __future__ import annotations
import functools
from collections.abc import Iterator, Sequence from collections.abc import Iterator, Sequence
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
@@ -42,6 +43,18 @@ from lerobot.datasets.io_utils import load_tasks
from lerobot.datasets.utils import DEFAULT_TASKS_PATH 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 @dataclass
class EpisodeRecord: class EpisodeRecord:
"""Per-episode record yielded by the reader.""" """Per-episode record yielded by the reader."""
@@ -61,10 +74,7 @@ class EpisodeRecord:
def frames_df(self): # type: ignore[no-untyped-def] def frames_df(self): # type: ignore[no-untyped-def]
"""Lazy-load the pandas slice for this episode (memoized).""" """Lazy-load the pandas slice for this episode (memoized)."""
if self._frames_df_cache is None: if self._frames_df_cache is None:
import pandas as pd # noqa: PLC0415 - deferred for optional dataset extra df = _read_parquet_as_pandas(self.data_path)
table = pq.read_table(self.data_path)
df: pd.DataFrame = table.to_pandas()
self._frames_df_cache = df.iloc[self.row_offset : self.row_offset + self.row_count].reset_index( self._frames_df_cache = df.iloc[self.row_offset : self.row_offset + self.row_count].reset_index(
drop=True drop=True
) )
@@ -39,6 +39,7 @@ _MODULES: tuple[ModuleName, ...] = (
"plan", "plan",
"interjections", "interjections",
"vqa", "vqa",
"advantage",
) )
+1
View File
@@ -32,6 +32,7 @@ DEFAULT_BINDINGS = {
"interjection": "emitted_at(t, style=interjection)", "interjection": "emitted_at(t, style=interjection)",
"vqa": "emitted_at(t, style=vqa, role=assistant)", "vqa": "emitted_at(t, style=vqa, role=assistant)",
"vqa_query": "emitted_at(t, style=vqa, role=user)", "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_]*)\}") 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 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. """Validates that all dataset metadata have consistent properties.
Ensures all datasets have the same fps, robot_type, and features to guarantee 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: Args:
all_metadata: List of LeRobotDatasetMetadata objects to validate. 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: 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: Raises:
ValueError: If any metadata has different fps, robot_type, or features 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 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}." 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): if not features_equal_for_merge(features, meta.features):
raise ValueError( if not lenient:
f"Same features is expected, but got features={meta.features} instead of {features}." 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 return fps, robot_type, features
@@ -289,6 +298,7 @@ def aggregate_datasets(
chunk_size: int | None = None, chunk_size: int | None = None,
concatenate_videos: bool = True, concatenate_videos: bool = True,
concatenate_data: bool = True, concatenate_data: bool = True,
lenient: bool = False,
): ):
"""Aggregates multiple LeRobot datasets into a single unified dataset. """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) LeRobotDatasetMetadata(repo_id, root=root) for repo_id, root in zip(repo_ids, roots, strict=False)
] ]
) )
fps, robot_type, _ = validate_all_metadata(all_metadata) fps, robot_type, union_features = validate_all_metadata(all_metadata, lenient=lenient)
features = merge_video_feature_info_for_aggregate(all_metadata) 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"] video_keys = [key for key in features if features[key]["dtype"] == "video"]
dst_meta = LeRobotDatasetMetadata.create( 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 = pd.read_parquet(src_path)
df = update_data_df(df, src_meta, dst_meta) 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 # Write data and get the actual destination file it was written to
# This avoids duplicating the rotation logic here # This avoids duplicating the rotation logic here
data_idx, (dst_chunk, dst_file) = append_or_create_parquet_file( 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, output_dir: str | Path | None = None,
concatenate_videos: bool = True, concatenate_videos: bool = True,
concatenate_data: bool = True, concatenate_data: bool = True,
lenient: bool = False,
) -> LeRobotDataset: ) -> LeRobotDataset:
"""Merge multiple LeRobotDatasets into a single dataset. """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. 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_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. 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: if not datasets:
raise ValueError("No datasets to merge") raise ValueError("No datasets to merge")
@@ -301,6 +303,7 @@ def merge_datasets(
aggr_root=output_dir, aggr_root=output_dir,
concatenate_videos=concatenate_videos, concatenate_videos=concatenate_videos,
concatenate_data=concatenate_data, concatenate_data=concatenate_data,
lenient=lenient,
) )
merged_dataset = LeRobotDataset( merged_dataset = LeRobotDataset(
+2 -2
View File
@@ -43,10 +43,10 @@ CORE_STYLES = {
# validation. Empty by default — populate from a downstream module that # validation. Empty by default — populate from a downstream module that
# also extends ``PERSISTENT_STYLES`` or ``EVENT_ONLY_STYLES`` to declare # also extends ``PERSISTENT_STYLES`` or ``EVENT_ONLY_STYLES`` to declare
# the new style's column. # the new style's column.
EXTENDED_STYLES: set[str] = set() EXTENDED_STYLES: set[str] = {"advantage"}
STYLE_REGISTRY = CORE_STYLES | EXTENDED_STYLES 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"} EVENT_ONLY_STYLES = {"interjection", "vqa", "trace"}
# Styles whose ``content`` is grounded in a specific camera view. Rows of these # 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 num_inference_steps: int | None = None
mask_action_dim_padding: bool = True mask_action_dim_padding: bool = True
enable_inference_cuda_graph: 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 # MolmoAct2-local eval option. When enabled, stochastic continuous action
# generation uses a rollout-local generator derived from eval_seed. # generation uses a rollout-local generator derived from eval_seed.
per_episode_seed: bool = False per_episode_seed: bool = False
@@ -497,6 +497,56 @@ def _weighted_per_example(
return loss_sum * float(batch_size) / global_weight_sum 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): class MolmoAct2Policy(PreTrainedPolicy):
"""MolmoAct2 policy wrapping the vendored HF model for LeRobot. """MolmoAct2 policy wrapping the vendored HF model for LeRobot.
@@ -1622,6 +1672,183 @@ class MolmoAct2Policy(PreTrainedPolicy):
metrics["loss"] = loss.detach().float().mean().item() metrics["loss"] = loss.detach().float().mean().item()
return loss, metrics 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() @torch.no_grad()
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor: def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
"""Generate an action chunk via continuous flow matching or discrete AR decoding.""" """Generate an action chunk via continuous flow matching or discrete AR decoding."""
@@ -1657,6 +1884,15 @@ class MolmoAct2Policy(PreTrainedPolicy):
model_inputs=model_inputs, model_inputs=model_inputs,
action_dim=action_dim, 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(): elif self._rtc_enabled():
actions = self._generate_actions_from_inputs_with_rtc( actions = self._generate_actions_from_inputs_with_rtc(
model_inputs=model_inputs, model_inputs=model_inputs,
@@ -359,6 +359,7 @@ def _build_robot_text(
add_setup_tokens: bool, add_setup_tokens: bool,
add_control_tokens: bool, add_control_tokens: bool,
num_images: int, num_images: int,
advantage: str = "",
) -> str: ) -> str:
setup_text = _wrap_setup_text(setup_type, add_setup_tokens=add_setup_tokens) 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) control_text = _wrap_control_text(control_mode, add_control_tokens=add_control_tokens)
@@ -375,7 +376,10 @@ def _build_robot_text(
image_prefix = "<|image|>" image_prefix = "<|image|>"
else: else:
image_prefix = "".join(f"Image {idx + 1}<|image|>" for idx in range(num_images)) 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]: def _as_text_list(value: Any, batch_size: int) -> list[str]:
@@ -695,6 +699,39 @@ class MolmoAct2ClampNormalizedProcessorStep(ProcessorStep):
return features 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") @ProcessorStepRegistry.register(name="molmoact2_pack_inputs")
@dataclass @dataclass
class MolmoAct2PackInputsProcessorStep(ProcessorStep): class MolmoAct2PackInputsProcessorStep(ProcessorStep):
@@ -715,6 +752,10 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
chunk_size: int = 30 chunk_size: int = 30
max_action_dim: int = 32 max_action_dim: int = 32
env_action_dim: int | None = None 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: def __post_init__(self) -> None:
require_package("transformers", extra="molmoact2") require_package("transformers", extra="molmoact2")
@@ -757,6 +798,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
"chunk_size": self.chunk_size, "chunk_size": self.chunk_size,
"max_action_dim": self.max_action_dim, "max_action_dim": self.max_action_dim,
"env_action_dim": self.env_action_dim, "env_action_dim": self.env_action_dim,
"advantage_prefix": self.advantage_prefix,
} }
def _resolve_max_sequence_length( def _resolve_max_sequence_length(
@@ -919,8 +961,40 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
if task_source is None: if task_source is None:
task_source = complementary.get("language_instruction") task_source = complementary.get("language_instruction")
tasks = _as_text_list(task_source, batch_size) 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 complementary["task"] = tasks
action_padded = None action_padded = None
@@ -953,6 +1027,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
add_setup_tokens=self.add_setup_tokens, add_setup_tokens=self.add_setup_tokens,
add_control_tokens=self.add_control_tokens, add_control_tokens=self.add_control_tokens,
num_images=len(images), num_images=len(images),
advantage=advantages[batch_idx],
) )
prompt_texts.append(prompt) prompt_texts.append(prompt)
if build_action_labels: if build_action_labels:
@@ -989,6 +1064,33 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
if build_action_labels: if build_action_labels:
inputs["labels"] = self._build_labels(inputs["input_ids"], inputs["attention_mask"]) 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.update(dict(inputs))
complementary["action_dim_is_pad"] = action_dim_is_pad complementary["action_dim_is_pad"] = action_dim_is_pad
if action_horizon_is_pad is not None: if action_horizon_is_pad is not None:
@@ -1164,28 +1266,49 @@ def make_molmoact2_pre_post_processors(
stats=masked_dataset_stats, stats=masked_dataset_stats,
), ),
MolmoAct2ClampNormalizedProcessorStep(normalization_masks=normalization_masks), 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] = [ output_steps: list[ProcessorStep] = [
MolmoAct2ClampActionProcessorStep(), MolmoAct2ClampActionProcessorStep(),
MolmoAct2MaskedUnnormalizerProcessorStep( MolmoAct2MaskedUnnormalizerProcessorStep(
@@ -87,6 +87,17 @@ class PI05Config(PreTrainedConfig):
freeze_vision_encoder: bool = False # Freeze only the vision encoder freeze_vision_encoder: bool = False # Freeze only the vision encoder
train_expert_only: bool = False # Freeze entire VLM, train only action expert and projections 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 settings: see openpi `AdamW`
optimizer_lr: float = 2.5e-5 # see openpi `CosineDecaySchedule: peak_lr` optimizer_lr: float = 2.5e-5 # see openpi `CosineDecaySchedule: peak_lr`
optimizer_betas: tuple[float, float] = (0.9, 0.95) optimizer_betas: tuple[float, float] = (0.9, 0.95)
+141 -2
View File
@@ -52,6 +52,8 @@ from lerobot.utils.constants import (
ACTION, ACTION,
OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS, OBS_LANGUAGE_TOKENS,
OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
OBS_LANGUAGE_UNCOND_TOKENS,
OPENPI_ATTENTION_MASK_VALUE, 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): def pad_vector(vector, new_dim):
"""Pad the last dimension of a vector to new_dim with zeros. """Pad the last dimension of a vector to new_dim with zeros.
@@ -797,9 +813,17 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
masks, masks,
noise=None, noise=None,
num_steps=None, num_steps=None,
uncond_tokens=None,
uncond_masks=None,
**kwargs: Unpack[ActionSelectKwargs], **kwargs: Unpack[ActionSelectKwargs],
) -> Tensor: ) -> 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: if num_steps is None:
num_steps = self.config.num_inference_steps 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 ) # Use config max_action_dim for internal processing
noise = self.sample_noise(actions_shape, device) 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_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_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks)
prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1 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, 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 dt = -1.0 / num_steps
x_t = noise 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) time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor): 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( return self.denoise_step(
prefix_pad_masks=prefix_pad_masks, prefix_pad_masks=prefix_pad_masks,
past_key_values=past_key_values, past_key_values=past_key_values,
@@ -907,6 +960,80 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
suffix_out = suffix_out.to(dtype=torch.float32) suffix_out = suffix_out.to(dtype=torch.float32)
return self.action_out_proj(suffix_out) 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): class PI05Policy(PreTrainedPolicy):
"""PI05 Policy for LeRobot.""" """PI05 Policy for LeRobot."""
@@ -1243,8 +1370,20 @@ class PI05Policy(PreTrainedPolicy):
images, img_masks = self._preprocess_images(batch) images, img_masks = self._preprocess_images(batch)
tokens, masks = batch[f"{OBS_LANGUAGE_TOKENS}"], batch[f"{OBS_LANGUAGE_ATTENTION_MASK}"] 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) # 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 # Unpad actions to actual action dimension
original_action_dim = self.config.output_features[ACTION].shape[0] 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.types import EnvTransition, TransitionKey
from lerobot.utils.constants import ( from lerobot.utils.constants import (
OBS_LANGUAGE_UNCOND_ATTENTION_MASK,
OBS_LANGUAGE_UNCOND_TOKENS,
OBS_STATE, OBS_STATE,
POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_POSTPROCESSOR_DEFAULT_NAME,
POLICY_PREPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME,
@@ -57,6 +59,7 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
max_state_dim: int = 32 max_state_dim: int = 32
task_key: str = "task" task_key: str = "task"
cfg_enabled: bool = False
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy() transition = transition.copy()
@@ -84,8 +87,25 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
full_prompts.append(full_prompt) full_prompts.append(full_prompt)
transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts 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 return transition
def transform_features( def transform_features(
@@ -111,9 +131,10 @@ def make_pi05_pre_post_processors(
1. Renaming features to match pretrained configurations. 1. Renaming features to match pretrained configurations.
2. Normalizing input and output features based on dataset statistics. 2. Normalizing input and output features based on dataset statistics.
3. Adding a batch dimension. 3. Adding a batch dimension.
4. Appending a newline character to the task description for tokenizer compatibility. 4. (Optional) Rendering language annotations via recipe YAML.
5. Tokenizing the text prompt using the PaliGemma tokenizer. 5. (Optional) Flattening rendered messages into the task string.
6. Moving all data to the specified device. 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: The post-processing pipeline handles the model's output by:
1. Moving data to the CPU. 1. Moving data to the CPU.
@@ -122,8 +143,6 @@ def make_pi05_pre_post_processors(
Args: Args:
config: The configuration object for the PI0 policy. config: The configuration object for the PI0 policy.
dataset_stats: A dictionary of statistics for normalization. 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: Returns:
A tuple containing the configured pre-processor and post-processor pipelines. 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, norm_map=config.normalization_mapping,
stats=dataset_stats, 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] = [ output_steps: list[ProcessorStep] = [
UnnormalizerProcessorStep( UnnormalizerProcessorStep(
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
+4
View File
@@ -164,6 +164,10 @@ _COMPLEMENTARY_KEYS = (
"messages", "messages",
"message_streams", "message_streams",
"target_message_indices", "target_message_indices",
"mc_return",
"is_terminal",
"next.success",
"intervention",
) )
+14 -2
View File
@@ -1054,8 +1054,20 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
try: try:
step_class = ProcessorStepRegistry.get(step_entry["registry_name"]) step_class = ProcessorStepRegistry.get(step_entry["registry_name"])
return step_class, step_entry["registry_name"] return step_class, step_entry["registry_name"]
except KeyError as e: except KeyError:
raise ImportError(f"Failed to load processor step from registry. {str(e)}") from e 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: else:
# Fallback to dynamic import using the full class path # Fallback to dynamic import using the full class path
full_class_path = step_entry["class"] full_class_path = step_entry["class"]
@@ -40,11 +40,14 @@ class RenderMessagesStep(ProcessorStep):
``message_streams`` / ``target_message_indices`` keys. ``message_streams`` / ``target_message_indices`` keys.
""" """
recipe: TrainingRecipe recipe: TrainingRecipe | None = None
dataset_ctx: Any | None = None dataset_ctx: Any | None = None
def __call__(self, transition: EnvTransition) -> EnvTransition | None: def __call__(self, transition: EnvTransition) -> EnvTransition | None:
"""Render messages for a single transition; return ``None`` to drop it.""" """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 {} complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
persistent = complementary_data.get(LANGUAGE_PERSISTENT) or [] persistent = complementary_data.get(LANGUAGE_PERSISTENT) or []
events = complementary_data.get(LANGUAGE_EVENTS) 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_side: str = "right"
padding: str = "max_length" padding: str = "max_length"
truncation: bool = True 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) # Internal tokenizer instance (not part of the config)
input_tokenizer: Any = field(default=None, init=False, repr=False) input_tokenizer: Any = field(default=None, init=False, repr=False)
@@ -201,8 +203,8 @@ class TokenizerProcessorStep(ObservationProcessorStep):
new_observation = dict(observation) new_observation = dict(observation)
# Add tokenized data to the observation # Add tokenized data to the observation
new_observation[OBS_LANGUAGE_TOKENS] = tokenized_prompt["input_ids"] new_observation[self.output_tokens_key] = tokenized_prompt["input_ids"]
new_observation[OBS_LANGUAGE_ATTENTION_MASK] = tokenized_prompt["attention_mask"].to(dtype=torch.bool) new_observation[self.output_mask_key] = tokenized_prompt["attention_mask"].to(dtype=torch.bool)
# Tokenize subtask if available # Tokenize subtask if available
subtask = self.get_subtask(self.transition) subtask = self.get_subtask(self.transition)
@@ -309,14 +311,14 @@ class TokenizerProcessorStep(ObservationProcessorStep):
The updated dictionary of policy features. The updated dictionary of policy features.
""" """
# Add a feature for the token IDs if it doesn't already exist # Add a feature for the token IDs if it doesn't already exist
if OBS_LANGUAGE_TOKENS not in features[PipelineFeatureType.OBSERVATION]: if self.output_tokens_key not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][OBS_LANGUAGE_TOKENS] = PolicyFeature( features[PipelineFeatureType.OBSERVATION][self.output_tokens_key] = PolicyFeature(
type=FeatureType.LANGUAGE, shape=(self.max_length,) type=FeatureType.LANGUAGE, shape=(self.max_length,)
) )
# Add a feature for the attention mask if it doesn't already exist # Add a feature for the attention mask if it doesn't already exist
if OBS_LANGUAGE_ATTENTION_MASK not in features[PipelineFeatureType.OBSERVATION]: if self.output_mask_key not in features[PipelineFeatureType.OBSERVATION]:
features[PipelineFeatureType.OBSERVATION][OBS_LANGUAGE_ATTENTION_MASK] = PolicyFeature( features[PipelineFeatureType.OBSERVATION][self.output_mask_key] = PolicyFeature(
type=FeatureType.LANGUAGE, shape=(self.max_length,) type=FeatureType.LANGUAGE, shape=(self.max_length,)
) )
+12
View File
@@ -13,23 +13,35 @@
# limitations under the License. # limitations under the License.
from .classifier.configuration_classifier import RewardClassifierConfig as RewardClassifierConfig from .classifier.configuration_classifier import RewardClassifierConfig as RewardClassifierConfig
from .distributional_value_function.configuration_distributional_value_function import (
DistributionalVFConfig as DistributionalVFConfig,
)
from .factory import ( from .factory import (
get_reward_model_class as get_reward_model_class, get_reward_model_class as get_reward_model_class,
make_reward_model as make_reward_model, make_reward_model as make_reward_model,
make_reward_model_config as make_reward_model_config, make_reward_model_config as make_reward_model_config,
make_reward_pre_post_processors as make_reward_pre_post_processors, make_reward_pre_post_processors as make_reward_pre_post_processors,
) )
from .nanovlm_value_function.configuration_nanovlm_value_function import (
NanoVLMVFConfig as NanoVLMVFConfig,
)
from .pretrained import PreTrainedRewardModel as PreTrainedRewardModel from .pretrained import PreTrainedRewardModel as PreTrainedRewardModel
from .robometer.configuration_robometer import RobometerConfig as RobometerConfig from .robometer.configuration_robometer import RobometerConfig as RobometerConfig
from .sarm.configuration_sarm import SARMConfig as SARMConfig from .sarm.configuration_sarm import SARMConfig as SARMConfig
from .temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
TemporalSiglipVFConfig as TemporalSiglipVFConfig,
)
from .topreward.configuration_topreward import TOPRewardConfig as TOPRewardConfig from .topreward.configuration_topreward import TOPRewardConfig as TOPRewardConfig
__all__ = [ __all__ = [
# Configuration classes # Configuration classes
"DistributionalVFConfig",
"NanoVLMVFConfig",
"RewardClassifierConfig", "RewardClassifierConfig",
"RobometerConfig", "RobometerConfig",
"SARMConfig", "SARMConfig",
"TOPRewardConfig", "TOPRewardConfig",
"TemporalSiglipVFConfig",
# Base class # Base class
"PreTrainedRewardModel", "PreTrainedRewardModel",
# Factory functions # Factory functions
@@ -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,124 @@
"""Shared distributional targets, loss, and metrics for VF experiments."""
from __future__ import annotations
import math
from typing import Any
import torch
import torch.nn.functional as F # noqa: N812
from torch import Tensor
class DistributionalValueMixin:
"""Mixin for models that expose ``_get_value_readout(batch)``."""
config: Any
value_head: Any
hl_gauss_sigma: float
def hl_gauss_target(self, target_value: Tensor) -> Tensor:
target_value = target_value.reshape(-1).clamp(
self.config.value_support_min, self.config.value_support_max
)
target_value = target_value.to(self.value_head.bin_centers.dtype)
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 = 0.5 * (
1.0
+ torch.erf((support_edges[None] - target_value[:, None]) / (self.hl_gauss_sigma * math.sqrt(2)))
)
normalization = (cdf[:, -1] - cdf[:, 0]).unsqueeze(-1).clamp_min(1e-10)
return (cdf[:, 1:] - cdf[:, :-1]) / normalization
def dirac_delta_target(self, target_value: Tensor) -> Tensor:
"""Two-hot scalar projection (legacy configuration name: ``dirac_delta``)."""
target_value = target_value.reshape(-1).clamp(
self.config.value_support_min, self.config.value_support_max
)
target_value = target_value.to(self.value_head.bin_centers.dtype)
bin_width = self.value_head.bin_centers[1] - self.value_head.bin_centers[0]
position = (target_value - self.config.value_support_min) / bin_width
lower = position.floor().long().clamp(0, self.config.num_value_bins - 1)
upper = position.ceil().long().clamp(0, self.config.num_value_bins - 1)
weight_upper = position - lower.float()
weight_lower = upper.float() - position
same = lower == upper
weight_upper = torch.where(same, torch.zeros_like(weight_upper), weight_upper)
weight_lower = torch.where(same, torch.ones_like(weight_lower), weight_lower)
distribution = torch.zeros(
target_value.shape[0],
self.config.num_value_bins,
device=target_value.device,
dtype=target_value.dtype,
)
rows = torch.arange(target_value.shape[0], device=target_value.device)
distribution[rows, lower] += weight_lower
distribution[rows, upper] += weight_upper
return distribution
def compute_target_distribution(self, target_value: Tensor, is_terminal: Tensor) -> Tensor:
if self.config.target_method == "hl_gauss":
base = self.hl_gauss_target(target_value)
elif self.config.target_method == "dirac_delta":
base = self.dirac_delta_target(target_value)
else:
raise ValueError(f"Unknown target method: {self.config.target_method}")
if not self.config.use_one_hot_terminal:
return base
is_terminal = is_terminal.reshape(-1)
if is_terminal.numel() != base.shape[0]:
raise ValueError(f"Expected {base.shape[0]} terminal flags, got {is_terminal.numel()}")
nearest = torch.argmin(
torch.abs(
self.value_head.bin_centers[None]
- target_value.reshape(-1, 1).to(self.value_head.bin_centers.dtype)
),
dim=-1,
)
terminal = F.one_hot(nearest, num_classes=self.config.num_value_bins).to(base.dtype)
return torch.where(is_terminal[:, None].bool(), terminal, base)
def _distributional_forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
readout = self._get_value_readout(batch)
logits = self.value_head(readout)
probabilities = logits.softmax(-1)
predicted_value = (probabilities * self.value_head.bin_centers.to(probabilities.dtype)).sum(-1)
targets = self.compute_target_distribution(batch["mc_return"], batch["is_terminal"])
loss = -(targets * logits.log_softmax(-1)).sum(-1).mean()
target_values = (
batch["mc_return"].reshape(-1).clamp(self.config.value_support_min, self.config.value_support_max)
)
return loss, {
"loss": loss.item(),
"predicted_value_mean": predicted_value.mean().item(),
"mc_return_mean": target_values.mean().item(),
"mae": (predicted_value - target_values).abs().mean().item(),
"acc_best": (logits.argmax(-1) == targets.argmax(-1)).float().mean().item(),
"acc_neighbor": _neighbor_accuracy(
logits.argmax(-1),
target_values,
self.value_head.bin_centers,
),
}
def compute_reward(self, batch: dict[str, Tensor]) -> Tensor:
logits = self.value_head(self._get_value_readout(batch))
probabilities = logits.softmax(-1)
return (probabilities * self.value_head.bin_centers.to(probabilities.dtype)).sum(-1)
def _neighbor_accuracy(predicted_bin: Tensor, target: Tensor, centers: Tensor) -> float:
width = centers[1] - centers[0]
position = (target - centers[0]) / width
lower = position.floor().long().clamp(0, len(centers) - 1)
upper = position.ceil().long().clamp(0, len(centers) - 1)
return ((predicted_bin == lower) | (predicted_bin == upper)).float().mean().item()
@@ -0,0 +1,121 @@
# 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
Architecture source of truth: "π0.6 Model Card", Section 2 (Model Design)
https://website.pi-asset.com/pi06star/PI06_model_card.pdf
Distributional value function V^{pi_ref}(o_t, l) (Section IV-A).
Architecture (~670M params):
Vision: SigLIP2-so400m 27 layers, 1152-dim, 1024 patches/image at 448px
LM: Gemma3-270M 18 layers, 640-dim
Proj: 2x2 pool RMSNorm Linear(1152, 640), 256 soft tokens/image
Readout: one-way learned value query 2-layer MLP 201 bins
Inputs: multi-camera images (3 x 256 soft tokens) + ``"Task: {task}."`` prompt
Targets: MC returns in [-1, 0], cross-entropy on Dirac delta (default) or HL-Gauss
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 Dirac delta (C51, default) or HL-Gauss soft targets,
with optional one-hot targets for terminal states.
Architecture: adapted from the native Gemma3 multimodal VLM design and
scaled to π0.6's ~670M value backbone:
448px SigLIP2-so400m images are pooled from 1024 patches to 256 soft
tokens, RMS-normalized, projected into Gemma3-270M, and followed by a
one-way learned value-query token. Image tokens attend bidirectionally;
text and the value query remain causal.
"""
# Backbone pretrained paths
siglip_path: str = "google/siglip2-so400m-patch14-384"
gemma3_path: str = "google/gemma-3-270m"
# Optional standard Gemma3ForConditionalGeneration checkpoint produced by
# standalone VLM alignment. When set, it supplies vision, connector, and LM.
vlm_pretrained_path: str | None = None
# Distributional head
num_value_bins: int = 201
value_support_min: float = -1.0
value_support_max: float = 0.0
# Stop Regressing (Farebrother et al., 2024) default: spreads most
# probability mass across approximately six neighboring bins.
hl_gauss_sigma_ratio: float = 0.75
# Target distribution method: "dirac_delta" (paper-faithful C51) or "hl_gauss" (soft)
target_method: str = "dirac_delta"
# 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] = (448, 448)
num_image_tokens: int = 256
# 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
# 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,602 @@
# 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
Architecture source of truth: "π0.6 Model Card", Section 2 (Model Design)
https://website.pi-asset.com/pi06star/PI06_model_card.pdf
Implements the distributional value function V^{pi_ref}(o_t, l) from Section IV-A.
It adapts the native Gemma3 multimodal VLM design to π0.6's smaller ~670M scale:
448px SigLIP images are pooled to 256 soft tokens, RMS-normalized, projected
into Gemma3-270M, and processed with bidirectional image / causal text
attention. A final learned value-query token supplies the 201-bin readout.
"""
from __future__ import annotations
import copy
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,
)
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 (
Gemma3Config,
Gemma3ForCausalLM,
Gemma3ForConditionalGeneration,
SiglipVisionModel,
)
from transformers.models.gemma3.modeling_gemma3 import (
Gemma3MultiModalProjector,
create_causal_mask_mapping,
)
else:
Gemma3Config = None # type: ignore[assignment]
Gemma3ForCausalLM = None # type: ignore[assignment]
Gemma3ForConditionalGeneration = None # type: ignore[assignment]
Gemma3MultiModalProjector = None # type: ignore[assignment]
SiglipVisionModel = None # type: ignore[assignment]
create_causal_mask_mapping = None # type: ignore[assignment]
class ValueHead(nn.Module):
"""Categorical value projection: hidden state → bin logits.
The 2-layer MLP topology is adapted from Robometer's prediction head:
Linear LayerNorm GELU Dropout Linear. Unlike Robometer's progress
and success heads, this head predicts RECAP's 201-bin MC-return
distribution over [-1, 0] from the final value-query representation.
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: adapted from the native Gemma3 multimodal VLM using a
448px SigLIP2-so400m vision tower, Gemma3 multimodal projector, and
Gemma3-270M language backbone. Each camera is represented by 256 soft
image tokens. Image tokens are bidirectional, text is causal, and a final
one-way value query supplies the hidden state consumed by the 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
if config.vlm_pretrained_path:
aligned_vlm = Gemma3ForConditionalGeneration.from_pretrained(config.vlm_pretrained_path)
self._vlm_config = aligned_vlm.config
self.vision_encoder = aligned_vlm.model.vision_tower
self.multi_modal_projector = aligned_vlm.model.multi_modal_projector
self.language_model = aligned_vlm.model.language_model
else:
self.vision_encoder = SiglipVisionModel.from_pretrained(config.siglip_path)
gemma3 = Gemma3ForCausalLM.from_pretrained(config.gemma3_path)
self.language_model = gemma3.model
# Adapt Gemma3's native multimodal connector to the 270M text
# backbone and π0.6's 448px input layout.
vision_config = copy.deepcopy(self.vision_encoder.config)
vision_config.image_size = config.image_resolution[0]
text_config = copy.deepcopy(gemma3.config)
self._vlm_config = Gemma3Config(
vision_config=vision_config,
text_config=text_config,
mm_tokens_per_image=config.num_image_tokens,
)
self.multi_modal_projector = Gemma3MultiModalProjector(self._vlm_config)
nn.init.normal_(
self.multi_modal_projector.mm_input_projection_weight,
mean=0.0,
std=self._vlm_config.initializer_range,
)
self._validate_vlm_config()
self.gemma3_hidden = self._vlm_config.text_config.hidden_size # 640
# One-way suffix query, analogous to PI05's suffix/action tokens.
self.value_query = nn.Embedding(1, self.gemma3_hidden)
nn.init.normal_(self.value_query.weight, std=0.02)
# Value head: value-query 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()
@property
def gemma3(self) -> nn.Module:
"""Backward-compatible access to the Gemma3 text backbone."""
return self.language_model
def _validate_vlm_config(self) -> None:
"""Validate the π0.6 448px → 256-token multimodal layout."""
vision_config = self._vlm_config.vision_config
image_size = self.config.image_resolution[0]
if vision_config.image_size != image_size:
raise ValueError(
f"VLM vision image_size ({vision_config.image_size}) does not match "
f"DistributionalVFConfig.image_resolution ({image_size})"
)
patches_per_side = image_size // vision_config.patch_size
tokens_per_side = int(self.config.num_image_tokens**0.5)
if tokens_per_side**2 != self.config.num_image_tokens:
raise ValueError("num_image_tokens must be a perfect square")
if patches_per_side % tokens_per_side:
raise ValueError(
f"{patches_per_side} patches/side cannot be evenly pooled to {tokens_per_side} tokens/side"
)
if self._vlm_config.mm_tokens_per_image != self.config.num_image_tokens:
raise ValueError(
f"VLM emits {self._vlm_config.mm_tokens_per_image} image tokens, "
f"expected {self.config.num_image_tokens}"
)
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.language_model.parameters():
param.requires_grad = False
self.language_model.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.language_model.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 with π0.6's Gemma3 visual connector.
Args:
image: [batch_size, channels, height, width] preprocessed image in [-1, 1].
Returns:
[B, 256, gemma3_hidden] pooled, normalized, projected features.
"""
vision_dtype = next(self.vision_encoder.parameters()).dtype
image_features = self.vision_encoder(
pixel_values=image.to(dtype=vision_dtype),
interpolate_pos_encoding=True,
).last_hidden_state
projected_features = self.multi_modal_projector(image_features)
if projected_features.shape[1] != self.config.num_image_tokens:
raise RuntimeError(
f"Expected {self.config.num_image_tokens} image tokens, got {projected_features.shape[1]}"
)
return projected_features
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.language_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, Tensor]:
"""Build [image soft tokens..., text] plus masks.
Returns:
embs: [B, total_prefix_len, hidden_dim]
pad_masks: [B, total_prefix_len] boolean
token_type_ids: [B, total_prefix_len], 1=image and 0=text
"""
embs: list[Tensor] = []
pad_masks: list[Tensor] = []
token_types: list[Tensor] = []
for img, img_mask in zip(images, img_masks, strict=True):
img_emb = self.embed_image(img)
bsize, num_image_tokens = img_emb.shape[:2]
embs.append(img_emb)
pad_masks.append(img_mask[:, None].expand(bsize, num_image_tokens))
token_types.append(torch.ones(bsize, num_image_tokens, dtype=torch.long, device=img_emb.device))
embs.append(text_embeddings)
pad_masks.append(text_padding_mask)
token_types.append(torch.zeros_like(text_padding_mask, dtype=torch.long))
return torch.cat(embs, dim=1), torch.cat(pad_masks, dim=1), torch.cat(token_types, 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.
"""
target_value = target_value.reshape(-1).clamp(
self.config.value_support_min, self.config.value_support_max
)
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:
"""Two-hot scalar projection between adjacent bins.
``dirac_delta`` is retained as the public configuration name for
checkpoint compatibility. This is the scalar two-hot projection, not
the full distributional Bellman projection used by C51.
Args:
target_value: [batch_size] or [batch_size, 1] target values.
Returns:
[batch_size, num_value_bins] target probability distribution.
"""
target_value = target_value.reshape(-1).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,
dtype=target_value.dtype,
)
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 at the nearest support bin for terminal states.
Args:
target_value: [batch_size] or [batch_size, 1] target values.
Returns:
[batch_size, num_value_bins] one-hot distribution at the nearest bin.
"""
target_value = target_value.reshape(-1).clamp(
self.config.value_support_min, self.config.value_support_max
)
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
is_terminal = is_terminal.reshape(-1)
if is_terminal.numel() != base_distribution.shape[0]:
raise ValueError(
f"Expected {base_distribution.shape[0]} terminal flags, got {is_terminal.numel()}"
)
terminal_distribution = self.one_hot_target(target_value)
return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution)
def _get_vlm_readout(self, batch: dict[str, Tensor]) -> Tensor:
"""Run Gemma3 image-bidirectional/text-causal attention plus value query."""
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, prefix_token_types = self.embed_prefix(
images, img_masks, text_embs, text_pad_mask
)
if self.config.stop_gradient_to_vlm:
prefix_embs = prefix_embs.detach()
batch_size, prefix_len = prefix_pad_masks.shape
device = prefix_embs.device
model_dtype = next(self.language_model.parameters()).dtype
query_ids = torch.zeros(batch_size, 1, dtype=torch.long, device=device)
query_emb = self.value_query(query_ids)
hidden_states = torch.cat([prefix_embs, query_emb], dim=1)
query_pad_mask = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
pad_masks = torch.cat([prefix_pad_masks, query_pad_mask], dim=1)
token_type_ids = torch.cat(
[prefix_token_types, torch.zeros(batch_size, 1, dtype=torch.long, device=device)],
dim=1,
)
prefix_position_ids = torch.cumsum(prefix_pad_masks.long(), dim=1) - 1
query_position_ids = prefix_pad_masks.sum(dim=1, keepdim=True).long()
position_ids = torch.cat([prefix_position_ids, query_position_ids], dim=1).clamp_min(0)
if hidden_states.dtype != model_dtype:
hidden_states = hidden_states.to(model_dtype)
attention_masks = create_causal_mask_mapping(
self._vlm_config,
inputs_embeds=hidden_states,
attention_mask=pad_masks,
past_key_values=None,
position_ids=position_ids,
token_type_ids=token_type_ids,
is_training=self.training,
)
outputs = self.language_model(
inputs_embeds=hidden_states,
attention_mask=attention_masks,
position_ids=position_ids,
use_cache=False,
)
return outputs.last_hidden_state[:, -1, :]
def _vlm_forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, Tensor]:
"""Run the VLM and value head.
Returns:
(value_logits [B, num_bins], predicted_value [B, 1])
"""
readout = self._get_vlm_readout(batch)
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,284 @@
# 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 448x448 (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.
Expects LeRobot's standard float image range [0, 1].
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] = (448, 448)
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
# Gemma3's SigLIP vision tower expects [-1, 1].
img = img * 2.0 - 1.0
if img.shape[1:3] != self.image_resolution:
img = resize_with_pad_torch(img, *self.image_resolution)
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, 448, 448] 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.vlm_pretrained_path or 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
+66
View File
@@ -20,13 +20,20 @@ from typing import Any
import torch import torch
from lerobot.configs import FeatureType
from lerobot.configs.rewards import RewardModelConfig from lerobot.configs.rewards import RewardModelConfig
from lerobot.processor import PolicyAction, PolicyProcessorPipeline from lerobot.processor import PolicyAction, PolicyProcessorPipeline
from lerobot.utils.feature_utils import dataset_to_policy_features
from .classifier.configuration_classifier import RewardClassifierConfig from .classifier.configuration_classifier import RewardClassifierConfig
from .distributional_value_function.configuration_distributional_value_function import DistributionalVFConfig
from .nanovlm_value_function.configuration_nanovlm_value_function import NanoVLMVFConfig
from .pretrained import PreTrainedRewardModel from .pretrained import PreTrainedRewardModel
from .robometer.configuration_robometer import RobometerConfig from .robometer.configuration_robometer import RobometerConfig
from .sarm.configuration_sarm import SARMConfig from .sarm.configuration_sarm import SARMConfig
from .temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
TemporalSiglipVFConfig,
)
from .topreward.configuration_topreward import TOPRewardConfig from .topreward.configuration_topreward import TOPRewardConfig
@@ -63,6 +70,24 @@ def get_reward_model_class(name: str) -> type[PreTrainedRewardModel]:
from lerobot.rewards.topreward.modeling_topreward import TOPRewardModel from lerobot.rewards.topreward.modeling_topreward import TOPRewardModel
return TOPRewardModel return TOPRewardModel
elif name == "distributional_value_function":
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
return DistributionalVFRewardModel
elif name == "temporal_siglip_value_function":
from lerobot.rewards.temporal_siglip_value_function.modeling_temporal_siglip_value_function import (
TemporalSiglipVFRewardModel,
)
return TemporalSiglipVFRewardModel
elif name == "nanovlm_value_function":
from lerobot.rewards.nanovlm_value_function.modeling_nanovlm_value_function import (
NanoVLMVFRewardModel,
)
return NanoVLMVFRewardModel
else: else:
try: try:
return _get_reward_model_cls_from_name(name=name) return _get_reward_model_cls_from_name(name=name)
@@ -96,6 +121,12 @@ def make_reward_model_config(reward_type: str, **kwargs) -> RewardModelConfig:
return RobometerConfig(**kwargs) return RobometerConfig(**kwargs)
elif reward_type == "topreward": elif reward_type == "topreward":
return TOPRewardConfig(**kwargs) return TOPRewardConfig(**kwargs)
elif reward_type == "distributional_value_function":
return DistributionalVFConfig(**kwargs)
elif reward_type == "temporal_siglip_value_function":
return TemporalSiglipVFConfig(**kwargs)
elif reward_type == "nanovlm_value_function":
return NanoVLMVFConfig(**kwargs)
else: else:
try: try:
config_cls = RewardModelConfig.get_choice_class(reward_type) config_cls = RewardModelConfig.get_choice_class(reward_type)
@@ -118,6 +149,13 @@ def make_reward_model(cfg: RewardModelConfig, **kwargs) -> PreTrainedRewardModel
Returns: Returns:
An instantiated and device-placed reward model. An instantiated and device-placed reward model.
""" """
dataset_meta = kwargs.get("dataset_meta")
if dataset_meta is not None and not cfg.input_features:
features = dataset_to_policy_features(dataset_meta.features)
cfg.input_features = {
key: feature for key, feature in features.items() if feature.type is not FeatureType.ACTION
}
reward_cls = get_reward_model_class(cfg.type) reward_cls = get_reward_model_class(cfg.type)
kwargs["config"] = cfg kwargs["config"] = cfg
@@ -192,6 +230,34 @@ def make_reward_pre_post_processors(
dataset_stats=kwargs.get("dataset_stats"), 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"),
)
elif isinstance(reward_cfg, TemporalSiglipVFConfig):
from lerobot.rewards.temporal_siglip_value_function.processor_temporal_siglip_value_function import (
make_temporal_siglip_vf_pre_post_processors,
)
return make_temporal_siglip_vf_pre_post_processors(
config=reward_cfg,
dataset_stats=kwargs.get("dataset_stats"),
)
elif isinstance(reward_cfg, NanoVLMVFConfig):
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
make_nanovlm_vf_pre_post_processors,
)
return make_nanovlm_vf_pre_post_processors(
config=reward_cfg,
dataset_stats=kwargs.get("dataset_stats"),
)
else: else:
try: try:
processors = _make_processors_from_reward_model_config( processors = _make_processors_from_reward_model_config(
@@ -0,0 +1,9 @@
from .configuration_nanovlm_value_function import NanoVLMVFConfig
from .modeling_nanovlm_value_function import NanoVLMVFRewardModel
from .processor_nanovlm_value_function import make_nanovlm_vf_pre_post_processors
__all__ = [
"NanoVLMVFConfig",
"NanoVLMVFRewardModel",
"make_nanovlm_vf_pre_post_processors",
]
@@ -0,0 +1,46 @@
"""Configuration for the pretrained nanoVLM-460M value-function experiment."""
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("nanovlm_value_function")
@dataclass
class NanoVLMVFConfig(RewardModelConfig):
nanovlm_pretrained_path: str = "lusxvr/nanoVLM-460M-8k"
nanovlm_code_path: str = "third_party/nanoVLM"
# The checkpoint was aligned with an 8k context. Native image tiling can
# require thousands of placeholder tokens for several robot cameras.
tokenizer_max_length: int = 8192
num_value_bins: int = 201
value_support_min: float = -1.0
value_support_max: float = 0.0
hl_gauss_sigma_ratio: float = 0.75
target_method: str = "dirac_delta"
use_one_hot_terminal: bool = True
value_dropout: float = 0.0
freeze_vision_encoder: bool = True
freeze_multimodal_projector: bool = True
freeze_language_model: bool = True
normalization_mapping: dict[str, NormalizationMode] = field(
default_factory=lambda: {"VISUAL": NormalizationMode.IDENTITY}
)
def validate_features(self) -> None:
if not any(feature.type == FeatureType.VISUAL for feature in self.input_features.values()):
raise ValueError("NanoVLMVFConfig requires visual input features")
def get_optimizer_preset(self) -> AdamWConfig:
return AdamWConfig(lr=1e-4, weight_decay=1e-4, grad_clip_norm=1.0)
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
return CosineDecayWithWarmupSchedulerConfig(
num_warmup_steps=500,
num_decay_steps=40000,
peak_lr=1e-4,
decay_lr=1e-6,
)
@@ -0,0 +1,115 @@
"""Distributional value head on the pretrained nanoVLM-460M checkpoint."""
from __future__ import annotations
import sys
from pathlib import Path
from typing import Any
import torch
from torch import Tensor, nn
from lerobot.rewards.distributional_value_function.common import DistributionalValueMixin
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import ValueHead
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
NANOVLM_ATTENTION_MASK,
NANOVLM_IMAGES,
NANOVLM_INPUT_IDS,
)
from lerobot.rewards.pretrained import PreTrainedRewardModel
from .configuration_nanovlm_value_function import NanoVLMVFConfig
class NanoVLMVFRewardModel(DistributionalValueMixin, PreTrainedRewardModel):
"""Use nanoVLM's aligned decoder readout for RECAP return classification."""
name = "nanovlm_value_function"
config_class = NanoVLMVFConfig
def __init__(self, config: NanoVLMVFConfig, **kwargs):
super().__init__(config)
self.config = config
config.validate_features()
code_path = Path(config.nanovlm_code_path)
if not code_path.is_absolute():
code_path = Path(__file__).resolve().parents[4] / code_path
if not code_path.exists():
raise FileNotFoundError(f"nanoVLM code not found at {code_path}")
if str(code_path) not in sys.path:
sys.path.insert(0, str(code_path))
from models.vision_language_model import VisionLanguageModel
self.nanovlm = VisionLanguageModel.from_pretrained(config.nanovlm_pretrained_path)
hidden_size = self.nanovlm.cfg.lm_hidden_dim
self.value_query = nn.Embedding(1, hidden_size)
nn.init.normal_(self.value_query.weight, std=0.02)
self.value_head = ValueHead(
hidden_size,
config.num_value_bins,
config.value_support_min,
config.value_support_max,
config.value_dropout,
)
bin_width = (config.value_support_max - config.value_support_min) / (config.num_value_bins - 1)
self.hl_gauss_sigma = config.hl_gauss_sigma_ratio * bin_width
self._set_requires_grad()
def _set_requires_grad(self):
if self.config.freeze_vision_encoder:
self.nanovlm.vision_encoder.requires_grad_(False).eval()
if self.config.freeze_multimodal_projector:
self.nanovlm.MP.requires_grad_(False).eval()
if self.config.freeze_language_model:
self.nanovlm.decoder.requires_grad_(False).eval()
def train(self, mode: bool = True):
super().train(mode)
if self.config.freeze_vision_encoder:
self.nanovlm.vision_encoder.eval()
if self.config.freeze_multimodal_projector:
self.nanovlm.MP.eval()
if self.config.freeze_language_model:
self.nanovlm.decoder.eval()
return self
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
return self._distributional_forward(batch)
def _get_value_readout(self, batch: dict[str, Tensor]) -> Tensor:
input_ids = batch[NANOVLM_INPUT_IDS]
attention_mask = batch[NANOVLM_ATTENTION_MASK].bool()
batch_size = input_ids.shape[0]
images = self.nanovlm._process_images(batch[NANOVLM_IMAGES], input_ids.device)
text_tokens = self.nanovlm.decoder.token_embedding(input_ids)
if images is not None:
image_tokens = self.nanovlm.MP(self.nanovlm.vision_encoder(images))
placeholder_count = (input_ids == self.nanovlm.tokenizer.image_token_id).sum().item()
image_token_count = image_tokens.shape[0] * image_tokens.shape[1]
if placeholder_count != image_token_count:
raise ValueError(
"nanoVLM image placeholders do not match projected image tokens: "
f"{placeholder_count} placeholders versus {image_token_count} tokens. "
"The prompt may have been truncated; increase tokenizer_max_length."
)
text_tokens = self.nanovlm._replace_img_tokens_with_embd(
input_ids,
text_tokens,
image_tokens,
)
query = self.value_query(torch.zeros(batch_size, 1, dtype=torch.long, device=text_tokens.device)).to(
text_tokens.dtype
)
inputs = torch.cat([text_tokens, query], dim=1)
attention_mask = torch.cat(
[
attention_mask,
torch.ones(batch_size, 1, dtype=torch.bool, device=text_tokens.device),
],
dim=1,
)
hidden, _ = self.nanovlm.decoder(inputs, attention_mask=attention_mask)
return hidden[:, -1]
def get_optim_params(self):
return [parameter for parameter in self.parameters() if parameter.requires_grad]
@@ -0,0 +1,214 @@
"""Processor using nanoVLM's native image splitting and chat-token layout."""
import json
import sys
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import torch
from torchvision.transforms.functional import to_pil_image
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
from lerobot.processor import (
AddBatchDimensionProcessorStep,
ComplementaryDataProcessorStep,
DeviceProcessorStep,
NormalizerProcessorStep,
PolicyAction,
PolicyProcessorPipeline,
ProcessorStepRegistry,
RenameObservationsProcessorStep,
batch_to_transition,
policy_action_to_transition,
transition_to_batch,
)
from lerobot.types import TransitionKey
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
from .configuration_nanovlm_value_function import NanoVLMVFConfig
NANOVLM_IMAGES = "observation.nanovlm.images"
NANOVLM_INPUT_IDS = "observation.nanovlm.input_ids"
NANOVLM_ATTENTION_MASK = "observation.nanovlm.attention_mask"
@ProcessorStepRegistry.register(name="nanovlm_native_processor")
@dataclass
class NanoVLMNativeProcessorStep(ComplementaryDataProcessorStep):
pretrained_path: str
code_path: str
image_keys: tuple[str, ...]
max_length: int
_tokenizer: Any = field(default=None, init=False, repr=False)
_image_processor: Any = field(default=None, init=False, repr=False)
_get_image_string: Any = field(default=None, init=False, repr=False)
_mp_image_token_length: int = field(default=64, init=False, repr=False)
def __post_init__(self):
code_path = Path(self.code_path)
if not code_path.is_absolute():
code_path = Path(__file__).resolve().parents[4] / code_path
if str(code_path) not in sys.path:
sys.path.insert(0, str(code_path))
from data.processors import get_image_processor, get_image_string, get_tokenizer
config_path = _resolve_checkpoint_file(self.pretrained_path, "config.json")
config = json.loads(Path(config_path).read_text())
if self.max_length > config["lm_max_length"]:
raise ValueError(
f"tokenizer_max_length={self.max_length} exceeds nanoVLM's "
f"lm_max_length={config['lm_max_length']}"
)
self._tokenizer = get_tokenizer(
config["lm_tokenizer"],
config["vlm_extra_tokens"],
config["lm_chat_template"],
)
self._image_processor = get_image_processor(
config["max_img_size"],
config["vit_img_size"],
config["resize_to_max_side_len"],
)
self._get_image_string = get_image_string
self._mp_image_token_length = config["mp_image_token_length"]
def complementary_data(self, complementary_data):
raw_tasks = complementary_data.get("task")
if raw_tasks is None:
raise ValueError("Task is required for nanoVLM value processing")
observation = self.transition[TransitionKey.OBSERVATION]
present_image_keys = [key for key in self.image_keys if key in observation]
if not present_image_keys:
raise ValueError("No configured nanoVLM image key is present in the observation")
batch_size = observation[present_image_keys[0]].shape[0]
tasks = [raw_tasks] * batch_size if isinstance(raw_tasks, str) else list(raw_tasks)
if len(tasks) != batch_size:
raise ValueError(f"Received {len(tasks)} tasks for an image batch of size {batch_size}")
processed_batch = []
input_rows = []
attention_rows = []
for batch_index in range(batch_size):
processed_images = []
split_counts = []
for key in self.image_keys:
if key not in observation:
continue
image = observation[key][batch_index]
if image.ndim != 3:
raise ValueError(f"nanoVLM expects CHW images, got {tuple(image.shape)} for {key}")
if image.shape[0] not in (1, 3, 4) and image.shape[-1] in (1, 3, 4):
image = image.permute(2, 0, 1)
if image.dtype != torch.uint8:
image = image.float()
if image.min() < -1e-6 or image.max() > 1.0 + 1e-6:
raise ValueError(
f"nanoVLM expects uint8 [0,255] or float [0,1] images; "
f"{key} has range [{image.min().item()}, {image.max().item()}]"
)
image = image.clamp(0, 1)
pil_image = to_pil_image(image.cpu()).convert("RGB")
processed, split_count = self._image_processor(pil_image)
processed_images.append(processed)
split_counts.append(split_count)
image_string = self._get_image_string(
self._tokenizer,
split_counts,
self._mp_image_token_length,
)
prompt = self._tokenizer.apply_chat_template(
[{"role": "user", "content": image_string + f"Task: {tasks[batch_index]}."}],
tokenize=False,
add_generation_prompt=True,
)
tokenized = self._tokenizer(
prompt,
truncation=False,
add_special_tokens=False,
)
if len(tokenized["input_ids"]) > self.max_length:
raise ValueError(
f"nanoVLM prompt has {len(tokenized['input_ids'])} tokens, exceeding "
f"tokenizer_max_length={self.max_length}. The native nanoVLM collator "
"discards over-length examples instead of truncating image placeholders."
)
input_rows.append(tokenized["input_ids"])
attention_rows.append(tokenized.get("attention_mask", [1] * len(tokenized["input_ids"])))
processed_batch.append(processed_images)
max_length = max(map(len, input_rows))
for input_ids, attention_mask in zip(input_rows, attention_rows, strict=True):
padding = max_length - len(input_ids)
input_ids[:0] = [self._tokenizer.pad_token_id] * padding
attention_mask[:0] = [0] * padding
observation = dict(observation)
observation[NANOVLM_IMAGES] = processed_batch
observation[NANOVLM_INPUT_IDS] = torch.tensor(input_rows, dtype=torch.long)
observation[NANOVLM_ATTENTION_MASK] = torch.tensor(attention_rows, dtype=torch.bool)
self.transition[TransitionKey.OBSERVATION] = observation
return complementary_data
def transform_features(
self,
features: dict[PipelineFeatureType, dict[str, PolicyFeature]],
):
return features
def get_config(self):
return {
"pretrained_path": self.pretrained_path,
"code_path": self.code_path,
"image_keys": self.image_keys,
"max_length": self.max_length,
}
def _resolve_checkpoint_file(repo_id_or_path: str, filename: str) -> str:
local_path = Path(repo_id_or_path) / filename
if local_path.exists():
return str(local_path)
from huggingface_hub import hf_hub_download
return hf_hub_download(repo_id=repo_id_or_path, filename=filename)
def make_nanovlm_vf_pre_post_processors(
config: NanoVLMVFConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction],
]:
image_keys = tuple(
key for key, feature in config.input_features.items() if feature.type == FeatureType.VISUAL
)
preprocessor = PolicyProcessorPipeline(
steps=[
RenameObservationsProcessorStep(rename_map={}),
AddBatchDimensionProcessorStep(),
NormalizerProcessorStep(
features={**config.input_features, **config.output_features},
norm_map=config.normalization_mapping,
stats=dataset_stats,
),
NanoVLMNativeProcessorStep(
pretrained_path=config.nanovlm_pretrained_path,
code_path=config.nanovlm_code_path,
image_keys=image_keys,
max_length=config.tokenizer_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
@@ -0,0 +1,9 @@
from .configuration_temporal_siglip_value_function import TemporalSiglipVFConfig
from .modeling_temporal_siglip_value_function import TemporalSiglipVFRewardModel
from .processor_temporal_siglip_value_function import make_temporal_siglip_vf_pre_post_processors
__all__ = [
"TemporalSiglipVFConfig",
"TemporalSiglipVFRewardModel",
"make_temporal_siglip_vf_pre_post_processors",
]
@@ -0,0 +1,57 @@
"""Configuration for the experimental temporal SigLIP2 value function."""
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("temporal_siglip_value_function")
@dataclass
class TemporalSiglipVFConfig(RewardModelConfig):
siglip_path: str = "google/siglip2-so400m-patch14-384"
image_resolution: tuple[int, int] = (384, 384)
tokenizer_max_length: int = 64
history_steps: int = 6
frame_gap: int = 30
state_key: str = "observation.state"
state_dim: int = 32
hidden_size: int = 512
num_layers: int = 4
num_heads: int = 8
dropout: float = 0.1
num_value_bins: int = 201
value_support_min: float = -1.0
value_support_max: float = 0.0
hl_gauss_sigma_ratio: float = 0.75
target_method: str = "dirac_delta"
use_one_hot_terminal: bool = True
normalization_mapping: dict[str, NormalizationMode] = field(
default_factory=lambda: {
"VISUAL": NormalizationMode.IDENTITY,
"STATE": NormalizationMode.MEAN_STD,
}
)
@property
def observation_delta_indices(self) -> list[int]:
return [-self.frame_gap * index for index in range(self.history_steps - 1, -1, -1)]
def validate_features(self) -> None:
if not any(feature.type == FeatureType.VISUAL for feature in self.input_features.values()):
raise ValueError("TemporalSiglipVFConfig requires visual input features")
if self.state_key not in self.input_features:
raise ValueError(f"TemporalSiglipVFConfig requires {self.state_key!r}")
def get_optimizer_preset(self) -> AdamWConfig:
return AdamWConfig(lr=1e-4, weight_decay=1e-4, grad_clip_norm=1.0)
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
return CosineDecayWithWarmupSchedulerConfig(
num_warmup_steps=500,
num_decay_steps=40000,
peak_lr=1e-4,
decay_lr=1e-6,
)
@@ -0,0 +1,176 @@
"""Past-only temporal SigLIP2 distributional value function."""
from __future__ import annotations
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.distributional_value_function.common import DistributionalValueMixin
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import ValueHead
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
)
from lerobot.rewards.pretrained import PreTrainedRewardModel
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
from lerobot.utils.import_utils import _transformers_available, require_package
from .configuration_temporal_siglip_value_function import TemporalSiglipVFConfig
if TYPE_CHECKING or _transformers_available:
from transformers import AutoModel
else:
AutoModel = None # type: ignore[assignment]
class TemporalSiglipVFRewardModel(DistributionalValueMixin, PreTrainedRewardModel):
"""Fuse three-camera history, task, and state before causal temporal attention."""
name = "temporal_siglip_value_function"
config_class = TemporalSiglipVFConfig
def __init__(self, config: TemporalSiglipVFConfig, **kwargs):
require_package("transformers", extra="recap")
super().__init__(config)
self.config = config
config.validate_features()
self.image_keys = [
key for key, feature in config.input_features.items() if feature.type == FeatureType.VISUAL
]
self.siglip = AutoModel.from_pretrained(config.siglip_path)
self.siglip.requires_grad_(False).eval()
vision_dim = self.siglip.config.vision_config.hidden_size
text_dim = self.siglip.config.text_config.hidden_size
hidden_size = config.hidden_size
self.camera_proj = nn.Linear(vision_dim, hidden_size)
self.camera_embedding = nn.Embedding(len(self.image_keys), hidden_size)
self.task_proj = nn.Linear(text_dim, hidden_size)
self.state_proj = nn.Linear(config.state_dim, hidden_size)
self.frame_fusion = nn.Sequential(
nn.LayerNorm((len(self.image_keys) + 2) * hidden_size),
nn.Linear((len(self.image_keys) + 2) * hidden_size, hidden_size),
nn.GELU(),
)
layer = nn.TransformerEncoderLayer(
d_model=hidden_size,
nhead=config.num_heads,
dim_feedforward=4 * hidden_size,
dropout=config.dropout,
batch_first=True,
activation="gelu",
)
self.temporal_transformer = nn.TransformerEncoder(
layer,
num_layers=config.num_layers,
norm=nn.LayerNorm(hidden_size),
)
self.time_embedding = nn.Embedding(config.history_steps, hidden_size)
self.value_head = ValueHead(
hidden_size,
config.num_value_bins,
config.value_support_min,
config.value_support_max,
config.dropout,
)
bin_width = (config.value_support_max - config.value_support_min) / (config.num_value_bins - 1)
self.hl_gauss_sigma = config.hl_gauss_sigma_ratio * bin_width
def train(self, mode: bool = True):
super().train(mode)
self.siglip.eval()
return self
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
return self._distributional_forward(batch)
def _get_value_readout(self, batch: dict[str, Tensor]) -> Tensor:
images = [batch[key] for key in self.image_keys]
masks = [batch[key + IMAGE_MASK_SUFFIX].bool() for key in self.image_keys]
state = batch[self.config.state_key]
if state.ndim == 2:
state = state[:, None]
batch_size, history_steps = images[0].shape[:2]
if history_steps != self.config.history_steps:
raise ValueError(f"Expected {self.config.history_steps} frames, got {history_steps}")
camera_tokens = []
vision_dtype = next(self.siglip.vision_model.parameters()).dtype
for camera_index, (image, mask) in enumerate(zip(images, masks, strict=True)):
with torch.no_grad():
features = self.siglip.vision_model(
pixel_values=image.flatten(0, 1).to(vision_dtype),
interpolate_pos_encoding=True,
return_dict=True,
).pooler_output
features = self.camera_proj(features).unflatten(0, (batch_size, history_steps))
camera_ids = torch.full(
(batch_size, history_steps),
camera_index,
dtype=torch.long,
device=features.device,
)
camera_tokens.append(
(features + self.camera_embedding(camera_ids)) * mask[..., None].to(features.dtype)
)
with torch.no_grad():
task_features = self.siglip.text_model(
input_ids=batch[OBS_LANGUAGE_TOKENS],
attention_mask=batch[OBS_LANGUAGE_ATTENTION_MASK],
return_dict=True,
).pooler_output
task_token = self.task_proj(task_features)[:, None].expand(-1, history_steps, -1)
state = self._fit_state_dim(state).to(task_token.dtype)
state_token = self.state_proj(state)
frame_tokens = self.frame_fusion(torch.cat([*camera_tokens, task_token, state_token], -1))
frame_tokens = (
frame_tokens + self.time_embedding(torch.arange(history_steps, device=frame_tokens.device))[None]
)
frame_valid = torch.stack(masks).any(0)
attention_mask = self._make_temporal_attention_mask(frame_valid)
hidden = self.temporal_transformer(
frame_tokens,
mask=attention_mask,
)
# The history window is ordered oldest→current and the current frame is
# always the final, non-padding element.
return hidden[:, -1]
def _make_temporal_attention_mask(self, frame_valid: Tensor) -> Tensor:
"""Combine causal and padding masks without fully masked padded queries.
A left-padded causal query has no valid past keys. PyTorch's optimized
eval path returns NaNs for such rows, which then contaminate later valid
tokens. Padded queries attend only to themselves; valid queries retain
causal attention and cannot attend to padded keys.
"""
batch_size, history_steps = frame_valid.shape
causal_mask = torch.triu(
torch.ones(history_steps, history_steps, dtype=torch.bool, device=frame_valid.device),
diagonal=1,
)
attention_mask = causal_mask[None].expand(batch_size, -1, -1) | (~frame_valid)[:, None, :]
padded_queries = (~frame_valid).nonzero(as_tuple=False)
attention_mask[
padded_queries[:, 0],
padded_queries[:, 1],
padded_queries[:, 1],
] = False
return (
attention_mask[:, None]
.expand(-1, self.config.num_heads, -1, -1)
.reshape(batch_size * self.config.num_heads, history_steps, history_steps)
)
def _fit_state_dim(self, state: Tensor) -> Tensor:
if state.shape[-1] > self.config.state_dim:
return state[..., : self.config.state_dim]
return F.pad(state, (0, self.config.state_dim - state.shape[-1]))
def get_optim_params(self):
return [parameter for parameter in self.parameters() if parameter.requires_grad]
@@ -0,0 +1,185 @@
"""Processor for past-only temporal SigLIP2 value inputs."""
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
import torch
from torch import Tensor
from lerobot.configs import FeatureType
from lerobot.processor import (
AddBatchDimensionProcessorStep,
DeviceProcessorStep,
NormalizerProcessorStep,
PolicyAction,
PolicyProcessorPipeline,
ProcessorStep,
ProcessorStepRegistry,
RenameObservationsProcessorStep,
batch_to_transition,
policy_action_to_transition,
transition_to_batch,
)
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
DistributionalVFPrepareTaskPromptStep,
resize_with_pad_torch,
)
from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import (
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
POLICY_POSTPROCESSOR_DEFAULT_NAME,
POLICY_PREPROCESSOR_DEFAULT_NAME,
)
from lerobot.utils.import_utils import _transformers_available, require_package
from .configuration_temporal_siglip_value_function import TemporalSiglipVFConfig
if TYPE_CHECKING or _transformers_available:
from transformers import AutoTokenizer
else:
AutoTokenizer = None
@ProcessorStepRegistry.register(name="temporal_siglip_vf_image_processor")
@dataclass
class TemporalSiglipImageProcessorStep(ProcessorStep):
image_resolution: tuple[int, int]
image_keys: tuple[str, ...]
history_steps: int
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
observation = dict(transition.get(TransitionKey.OBSERVATION, {}))
batch_size = self._batch_size(observation)
for key in self.image_keys:
if key not in observation:
height, width = self.image_resolution
observation[key] = torch.full((batch_size, self.history_steps, 3, height, width), -1.0)
observation[key + IMAGE_MASK_SUFFIX] = torch.zeros(
batch_size, self.history_steps, dtype=torch.bool
)
continue
image = observation[key]
if image.ndim == 4:
image = image[:, None]
if image.ndim != 5 or image.shape[2] != 3:
raise ValueError(f"Expected {key} as [B,T,3,H,W], got {tuple(image.shape)}")
batch_size, history = image.shape[:2]
image = image.float() / 127.5 - 1.0 if image.dtype == torch.uint8 else image.float() * 2.0 - 1.0
image = image.flatten(0, 1).permute(0, 2, 3, 1)
if image.shape[1:3] != self.image_resolution:
image = resize_with_pad_torch(image, *self.image_resolution)
observation[key] = image.permute(0, 3, 1, 2).unflatten(0, (batch_size, history))
padding_key = f"{key}_is_pad"
if padding_key in observation:
observation[key + IMAGE_MASK_SUFFIX] = ~observation[padding_key].bool()
else:
observation[key + IMAGE_MASK_SUFFIX] = torch.ones(
batch_size, history, dtype=torch.bool, device=image.device
)
transition[TransitionKey.OBSERVATION] = observation
return transition
@staticmethod
def _batch_size(observation: dict[str, Any]) -> int:
for value in observation.values():
if isinstance(value, Tensor) and value.ndim >= 2:
return value.shape[0]
return 1
def transform_features(self, features):
return features
def get_config(self):
return {
"image_resolution": self.image_resolution,
"image_keys": self.image_keys,
"history_steps": self.history_steps,
}
@ProcessorStepRegistry.register(name="temporal_siglip_vf_tokenizer")
@dataclass
class TemporalSiglipTokenizerStep(ProcessorStep):
tokenizer_name: str
max_length: int
_tokenizer: Any = field(default=None, init=False, repr=False)
def __post_init__(self):
require_package("transformers", extra="recap")
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_name)
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
task = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get("task")
if task is None:
raise ValueError("Task is required for TemporalSiglipTokenizerStep")
task = [task] if isinstance(task, str) else list(task)
tokenized = self._tokenizer(
task,
padding="max_length",
max_length=self.max_length,
truncation=True,
return_tensors="pt",
)
input_ids = tokenized["input_ids"]
attention_mask = tokenized.get("attention_mask")
if attention_mask is None:
attention_mask = input_ids.ne(self._tokenizer.pad_token_id)
observation = dict(transition.get(TransitionKey.OBSERVATION, {}))
observation[OBS_LANGUAGE_TOKENS] = input_ids
observation[OBS_LANGUAGE_ATTENTION_MASK] = attention_mask.bool()
transition[TransitionKey.OBSERVATION] = observation
return transition
def transform_features(self, features):
return features
def get_config(self):
return {"tokenizer_name": self.tokenizer_name, "max_length": self.max_length}
def make_temporal_siglip_vf_pre_post_processors(
config: TemporalSiglipVFConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction],
]:
image_keys = tuple(
key for key, feature in config.input_features.items() if feature.type == FeatureType.VISUAL
)
preprocessor = PolicyProcessorPipeline(
steps=[
RenameObservationsProcessorStep(rename_map={}),
AddBatchDimensionProcessorStep(),
NormalizerProcessorStep(
features={**config.input_features, **config.output_features},
norm_map=config.normalization_mapping,
stats=dataset_stats,
),
TemporalSiglipImageProcessorStep(
image_resolution=config.image_resolution,
image_keys=image_keys,
history_steps=config.history_steps,
),
DistributionalVFPrepareTaskPromptStep(),
TemporalSiglipTokenizerStep(
tokenizer_name=config.siglip_path,
max_length=config.tokenizer_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
@@ -182,7 +182,7 @@ class SOFollower(Robot):
def get_observation(self) -> RobotObservation: def get_observation(self) -> RobotObservation:
# Read arm position # Read arm position
start = time.perf_counter() 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()} obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
dt_ms = (time.perf_counter() - start) * 1e3 dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read state: {dt_ms:.1f}ms") 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. # Cap goal position when too far away from present position.
# /!\ Slower fps expected due to reading from the follower. # /!\ Slower fps expected due to reading from the follower.
if self.config.max_relative_target is not None: 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_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) goal_pos = ensure_safe_goal_position(goal_present_pos, self.config.max_relative_target)
# Send goal position to the arm # 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()} return {f"{motor}.pos": val for motor, val in goal_pos.items()}
@check_if_not_connected @check_if_not_connected
+2
View File
@@ -121,6 +121,8 @@ class DAggerPedalConfig:
pause_resume: str = "KEY_A" pause_resume: str = "KEY_A"
correction: str = "KEY_B" correction: str = "KEY_B"
upload: str = "KEY_C" upload: str = "KEY_C"
success: str = "KEY_D"
failure: str = "KEY_E"
@RolloutStrategyConfig.register_subclass("episodic") @RolloutStrategyConfig.register_subclass("episodic")
+26 -1
View File
@@ -118,6 +118,9 @@ class DAggerEvents:
# Episode success labeling # Episode success labeling
self._episode_success: bool | None = None self._episode_success: bool | None = None
# Episode success labeling
self._episode_success: bool | None = None
# -- Thread-safe phase access ------------------------------------------ # -- Thread-safe phase access ------------------------------------------
@property @property
@@ -181,6 +184,23 @@ class DAggerEvents:
self._episode_success = None self._episode_success = None
return result 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
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Input device handlers # Input device handlers
@@ -243,6 +263,12 @@ def _init_dagger_pedal(events: DAggerEvents, cfg: DAggerPedalConfig):
events.request_transition(code_to_event[code]) events.request_transition(code_to_event[code])
if code == cfg.upload: if code == cfg.upload:
events.upload_requested.set() 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) logger.info("Initializing DAgger foot pedal listener (device=%s)", cfg.device_path)
return start_pedal_listener(on_press, device_path=cfg.device_path) return start_pedal_listener(on_press, device_path=cfg.device_path)
@@ -365,7 +391,6 @@ class DAggerStrategy(RolloutStrategy):
return return
label = self._events.consume_episode_success() label = self._events.consume_episode_success()
logger.info("_stamp_episode_success: label=%s, buffer_len=%d", label, len(success_buf))
if label: if label:
success_buf[-1] = np.array([True], dtype=bool) success_buf[-1] = np.array([True], dtype=bool)
+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.executor import Executor
from lerobot.annotations.steerable_pipeline.frames import make_frame_provider from lerobot.annotations.steerable_pipeline.frames import make_frame_provider
from lerobot.annotations.steerable_pipeline.modules import ( from lerobot.annotations.steerable_pipeline.modules import (
AdvantageModule,
GeneralVqaModule, GeneralVqaModule,
InterjectionsAndSpeechModule, InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule, PlanSubtasksMemoryModule,
@@ -47,6 +48,7 @@ logger = logging.getLogger(__name__)
def _resolve_root(cfg: AnnotationPipelineConfig) -> Path: def _resolve_root(cfg: AnnotationPipelineConfig) -> Path:
"""Resolve the dataset root by downloading the full snapshot if needed."""
if cfg.root is not None: if cfg.root is not None:
return Path(cfg.root) return Path(cfg.root)
if cfg.repo_id is not None: if cfg.repo_id is not None:
@@ -63,17 +65,24 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
root = _resolve_root(cfg) root = _resolve_root(cfg)
logger.info("annotate: root=%s", root) logger.info("annotate: root=%s", root)
vlm = make_vlm_client(cfg.vlm) needs_vlm = cfg.plan.enabled or cfg.interjections.enabled or cfg.vqa.enabled
frame_provider = make_frame_provider(root, camera_key=cfg.vlm.camera_key, video_backend=cfg.video_backend) 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 # 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 # is obvious in job output rather than discovered post-hoc by counting
# parquet rows. # parquet rows.
cam_keys = list(getattr(frame_provider, "camera_keys", []) or []) cam_keys = list(getattr(frame_provider, "camera_keys", []) or []) if frame_provider else []
logger.info( if frame_provider:
"annotate: frame_provider default camera=%r, all cameras=%s", logger.info(
getattr(frame_provider, "camera_key", None), "annotate: frame_provider default camera=%r, all cameras=%s",
cam_keys, getattr(frame_provider, "camera_key", None),
) cam_keys,
)
if cfg.vqa.enabled and not cam_keys: if cfg.vqa.enabled and not cam_keys:
logger.warning( logger.warning(
"annotate: the vqa module is enabled but no cameras were " "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 " "meta/info.json for observation.images.* features, or pass "
"--vlm.camera_key=<key> to seed the cameras list." "--vlm.camera_key=<key> to seed the cameras list."
) )
plan = PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider) plan = (
interjections = InterjectionsAndSpeechModule( PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan, frame_provider=frame_provider)
vlm=vlm, config=cfg.interjections, seed=cfg.seed, 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() writer = LanguageColumnsWriter()
validator = StagingValidator( validator = StagingValidator(
dataset_camera_keys=tuple(getattr(frame_provider, "camera_keys", []) or []) or None, dataset_camera_keys=tuple(cam_keys) or None,
) )
executor = Executor( executor = Executor(
@@ -96,6 +121,7 @@ def annotate(cfg: AnnotationPipelineConfig) -> None:
plan=plan, plan=plan,
interjections=interjections, interjections=interjections,
vqa=vqa, vqa=vqa,
advantage=advantage,
writer=writer, writer=writer,
validator=validator, 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()
@@ -0,0 +1,637 @@
#!/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.
"""Create a modern episode video with RECAP value and advantage overlays."""
from __future__ import annotations
import argparse
import logging
import random
import shutil
import subprocess
from dataclasses import dataclass
from pathlib import Path
import cv2
import numpy as np
import pandas as pd
from huggingface_hub import hf_hub_download, snapshot_download
from lerobot.datasets import LeRobotDatasetMetadata
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class OverlayTheme:
background: tuple[int, int, int] = (22, 24, 28)
surface: tuple[int, int, int] = (34, 37, 43)
text: tuple[int, int, int] = (242, 244, 248)
muted: tuple[int, int, int] = (158, 164, 176)
positive: tuple[int, int, int] = (103, 198, 94)
negative: tuple[int, int, int] = (91, 91, 235)
intervention: tuple[int, int, int] = (72, 184, 238)
value: tuple[int, int, int] = (230, 178, 76)
target: tuple[int, int, int] = (190, 194, 204)
raw_value: tuple[int, int, int] = (104, 111, 125)
marker: tuple[int, int, int] = (255, 255, 255)
THEME = OverlayTheme()
def _fit_text(text: str, max_width: int, font_scale: float, thickness: int = 1) -> str:
if cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, font_scale, thickness)[0][0] <= max_width:
return text
suffix = "..."
while (
text
and cv2.getTextSize(
text + suffix,
cv2.FONT_HERSHEY_SIMPLEX,
font_scale,
thickness,
)[0][0]
> max_width
):
text = text[:-1]
return text + suffix
def _draw_text(
image: np.ndarray,
text: str,
position: tuple[int, int],
*,
scale: float,
color: tuple[int, int, int],
thickness: int = 1,
) -> None:
cv2.putText(
image,
text,
position,
cv2.FONT_HERSHEY_SIMPLEX,
scale,
color,
thickness,
cv2.LINE_AA,
)
def _rounded_rectangle(
image: np.ndarray,
top_left: tuple[int, int],
bottom_right: tuple[int, int],
color: tuple[int, int, int],
radius: int,
) -> None:
x0, y0 = top_left
x1, y1 = bottom_right
radius = min(radius, (x1 - x0) // 2, (y1 - y0) // 2)
cv2.rectangle(image, (x0 + radius, y0), (x1 - radius, y1), color, -1)
cv2.rectangle(image, (x0, y0 + radius), (x1, y1 - radius), color, -1)
for center in (
(x0 + radius, y0 + radius),
(x1 - radius, y0 + radius),
(x0 + radius, y1 - radius),
(x1 - radius, y1 - radius),
):
cv2.circle(image, center, radius, color, -1, cv2.LINE_AA)
def _alpha_rounded_rectangle(
image: np.ndarray,
top_left: tuple[int, int],
bottom_right: tuple[int, int],
color: tuple[int, int, int],
*,
alpha: float,
radius: int,
) -> None:
overlay = image.copy()
_rounded_rectangle(overlay, top_left, bottom_right, color, radius)
cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, dst=image)
def _contiguous_segments(labels: np.ndarray) -> list[tuple[int, int, str]]:
if len(labels) == 0:
return []
segments: list[tuple[int, int, str]] = []
start = 0
for index in range(1, len(labels)):
if labels[index] != labels[start]:
segments.append((start, index, str(labels[start])))
start = index
segments.append((start, len(labels), str(labels[start])))
return segments
def _draw_timeline(
dashboard: np.ndarray,
labels: np.ndarray,
interventions: np.ndarray,
current_index: int,
*,
x0: int,
x1: int,
y0: int,
y1: int,
) -> None:
_rounded_rectangle(dashboard, (x0, y0), (x1, y1), THEME.surface, 5)
width = x1 - x0
count = max(len(labels), 1)
for start, end, label in _contiguous_segments(labels):
segment_x0 = x0 + round(width * start / count)
segment_x1 = x0 + round(width * end / count)
color = THEME.positive if label == "positive" else THEME.negative
cv2.rectangle(
dashboard,
(segment_x0, y0 + 2),
(max(segment_x0 + 1, segment_x1), y1 - 2),
color,
-1,
)
for index in np.flatnonzero(interventions):
x = x0 + round(width * index / count)
cv2.line(dashboard, (x, y0), (x, y1), THEME.intervention, 2, cv2.LINE_AA)
marker_x = x0 + round(width * current_index / max(count - 1, 1))
cv2.line(dashboard, (marker_x, y0 - 3), (marker_x, y1 + 3), THEME.marker, 2, cv2.LINE_AA)
def _value_to_y(value: float, y0: int, y1: int) -> int:
normalized = float(np.clip(value, -1.0, 0.0) + 1.0)
return round(y1 - normalized * (y1 - y0))
def _draw_curve(
image: np.ndarray,
values: np.ndarray,
color: tuple[int, int, int],
*,
x0: int,
x1: int,
y0: int,
y1: int,
thickness: int,
) -> None:
if len(values) < 2:
return
x = np.linspace(x0, x1, len(values)).astype(np.int32)
y = np.asarray([_value_to_y(float(value), y0, y1) for value in values], dtype=np.int32)
points = np.stack((x, y), axis=1).reshape(-1, 1, 2)
cv2.polylines(image, [points], False, color, thickness, cv2.LINE_AA)
def _draw_dashboard(
width: int,
height: int,
episode: pd.DataFrame,
current_index: int,
fps: float,
task: str,
) -> np.ndarray:
dashboard = np.full((height, width, 3), THEME.background, dtype=np.uint8)
padding = max(14, width // 45)
value = float(episode.iloc[current_index]["predicted_value"])
advantage = float(episode.iloc[current_index]["advantage"])
time_s = current_index / fps
duration_s = max((len(episode) - 1) / fps, 0)
task_text = _fit_text(task, int(width * 0.52), 0.48)
_draw_text(
dashboard,
task_text,
(padding, 24),
scale=0.48,
color=THEME.text,
thickness=1,
)
metric_text = f"VALUE {value:+.3f} ADV {advantage:+.3f}"
_draw_text(
dashboard,
metric_text,
(padding, 49),
scale=0.5,
color=THEME.text,
thickness=1,
)
time_text = f"{time_s:05.1f}s / {duration_s:05.1f}s"
time_width = cv2.getTextSize(time_text, cv2.FONT_HERSHEY_SIMPLEX, 0.48, 1)[0][0]
_draw_text(
dashboard,
time_text,
(width - padding - time_width, 24),
scale=0.48,
color=THEME.muted,
)
labels = episode["advantage_label"].astype(str).to_numpy()
interventions = (
episode["intervention"].astype(bool).to_numpy()
if "intervention" in episode
else np.zeros(len(episode), dtype=bool)
)
timeline_y0 = 59
timeline_y1 = 73
_draw_timeline(
dashboard,
labels,
interventions,
current_index,
x0=padding,
x1=width - padding,
y0=timeline_y0,
y1=timeline_y1,
)
curve_y0 = 82
curve_y1 = height - 13
cv2.line(
dashboard,
(padding, _value_to_y(0.0, curve_y0, curve_y1)),
(width - padding, _value_to_y(0.0, curve_y0, curve_y1)),
THEME.surface,
1,
cv2.LINE_AA,
)
_draw_curve(
dashboard,
episode["mc_return"].to_numpy(float),
THEME.target,
x0=padding,
x1=width - padding,
y0=curve_y0,
y1=curve_y1,
thickness=1,
)
if "predicted_value_raw" in episode:
_draw_curve(
dashboard,
episode["predicted_value_raw"].to_numpy(float),
THEME.raw_value,
x0=padding,
x1=width - padding,
y0=curve_y0,
y1=curve_y1,
thickness=1,
)
_draw_curve(
dashboard,
episode["predicted_value"].to_numpy(float),
THEME.value,
x0=padding,
x1=width - padding,
y0=curve_y0,
y1=curve_y1,
thickness=2,
)
marker_x = padding + round((width - 2 * padding) * current_index / max(len(episode) - 1, 1))
marker_y = _value_to_y(value, curve_y0, curve_y1)
cv2.circle(dashboard, (marker_x, marker_y), 4, THEME.marker, -1, cv2.LINE_AA)
return dashboard
def _draw_label_timeline(
frame: np.ndarray,
labels: np.ndarray,
current_index: int,
) -> None:
height, width = frame.shape[:2]
horizontal_margin = max(18, round(width * 0.04))
bar_height = max(26, round(height * 0.055))
bottom_margin = max(14, round(height * 0.03))
x0 = horizontal_margin
x1 = width - horizontal_margin
y1 = height - bottom_margin
y0 = y1 - bar_height
radius = max(6, bar_height // 3)
_alpha_rounded_rectangle(
frame,
(x0, y0),
(x1, y1),
THEME.background,
alpha=0.36,
radius=radius,
)
inset = max(3, bar_height // 8)
inner_x0, inner_x1 = x0 + inset, x1 - inset
inner_y0, inner_y1 = y0 + inset, y1 - inset
timeline_width = inner_x1 - inner_x0
count = max(len(labels), 1)
overlay = frame.copy()
for start, end, label in _contiguous_segments(labels):
segment_x0 = inner_x0 + round(timeline_width * start / count)
segment_x1 = inner_x0 + round(timeline_width * end / count)
color = THEME.positive if label == "positive" else THEME.negative
cv2.rectangle(
overlay,
(segment_x0, inner_y0),
(max(segment_x0 + 1, segment_x1), inner_y1),
color,
-1,
)
cv2.addWeighted(overlay, 0.56, frame, 0.44, 0, dst=frame)
marker_x = inner_x0 + round(timeline_width * current_index / max(len(labels) - 1, 1))
cv2.line(
frame,
(marker_x, y0 - 3),
(marker_x, y1 + 3),
THEME.marker,
2,
cv2.LINE_AA,
)
def _resolve_dataset(
repo_id: str,
root: Path | None,
) -> tuple[Path, LeRobotDatasetMetadata]:
if root is None:
root = Path(
snapshot_download(
repo_id=repo_id,
repo_type="dataset",
allow_patterns=["meta/**"],
)
)
metadata = LeRobotDatasetMetadata(repo_id="local", root=root)
return root, metadata
def _resolve_video_path(
repo_id: str,
root: Path,
relative_path: Path,
local_root: bool,
) -> Path:
path = root / relative_path
if path.is_file():
return path
if local_root:
raise FileNotFoundError(f"Video not found: {path}")
return Path(
hf_hub_download(
repo_id=repo_id,
repo_type="dataset",
filename=str(relative_path),
)
)
def _encode_h264(temp_path: Path, output_path: Path) -> None:
ffmpeg = shutil.which("ffmpeg")
if ffmpeg is None:
temp_path.replace(output_path)
logger.warning("ffmpeg not found; kept OpenCV mp4v output at %s", output_path)
return
result = subprocess.run( # nosec B603
[
ffmpeg,
"-y",
"-i",
str(temp_path),
"-an",
"-c:v",
"libx264",
"-preset",
"medium",
"-crf",
"18",
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
str(output_path),
],
capture_output=True,
text=True,
)
if result.returncode != 0:
logger.warning("H.264 encoding failed; retaining mp4v output: %s", result.stderr[-500:])
temp_path.replace(output_path)
return
temp_path.unlink(missing_ok=True)
def _create_decode_proxy(
source_path: Path,
proxy_path: Path,
*,
start_timestamp: float,
end_timestamp: float,
fps: float,
expected_frames: int,
) -> Path:
"""Decode an episode segment with system FFmpeg into an OpenCV-safe proxy."""
ffmpeg = shutil.which("ffmpeg")
if ffmpeg is None:
raise RuntimeError(
"System ffmpeg is required to decode this dataset's video codec. "
"Install ffmpeg and rerun the command."
)
duration = end_timestamp - start_timestamp
result = subprocess.run( # nosec B603
[
ffmpeg,
"-y",
"-ss",
f"{start_timestamp:.9f}",
"-i",
str(source_path),
"-t",
f"{duration:.9f}",
"-an",
"-vf",
f"fps={fps:.9f}",
"-frames:v",
str(expected_frames),
"-c:v",
"libx264",
"-preset",
"ultrafast",
"-crf",
"12",
"-pix_fmt",
"yuv420p",
str(proxy_path),
],
capture_output=True,
text=True,
)
if result.returncode != 0 or not proxy_path.is_file():
raise RuntimeError(f"FFmpeg episode decode failed:\n{result.stderr[-1500:]}")
return proxy_path
def create_advantage_video(
*,
repo_id: str,
predictions_path: Path,
episode_index: int,
camera_key: str | None,
output_dir: Path,
root: Path | None = None,
) -> Path:
predictions = pd.read_csv(predictions_path)
required = {
"episode_index",
"frame_index",
"mc_return",
"predicted_value",
"advantage",
"advantage_label",
}
missing = required.difference(predictions.columns)
if missing:
raise ValueError(f"Predictions CSV is missing required columns: {sorted(missing)}")
episode = (
predictions[predictions["episode_index"] == episode_index]
.sort_values("frame_index")
.reset_index(drop=True)
)
if episode.empty:
raise ValueError(f"Episode {episode_index} is absent from {predictions_path}")
local_root = root is not None
root, metadata = _resolve_dataset(repo_id, root)
available_cameras = list(metadata.video_keys)
if not available_cameras:
raise ValueError("Dataset has no video features")
if camera_key is None:
camera_key = available_cameras[0]
if camera_key not in available_cameras:
raise ValueError(f"Unknown camera {camera_key!r}; available cameras: {available_cameras}")
relative_video_path = metadata.get_video_file_path(episode_index, camera_key)
video_path = _resolve_video_path(repo_id, root, relative_video_path, local_root)
episode_metadata = metadata.episodes[episode_index]
start_timestamp = float(episode_metadata[f"videos/{camera_key}/from_timestamp"])
end_timestamp = float(episode_metadata[f"videos/{camera_key}/to_timestamp"])
fps = float(metadata.fps)
output_dir.mkdir(parents=True, exist_ok=True)
camera_name = camera_key.replace(".", "_").replace("/", "_")
output_path = output_dir / f"episode_{episode_index:06d}_{camera_name}_advantage.mp4"
decode_proxy_path = output_path.with_name(output_path.stem + "_decode_proxy.mp4")
_create_decode_proxy(
video_path,
decode_proxy_path,
start_timestamp=start_timestamp,
end_timestamp=end_timestamp,
fps=fps,
expected_frames=len(episode),
)
capture = cv2.VideoCapture(str(decode_proxy_path))
width = int(capture.get(cv2.CAP_PROP_FRAME_WIDTH))
frame_height = int(capture.get(cv2.CAP_PROP_FRAME_HEIGHT))
if width <= 0 or frame_height <= 0:
capture.release()
decode_proxy_path.unlink(missing_ok=True)
raise RuntimeError(f"Could not read video dimensions from {video_path}")
temp_path = output_path.with_name(output_path.stem + "_temp.mp4")
writer = cv2.VideoWriter(
str(temp_path),
cv2.VideoWriter_fourcc(*"mp4v"),
fps,
(width, frame_height),
)
if not writer.isOpened():
capture.release()
decode_proxy_path.unlink(missing_ok=True)
raise RuntimeError(f"Could not open video writer for {temp_path}")
expected_frames = len(episode)
logger.info(
"Rendering episode %d, camera=%s, frames=%d, source interval=%.3f%.3fs",
episode_index,
camera_key,
expected_frames,
start_timestamp,
end_timestamp,
)
written = 0
labels = episode["advantage_label"].astype(str).to_numpy()
try:
for index in range(expected_frames):
ok, frame = capture.read()
if not ok:
logger.warning("Video ended after %d/%d frames", index, expected_frames)
break
_draw_label_timeline(
frame,
labels,
index,
)
writer.write(frame)
written += 1
finally:
writer.release()
capture.release()
decode_proxy_path.unlink(missing_ok=True)
if written != expected_frames:
raise RuntimeError(
f"Rendered {written} frames but predictions contain {expected_frames}; "
"video/prediction alignment is incomplete"
)
_encode_h264(temp_path, output_path)
logger.info("Wrote %s", output_path)
return output_path
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--repo-id", required=True)
parser.add_argument("--predictions-path", type=Path, required=True)
parser.add_argument("--episode", type=int)
parser.add_argument("--camera-key")
parser.add_argument("--all-cameras", action="store_true")
parser.add_argument("--root", type=Path)
parser.add_argument("--output-dir", type=Path, default=Path("advantage_videos"))
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
predictions = pd.read_csv(args.predictions_path, usecols=["episode_index"])
episode = args.episode
if episode is None:
episode = random.Random(args.seed).choice(sorted(predictions["episode_index"].unique()))
logger.info("Randomly selected episode %d (seed=%d)", episode, args.seed)
if args.all_cameras:
_, metadata = _resolve_dataset(args.repo_id, args.root)
camera_keys: list[str | None] = list(metadata.video_keys)
else:
camera_keys = [args.camera_key]
for camera_key in camera_keys:
create_advantage_video(
repo_id=args.repo_id,
predictions_path=args.predictions_path,
episode_index=episode,
camera_key=camera_key,
output_dir=args.output_dir,
root=args.root,
)
if __name__ == "__main__":
main()
@@ -290,6 +290,8 @@ class MergeConfig(OperationConfig):
# When False, keep one file per source file instead of packing into shards. # When False, keep one file per source file instead of packing into shards.
concatenate_videos: bool = True concatenate_videos: bool = True
concatenate_data: bool = True concatenate_data: bool = True
# Allow merging datasets with different feature sets (union + fill defaults).
lenient: bool = False
@OperationConfig.register_subclass("remove_feature") @OperationConfig.register_subclass("remove_feature")
@@ -498,6 +500,7 @@ def handle_merge(cfg: EditDatasetConfig) -> None:
output_dir=output_dir, output_dir=output_dir,
concatenate_videos=cfg.operation.concatenate_videos, concatenate_videos=cfg.operation.concatenate_videos,
concatenate_data=cfg.operation.concatenate_data, concatenate_data=cfg.operation.concatenate_data,
lenient=cfg.operation.lenient,
) )
logging.info(f"Merged dataset saved to {output_dir}") logging.info(f"Merged dataset saved to {output_dir}")
@@ -0,0 +1,362 @@
#!/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.
"""Evaluate a trained reward model against per-frame MC returns.
The output CSV is also accepted by ``lerobot-annotate
--advantage.predictions_path=...``. This keeps model inference faithful to the
normal LeRobot dataset and processor path, including multi-camera observations,
state normalization, and temporal delta windows.
"""
from __future__ import annotations
import argparse
import logging
import math
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from huggingface_hub import snapshot_download
from torch.utils.data import DataLoader
from tqdm import tqdm
from lerobot.configs.rewards import RewardModelConfig
from lerobot.datasets import LeRobotDataset, LeRobotDatasetMetadata
from lerobot.datasets.factory import resolve_delta_timestamps
from lerobot.processor import PolicyProcessorPipeline, batch_to_transition, transition_to_batch
from lerobot.rewards import make_reward_model
from lerobot.utils.collate import lerobot_collate_fn
from lerobot.utils.constants import ACTION, OBS_STATE, POLICY_PREPROCESSOR_DEFAULT_NAME
from lerobot.utils.hub import find_latest_hub_checkpoint
logger = logging.getLogger(__name__)
def _resolve_model_path(path_or_repo: str) -> Path:
path = Path(path_or_repo)
if path.is_dir():
if (path / "pretrained_model").is_dir():
return path / "pretrained_model"
return path
latest = find_latest_hub_checkpoint(path_or_repo)
if latest is None:
snapshot = Path(snapshot_download(repo_id=path_or_repo, repo_type="model"))
return snapshot
snapshot = Path(
snapshot_download(
repo_id=path_or_repo,
repo_type="model",
allow_patterns=f"{latest}/pretrained_model/*",
)
)
return snapshot / latest / "pretrained_model"
def _correlation(x: np.ndarray, y: np.ndarray) -> float:
if len(x) < 2 or np.std(x) == 0 or np.std(y) == 0:
return float("nan")
return float(np.corrcoef(x, y)[0, 1])
def _spearman(x: np.ndarray, y: np.ndarray) -> float:
return _correlation(
pd.Series(x).rank(method="average").to_numpy(), pd.Series(y).rank(method="average").to_numpy()
)
def _binary_auc(scores: np.ndarray, labels: np.ndarray) -> float:
labels = labels.astype(bool)
num_positive = int(labels.sum())
num_negative = len(labels) - num_positive
if num_positive == 0 or num_negative == 0:
return float("nan")
ranks = pd.Series(scores).rank(method="average").to_numpy()
rank_sum_positive = float(ranks[labels].sum())
return (rank_sum_positive - num_positive * (num_positive + 1) / 2) / (num_positive * num_negative)
def _held_out_episodes(metadata: LeRobotDatasetMetadata, eval_split: float) -> list[int] | None:
if eval_split == 0:
return None
if not 0 < eval_split < 1:
raise ValueError(f"eval_split must be in [0,1), got {eval_split}")
task_to_episodes: dict[str, list[int]] = {}
episode_tasks = metadata.episodes["tasks"]
for episode_index in range(metadata.total_episodes):
task = episode_tasks[episode_index][0] if episode_tasks[episode_index] else ""
task_to_episodes.setdefault(task, []).append(episode_index)
held_out: list[int] = []
for episodes in task_to_episodes.values():
count = math.ceil(len(episodes) * eval_split)
held_out.extend(episodes[-count:])
return held_out
def _compute_advantages(
target: np.ndarray,
prediction: np.ndarray,
episode_index: np.ndarray,
n_step: int | None,
) -> np.ndarray:
if n_step is None:
return target - prediction
advantage = np.empty_like(target)
for episode in np.unique(episode_index):
indices = np.flatnonzero(episode_index == episode)
for local_index, index in enumerate(indices):
bootstrap_index = local_index + n_step
if bootstrap_index >= len(indices):
advantage[index] = target[index] - prediction[index]
else:
future = indices[bootstrap_index]
advantage[index] = target[index] - target[future] + prediction[future] - prediction[index]
return advantage
def _target_distribution(model, target: torch.Tensor, is_terminal: torch.Tensor) -> torch.Tensor:
try:
return model.compute_target_distribution(
target,
is_terminal,
method=model.config.target_method,
use_one_hot_terminal=model.config.use_one_hot_terminal,
)
except TypeError:
return model.compute_target_distribution(target, is_terminal)
def _predict_logits_and_value(model, batch: dict) -> tuple[torch.Tensor, torch.Tensor]:
if hasattr(model, "_vlm_forward"):
logits, predicted = model._vlm_forward(batch)
return logits, predicted.reshape(-1)
if hasattr(model, "_get_value_readout"):
logits = model.value_head(model._get_value_readout(batch))
probabilities = logits.softmax(-1)
centers = model.value_head.bin_centers.to(probabilities.dtype)
return logits, (probabilities * centers).sum(-1)
raise TypeError(f"{type(model).__name__} does not expose a distributional value readout")
def evaluate(args: argparse.Namespace) -> None:
device = torch.device(args.device)
model_path = _resolve_model_path(args.reward_model_path)
logger.info("Loading reward model from %s", model_path)
config = RewardModelConfig.from_pretrained(model_path)
config.pretrained_path = str(model_path)
config.device = device.type
metadata = LeRobotDatasetMetadata(args.dataset_repo_id, root=args.root)
delta_timestamps = resolve_delta_timestamps(config, metadata)
episodes = _held_out_episodes(metadata, args.eval_split)
if episodes is not None:
logger.info("Evaluating %d held-out episode(s) (eval_split=%s)", len(episodes), args.eval_split)
dataset = LeRobotDataset(
args.dataset_repo_id,
root=args.root,
episodes=episodes,
delta_timestamps=delta_timestamps,
video_backend=args.video_backend,
return_uint8=True,
)
dataloader = DataLoader(
dataset,
batch_size=args.batch_size,
shuffle=False,
num_workers=args.num_workers,
collate_fn=lerobot_collate_fn if dataset.meta.has_language_columns else None,
pin_memory=device.type == "cuda",
persistent_workers=args.num_workers > 0,
)
model = make_reward_model(config, dataset_meta=metadata).eval()
preprocessor = PolicyProcessorPipeline.from_pretrained(
model_path,
config_filename=f"{POLICY_PREPROCESSOR_DEFAULT_NAME}.json",
overrides={"device_processor": {"device": device.type}},
to_transition=batch_to_transition,
to_output=transition_to_batch,
)
predictions: list[np.ndarray] = []
targets: list[np.ndarray] = []
terminals: list[np.ndarray] = []
episode_indices: list[np.ndarray] = []
frame_indices: list[np.ndarray] = []
actions: list[np.ndarray] = []
states: list[np.ndarray] = []
interventions: list[np.ndarray] = []
nll_sum = 0.0
target_entropy_sum = 0.0
count = 0
with torch.inference_mode():
for batch in tqdm(dataloader, desc="Evaluating reward model"):
if args.max_samples > 0 and count >= args.max_samples:
break
for camera_key in metadata.camera_keys:
if camera_key in batch and batch[camera_key].dtype == torch.uint8:
batch[camera_key] = batch[camera_key].float().div_(255)
raw_target = batch["mc_return"].reshape(-1)
raw_terminal = batch["is_terminal"].reshape(-1).bool()
raw_episode = batch["episode_index"].reshape(-1)
raw_frame = batch["frame_index"].reshape(-1)
if args.max_samples > 0:
keep = min(len(raw_target), args.max_samples - count)
if keep < len(raw_target):
batch = {
key: value[:keep]
if isinstance(value, torch.Tensor)
else value[:keep]
if isinstance(value, list)
else value
for key, value in batch.items()
}
raw_target = raw_target[:keep]
raw_terminal = raw_terminal[:keep]
raw_episode = raw_episode[:keep]
raw_frame = raw_frame[:keep]
if ACTION in batch and isinstance(batch[ACTION], torch.Tensor):
actions.append(batch[ACTION].reshape(len(raw_target), -1).float().cpu().numpy())
if OBS_STATE in batch and isinstance(batch[OBS_STATE], torch.Tensor):
state = batch[OBS_STATE]
if state.ndim >= 3:
state = state[:, -1]
states.append(state.reshape(len(raw_target), -1).float().cpu().numpy())
if "intervention" in batch and isinstance(batch["intervention"], torch.Tensor):
interventions.append(batch["intervention"].reshape(-1).bool().cpu().numpy())
processed = preprocessor(batch)
with torch.autocast(device_type=device.type, dtype=torch.bfloat16, enabled=device.type == "cuda"):
logits, predicted = _predict_logits_and_value(model, processed)
target_on_device = processed["mc_return"].reshape(-1)
terminal_on_device = processed["is_terminal"].reshape(-1).bool()
target_dist = _target_distribution(model, target_on_device, terminal_on_device)
per_sample_nll = -(target_dist * logits.log_softmax(-1)).sum(-1)
entropy = -(target_dist * target_dist.clamp_min(1e-12).log()).sum(-1)
batch_count = len(predicted)
nll_sum += float(per_sample_nll.sum())
target_entropy_sum += float(entropy.sum())
count += batch_count
predictions.append(predicted.float().cpu().numpy())
targets.append(raw_target.float().cpu().numpy())
terminals.append(raw_terminal.cpu().numpy())
episode_indices.append(raw_episode.cpu().numpy())
frame_indices.append(raw_frame.cpu().numpy())
prediction = np.concatenate(predictions)
target = np.concatenate(targets)
terminal = np.concatenate(terminals)
episode_index = np.concatenate(episode_indices).astype(np.int64)
frame_index = np.concatenate(frame_indices).astype(np.int64)
residual = prediction - target
advantage = _compute_advantages(target, prediction, episode_index, args.n_step)
intervention = (
np.concatenate(interventions)
if interventions and sum(map(len, interventions)) == len(prediction)
else np.zeros(len(prediction), dtype=bool)
)
threshold_source = advantage[~intervention]
threshold = float(np.percentile(threshold_source, args.threshold_percentile * 100))
raw_advantage_label = np.where(advantage > threshold, "positive", "negative")
advantage_label = np.where(intervention, "positive", raw_advantage_label)
terminal_success = np.isclose(target[terminal], 0.0, atol=1e-6)
metrics = {
"samples": len(prediction),
"nll": nll_sum / count,
"target_entropy": target_entropy_sum / count,
"excess_nll": (nll_sum - target_entropy_sum) / count,
"mae": float(np.mean(np.abs(residual))),
"rmse": float(np.sqrt(np.mean(np.square(residual)))),
"bias": float(np.mean(residual)),
"prediction_std": float(np.std(prediction)),
"target_std": float(np.std(target)),
"pearson": _correlation(prediction, target),
"spearman": _spearman(prediction, target),
"terminal_success_auc": _binary_auc(prediction[terminal], terminal_success),
"advantage_threshold": threshold,
"positive_fraction_non_intervention": float(
np.mean(raw_advantage_label[~intervention] == "positive")
),
"positive_fraction_after_intervention_override": float(np.mean(advantage_label == "positive")),
}
for name, value in metrics.items():
logger.info("%s: %s", name, f"{value:.6f}" if isinstance(value, float) else value)
output_path = Path(args.output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
output = {
"episode_index": episode_index,
"frame_index": frame_index,
"mc_return": target,
"predicted_value": prediction,
"residual": residual,
"advantage": advantage,
"advantage_label_raw": raw_advantage_label,
"advantage_label": advantage_label,
"is_terminal": terminal,
}
if actions and states:
action = np.concatenate(actions)
state = np.concatenate(states)
common_dim = min(action.shape[1], state.shape[1])
output["command_delta_norm"] = np.linalg.norm(
action[:, :common_dim] - state[:, :common_dim],
axis=1,
)
state_motion = np.zeros(len(state), dtype=np.float32)
same_episode_indices = np.flatnonzero(episode_index[1:] == episode_index[:-1]) + 1
state_motion[same_episode_indices] = np.linalg.norm(
state[same_episode_indices] - state[same_episode_indices - 1],
axis=1,
)
output["state_motion_norm"] = state_motion
if interventions:
output["intervention"] = intervention
pd.DataFrame(output).to_csv(output_path, index=False)
logger.info("Wrote per-frame predictions to %s", output_path)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--reward-model-path", required=True)
parser.add_argument("--dataset-repo-id", required=True)
parser.add_argument("--root")
parser.add_argument("--output-path", required=True)
parser.add_argument("--device", default="cuda")
parser.add_argument("--video-backend", default="torchcodec")
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--num-workers", type=int, default=4)
parser.add_argument("--max-samples", type=int, default=0)
parser.add_argument("--eval-split", type=float, default=0.0)
parser.add_argument("--n-step", type=int)
parser.add_argument("--threshold-percentile", type=float, default=0.6)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
evaluate(args)
if __name__ == "__main__":
main()
+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. # 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: 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) 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: else:
unwrapped_model.push_model_to_hub(cfg, state_dict=model_state_dict, dataset_meta=dataset.meta) unwrapped_model.push_model_to_hub(cfg, state_dict=model_state_dict, dataset_meta=dataset.meta)
preprocessor.push_to_hub(active_cfg.repo_id) preprocessor.push_to_hub(active_cfg.repo_id)
@@ -145,7 +145,7 @@ class SOLeader(Teleoperator):
@check_if_not_connected @check_if_not_connected
def get_action(self) -> dict[str, float]: def get_action(self) -> dict[str, float]:
start = time.perf_counter() 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()} action = {f"{motor}.pos": val for motor, val in action.items()}
dt_ms = (time.perf_counter() - start) * 1e3 dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read action: {dt_ms:.1f}ms") 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: def send_feedback(self, feedback: dict[str, float]) -> None:
goals = {k.removesuffix(".pos"): v for k, v in feedback.items() if k.endswith(".pos")} goals = {k.removesuffix(".pos"): v for k, v in feedback.items() if k.endswith(".pos")}
if goals: if goals:
self.bus.sync_write("Goal_Position", goals) self.bus.sync_write("Goal_Position", goals, num_retry=3)
@check_if_not_connected @check_if_not_connected
def disconnect(self) -> None: def disconnect(self) -> None:
+3
View File
@@ -26,6 +26,9 @@ OBS_IMAGES = OBS_IMAGE + "s"
OBS_LANGUAGE = OBS_STR + ".language" OBS_LANGUAGE = OBS_STR + ".language"
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens" OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask" 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 = OBS_STR + ".subtask"
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens" OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask" OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
+3 -1
View File
@@ -28,9 +28,10 @@ import sys
import tempfile import tempfile
from pathlib import Path 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.executor import Executor
from lerobot.annotations.steerable_pipeline.modules import ( from lerobot.annotations.steerable_pipeline.modules import (
AdvantageModule,
GeneralVqaModule, GeneralVqaModule,
InterjectionsAndSpeechModule, InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule, PlanSubtasksMemoryModule,
@@ -85,6 +86,7 @@ def main() -> int:
plan=PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan), plan=PlanSubtasksMemoryModule(vlm=vlm, config=cfg.plan),
interjections=InterjectionsAndSpeechModule(vlm=vlm, config=cfg.interjections, seed=cfg.seed), interjections=InterjectionsAndSpeechModule(vlm=vlm, config=cfg.interjections, seed=cfg.seed),
vqa=GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed), vqa=GeneralVqaModule(vlm=vlm, config=cfg.vqa, seed=cfg.seed),
advantage=AdvantageModule(config=AdvantageConfig(enabled=False)),
writer=LanguageColumnsWriter(), writer=LanguageColumnsWriter(),
validator=StagingValidator(), validator=StagingValidator(),
) )
+328
View File
@@ -0,0 +1,328 @@
#!/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
assert cfg.threshold_percentile == 0.6
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_run_episode_uses_precomputed_predictions(staging: EpisodeStaging, tmp_path: Path):
predictions_path = tmp_path / "predictions.csv"
predictions_path.write_text(
"episode_index,frame_index,predicted_value\n0,0,-0.4\n0,1,-0.4\n0,2,-0.4\n0,3,-0.4\n"
)
record = _make_record(
num_frames=4,
mc_returns=np.array([-0.7, -0.4, -0.2, -0.1], dtype=np.float32),
)
module = AdvantageModule(
config=AdvantageConfig(
predictions_path=str(predictions_path),
threshold_percentile=0.5,
)
)
module.run_episode(record, staging)
rows = staging.read("advantage")
assert [row["content"] for row in rows] == ["negative", "negative", "positive", "positive"]
assert module._model is None
def test_predictions_csv_requires_every_frame(staging: EpisodeStaging, tmp_path: Path):
predictions_path = tmp_path / "predictions.csv"
predictions_path.write_text("episode_index,frame_index,predicted_value\n0,0,-0.4\n")
module = AdvantageModule(config=AdvantageConfig(predictions_path=str(predictions_path)))
with pytest.raises(KeyError, match="missing 1 frame"):
module.run_episode(_make_record(num_frames=2), staging)
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 import pyarrow.parquet as pq # noqa: E402
from lerobot.annotations.steerable_pipeline.config import ( # noqa: E402 from lerobot.annotations.steerable_pipeline.config import ( # noqa: E402
AdvantageConfig,
AnnotationPipelineConfig, AnnotationPipelineConfig,
InterjectionsConfig, InterjectionsConfig,
PlanConfig, 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.executor import Executor # noqa: E402
from lerobot.annotations.steerable_pipeline.modules import ( # noqa: E402 from lerobot.annotations.steerable_pipeline.modules import ( # noqa: E402
AdvantageModule,
GeneralVqaModule, GeneralVqaModule,
InterjectionsAndSpeechModule, InterjectionsAndSpeechModule,
PlanSubtasksMemoryModule, PlanSubtasksMemoryModule,
@@ -132,6 +134,7 @@ def _build_executor() -> Executor:
plan=PlanSubtasksMemoryModule(vlm=vlm, config=config.plan), plan=PlanSubtasksMemoryModule(vlm=vlm, config=config.plan),
interjections=InterjectionsAndSpeechModule(vlm=vlm, config=config.interjections, seed=config.seed), interjections=InterjectionsAndSpeechModule(vlm=vlm, config=config.interjections, seed=config.seed),
vqa=GeneralVqaModule(vlm=vlm, config=config.vqa, seed=config.seed), vqa=GeneralVqaModule(vlm=vlm, config=config.vqa, seed=config.seed),
advantage=AdvantageModule(config=AdvantageConfig(enabled=False)),
writer=LanguageColumnsWriter(), writer=LanguageColumnsWriter(),
validator=StagingValidator(), 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,747 @@
# 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 = 1
NUM_BINS = 201
IMAGE_SIZE = 448
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": (IMAGE_SIZE, IMAGE_SIZE),
}
defaults.update(overrides)
config = DistributionalVFConfig(**defaults)
config.input_features = {
IMAGE_KEY: PolicyFeature(type=FeatureType.VISUAL, shape=(3, IMAGE_SIZE, IMAGE_SIZE)),
IMAGE_KEY_WRIST_LEFT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, IMAGE_SIZE, IMAGE_SIZE)),
IMAGE_KEY_WRIST_RIGHT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, IMAGE_SIZE, IMAGE_SIZE)),
}
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, IMAGE_SIZE, IMAGE_SIZE, 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, IMAGE_SIZE, IMAGE_SIZE, 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, IMAGE_SIZE, IMAGE_SIZE, 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
def test_config_defaults_match_pi06_gemma3_layout():
config = DistributionalVFConfig()
assert config.image_resolution == (448, 448)
assert config.num_image_tokens == 256
assert config.target_method == "dirac_delta"
assert config.hl_gauss_sigma_ratio == 0.75
# ------------------------------------------------------------------
# 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_hl_gauss_clamps_values_outside_support():
model = _make_model()
outside = model.hl_gauss_target(torch.tensor([-1.5, 0.5]))
boundaries = model.hl_gauss_target(torch.tensor([-1.0, 0.0]))
torch.testing.assert_close(outside, boundaries)
@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_terminal_column_vector_preserves_batch_shape():
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=True
)
assert dist.shape == (2, NUM_BINS)
assert (dist[0] > 0).sum() > 2
assert (dist[1] > 0).sum() == 1
@skip_if_package_missing("transformers")
def test_terminal_count_must_match_batch_size():
model = _make_model()
with pytest.raises(ValueError, match="Expected 2 terminal flags, got 1"):
model.compute_target_distribution(
torch.tensor([-0.5, -0.3]),
torch.tensor([True]),
method="hl_gauss",
use_one_hot_terminal=True,
)
@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, "multi_modal_projector")
assert hasattr(model, "value_head")
assert hasattr(model, "value_query")
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_value_query_is_nn_embedding():
"""Value query is nn.Embedding (FSDP-safe) with correct shape."""
model = _make_model()
from torch import nn
assert isinstance(model.value_query, nn.Embedding)
assert model.value_query.num_embeddings == 1
assert model.value_query.embedding_dim == model.gemma3_hidden
@skip_if_package_missing("transformers")
def test_multimodal_projector_dimensions_and_pooling():
"""Gemma3 connector pools 448px patches to 256 tokens and projects to LM width."""
model = _make_model()
siglip_hidden = model.vision_encoder.config.hidden_size
projector = model.multi_modal_projector
assert projector.mm_input_projection_weight.shape == (siglip_hidden, model.gemma3_hidden)
assert projector.patches_per_image == 32
assert projector.tokens_per_side == 16
assert projector.kernel_size == 2
# ------------------------------------------------------------------
# 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=1)
with torch.no_grad():
values = model.compute_reward(batch)
assert values.shape == (1,)
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=1)
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_value_query():
"""Backprop produces non-zero gradients on the learned value query."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.value_query.weight.grad is not None
assert not torch.all(model.value_query.weight.grad == 0)
@skip_if_package_missing("transformers")
def test_gradient_flows_through_multimodal_projector():
"""Backprop produces non-zero gradients on the Gemma3 multimodal projection."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
projector_weight = model.multi_modal_projector.mm_input_projection_weight
assert projector_weight.grad is not None
assert not torch.all(projector_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_value_query_grad():
"""With stop_gradient_to_vlm, the value query 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.value_query.weight.grad is not None
assert not torch.all(model.value_query.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, IMAGE_SIZE, IMAGE_SIZE)
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.full((2, 3, 320, 240), 0.5), # 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
# Content value 0.5 must map to 0.0. This also verifies normalization
# happens before resize padding introduces -1 values.
assert torch.allclose(obs[IMAGE_KEY][:, :, 112, 112], torch.zeros(2, 3), atol=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}")
# ------------------------------------------------------------------
# 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)
@@ -0,0 +1,156 @@
import json
import sys
from types import ModuleType, SimpleNamespace
import torch
from PIL import Image
from torch import nn
from lerobot.configs import FeatureType, PolicyFeature
from lerobot.rewards.factory import get_reward_model_class, make_reward_model_config
from lerobot.rewards.nanovlm_value_function.configuration_nanovlm_value_function import (
NanoVLMVFConfig,
)
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
NANOVLM_ATTENTION_MASK,
NANOVLM_IMAGES,
NANOVLM_INPUT_IDS,
NanoVLMNativeProcessorStep,
)
from lerobot.types import TransitionKey
CAMERA = "observation.images.top"
def test_config_and_factory_registration():
config = make_reward_model_config("nanovlm_value_function")
assert isinstance(config, NanoVLMVFConfig)
assert config.tokenizer_max_length == 8192
assert get_reward_model_class("nanovlm_value_function").__name__ == "NanoVLMVFRewardModel"
def test_nanovlm_model_forward(monkeypatch):
from lerobot.rewards.nanovlm_value_function.modeling_nanovlm_value_function import (
NanoVLMVFRewardModel,
)
class FakeVision(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.ones(1))
def forward(self, image):
return torch.ones(image.shape[0], 4, 6)
class FakeProjector(nn.Module):
def __init__(self):
super().__init__()
self.proj = nn.Linear(6, 8)
def forward(self, features):
return self.proj(features)
class FakeDecoder(nn.Module):
def __init__(self):
super().__init__()
self.token_embedding = nn.Embedding(100, 8)
def forward(self, inputs, attention_mask=None):
return inputs, None
class FakeNano(nn.Module):
def __init__(self):
super().__init__()
self.cfg = SimpleNamespace(lm_hidden_dim=8)
self.vision_encoder = FakeVision()
self.MP = FakeProjector()
self.decoder = FakeDecoder()
self.tokenizer = SimpleNamespace(image_token_id=99)
def _process_images(self, images, device):
return torch.cat([image for sample in images for image in sample]).to(device)
def _replace_img_tokens_with_embd(self, input_ids, token_embd, image_embd):
token_embd = token_embd.clone()
token_embd[input_ids == self.tokenizer.image_token_id] = image_embd.flatten(0, 1)
return token_embd
@classmethod
def from_pretrained(cls, path):
return cls()
fake_module = ModuleType("models.vision_language_model")
fake_module.VisionLanguageModel = FakeNano
monkeypatch.setitem(sys.modules, "models.vision_language_model", fake_module)
config = NanoVLMVFConfig(
device="cpu",
nanovlm_code_path="third_party/nanoVLM",
)
config.input_features = {CAMERA: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 16, 16))}
model = NanoVLMVFRewardModel(config)
batch = {
NANOVLM_IMAGES: [[torch.rand(1, 3, 16, 16)]],
NANOVLM_INPUT_IDS: torch.tensor([[99, 99, 99, 99, 1]]),
NANOVLM_ATTENTION_MASK: torch.ones(1, 5, dtype=torch.bool),
"mc_return": torch.tensor([-0.5]),
"is_terminal": torch.tensor([False]),
}
loss, metrics = model(batch)
assert torch.isfinite(loss)
assert -1.0 <= metrics["predicted_value_mean"] <= 0.0
def test_native_processor_uses_checkpoint_layout_and_left_padding(monkeypatch, tmp_path):
config = {
"lm_tokenizer": "fake",
"vlm_extra_tokens": {},
"lm_chat_template": "fake",
"lm_max_length": 8192,
"max_img_size": 2048,
"vit_img_size": 512,
"resize_to_max_side_len": True,
"mp_image_token_length": 4,
}
(tmp_path / "config.json").write_text(json.dumps(config))
class FakeTokenizer:
pad_token_id = 0
image_token_id = 99
def apply_chat_template(self, messages, tokenize, add_generation_prompt):
assert not tokenize and add_generation_prompt
return messages[0]["content"]
def __call__(self, prompt, truncation, add_special_tokens):
assert not truncation and not add_special_tokens
suffix = [1, 2] if "long" in prompt else [1]
return {"input_ids": [99] * 4 + suffix, "attention_mask": [1] * (4 + len(suffix))}
def fake_image_processor(image):
assert isinstance(image, Image.Image) and image.mode == "RGB"
return torch.rand(1, 3, 512, 512), (1, 1)
processors = ModuleType("data.processors")
processors.get_tokenizer = lambda *args: FakeTokenizer()
processors.get_image_processor = lambda *args: fake_image_processor
processors.get_image_string = lambda *args: "<image>"
monkeypatch.setitem(sys.modules, "data.processors", processors)
step = NanoVLMNativeProcessorStep(
pretrained_path=str(tmp_path),
code_path="third_party/nanoVLM",
image_keys=(CAMERA,),
max_length=8192,
)
transition = {
TransitionKey.OBSERVATION: {CAMERA: torch.rand(2, 3, 16, 16)},
TransitionKey.COMPLEMENTARY_DATA: {"task": ["short", "long"]},
}
output = step(transition)[TransitionKey.OBSERVATION]
assert len(output[NANOVLM_IMAGES]) == 2
assert output[NANOVLM_INPUT_IDS].shape == (2, 6)
assert output[NANOVLM_INPUT_IDS][0, 0] == 0
assert not output[NANOVLM_ATTENTION_MASK][0, 0]
assert output[NANOVLM_ATTENTION_MASK][1].all()
@@ -0,0 +1,46 @@
from types import SimpleNamespace
from torch import nn
from lerobot.configs import FeatureType
from lerobot.rewards.factory import make_reward_model
from lerobot.rewards.temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
TemporalSiglipVFConfig,
)
def test_reward_factory_populates_input_features_from_dataset_meta(monkeypatch):
from lerobot.rewards import factory
class FakeReward(nn.Module):
def __init__(self, config, **kwargs):
super().__init__()
self.config = config
monkeypatch.setattr(factory, "get_reward_model_class", lambda name: FakeReward)
metadata = SimpleNamespace(
features={
"observation.images.top": {
"dtype": "video",
"shape": [480, 640, 3],
"names": ["height", "width", "channel"],
},
"observation.state": {
"dtype": "float32",
"shape": [14],
"names": [f"joint_{index}" for index in range(14)],
},
"action": {
"dtype": "float32",
"shape": [14],
"names": [f"joint_{index}" for index in range(14)],
},
}
)
config = TemporalSiglipVFConfig(device="cpu")
model = make_reward_model(config, dataset_meta=metadata)
assert model.config.input_features["observation.images.top"].type is FeatureType.VISUAL
assert model.config.input_features["observation.images.top"].shape == (3, 480, 640)
assert model.config.input_features["observation.state"].type is FeatureType.STATE
assert "action" not in model.config.input_features
@@ -0,0 +1,141 @@
from types import SimpleNamespace
import torch
from torch import nn
from lerobot.configs import FeatureType, PolicyFeature
from lerobot.rewards.factory import get_reward_model_class, make_reward_model_config
from lerobot.rewards.temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
TemporalSiglipVFConfig,
)
from lerobot.rewards.temporal_siglip_value_function.processor_temporal_siglip_value_function import (
TemporalSiglipImageProcessorStep,
)
from lerobot.types import TransitionKey
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE
CAMERAS = ("observation.images.top", "observation.images.left", "observation.images.right")
def _config(**kwargs):
config = TemporalSiglipVFConfig(
device="cpu",
hidden_size=8,
num_layers=1,
num_heads=2,
history_steps=2,
state_dim=4,
**kwargs,
)
config.input_features = {
**{key: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 16, 16)) for key in CAMERAS},
OBS_STATE: PolicyFeature(type=FeatureType.STATE, shape=(4,)),
}
return config
def test_config_and_factory_registration():
config = make_reward_model_config("temporal_siglip_value_function")
assert isinstance(config, TemporalSiglipVFConfig)
assert get_reward_model_class("temporal_siglip_value_function").__name__ == "TemporalSiglipVFRewardModel"
def test_history_offsets_are_past_only():
config = TemporalSiglipVFConfig(history_steps=4, frame_gap=10)
assert config.observation_delta_indices == [-30, -20, -10, 0]
def test_temporal_image_processor():
step = TemporalSiglipImageProcessorStep(
image_resolution=(32, 32),
image_keys=(CAMERAS[0],),
history_steps=2,
)
transition = {
TransitionKey.OBSERVATION: {
CAMERAS[0]: torch.full((1, 2, 3, 20, 16), 128, dtype=torch.uint8),
CAMERAS[0] + "_is_pad": torch.tensor([[True, False]]),
}
}
observation = step(transition)[TransitionKey.OBSERVATION]
assert observation[CAMERAS[0]].shape == (1, 2, 3, 32, 32)
assert observation[CAMERAS[0] + ".mask"].shape == (1, 2)
assert observation[CAMERAS[0] + ".mask"].tolist() == [[False, True]]
assert -1.0 <= observation[CAMERAS[0]].min() <= observation[CAMERAS[0]].max() <= 1.0
def test_siglip_tokenizer_builds_missing_attention_mask(monkeypatch):
from lerobot.rewards.temporal_siglip_value_function import (
processor_temporal_siglip_value_function as processing,
)
class FakeTokenizer:
pad_token_id = 0
def __call__(self, *args, **kwargs):
return {"input_ids": torch.tensor([[5, 1, 0, 0]])}
monkeypatch.setattr(
processing.AutoTokenizer,
"from_pretrained",
lambda *args, **kwargs: FakeTokenizer(),
)
step = processing.TemporalSiglipTokenizerStep("fake", max_length=4)
transition = {
TransitionKey.OBSERVATION: {},
TransitionKey.COMPLEMENTARY_DATA: {"task": ["Task: stack blocks."]},
}
observation = step(transition)[TransitionKey.OBSERVATION]
assert observation[OBS_LANGUAGE_ATTENTION_MASK].tolist() == [[True, True, False, False]]
def test_temporal_model_forward(monkeypatch):
from lerobot.rewards.temporal_siglip_value_function import (
modeling_temporal_siglip_value_function as modeling,
)
class FakeEncoder(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.ones(1))
def forward(self, pixel_values=None, input_ids=None, **kwargs):
batch = pixel_values.shape[0] if pixel_values is not None else input_ids.shape[0]
return SimpleNamespace(pooler_output=torch.ones(batch, 8))
class FakeSiglip(nn.Module):
def __init__(self):
super().__init__()
self.vision_model = FakeEncoder()
self.text_model = FakeEncoder()
self.config = SimpleNamespace(
vision_config=SimpleNamespace(hidden_size=8),
text_config=SimpleNamespace(hidden_size=8),
)
monkeypatch.setattr(modeling.AutoModel, "from_pretrained", lambda *args, **kwargs: FakeSiglip())
model = modeling.TemporalSiglipVFRewardModel(_config())
batch = {
**{key: torch.rand(1, 2, 3, 16, 16) for key in CAMERAS},
**{key + ".mask": torch.tensor([[False, True]]) for key in CAMERAS},
OBS_STATE: torch.rand(1, 2, 4),
OBS_LANGUAGE_TOKENS: torch.ones(1, 4, dtype=torch.long),
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(1, 4, dtype=torch.bool),
"mc_return": torch.tensor([-0.5]),
"is_terminal": torch.tensor([False]),
}
loss, metrics = model(batch)
assert torch.isfinite(loss)
assert -1.0 <= metrics["predicted_value_mean"] <= 0.0
model.eval()
with torch.no_grad():
eval_loss, eval_metrics = model(batch)
assert torch.isfinite(eval_loss)
assert -1.0 <= eval_metrics["predicted_value_mean"] <= 0.0
targets = model.compute_target_distribution(
torch.tensor([[-0.5], [-0.3]]),
torch.tensor([[False], [True]]),
)
assert targets.shape == (2, model.config.num_value_bins)
+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"
@@ -0,0 +1,102 @@
import shutil
import subprocess
import cv2
import numpy as np
import pandas as pd
import pytest
from lerobot.scripts.lerobot_create_advantage_video import (
_contiguous_segments,
_create_decode_proxy,
_draw_dashboard,
_draw_label_timeline,
)
def _episode_frame() -> pd.DataFrame:
return pd.DataFrame(
{
"predicted_value": [-0.8, -0.6, -0.4, -0.2],
"predicted_value_raw": [-0.81, -0.55, -0.45, -0.18],
"mc_return": [-0.9, -0.6, -0.3, 0.0],
"advantage": [-0.1, 0.1, -0.05, 0.2],
"advantage_label": ["negative", "positive", "negative", "positive"],
"intervention": [False, False, True, False],
}
)
def test_contiguous_advantage_segments():
labels = np.array(["negative", "negative", "positive", "positive", "negative"])
assert _contiguous_segments(labels) == [
(0, 2, "negative"),
(2, 4, "positive"),
(4, 5, "negative"),
]
def test_advantage_dashboard_shape():
dashboard = _draw_dashboard(
width=640,
height=140,
episode=_episode_frame(),
current_index=2,
fps=30,
task="stack the yellow cube on the red cube",
)
assert dashboard.shape == (140, 640, 3)
assert dashboard.dtype == np.uint8
assert dashboard.var() > 0
def test_label_timeline_modifies_only_lower_frame_region():
frame = np.zeros((480, 640, 3), dtype=np.uint8)
labels = np.array(["negative", "negative", "positive", "positive"])
_draw_label_timeline(frame, labels, current_index=2)
assert frame.sum() > 0
assert frame[:350].sum() == 0
@pytest.mark.skipif(shutil.which("ffmpeg") is None, reason="ffmpeg is required")
def test_ffmpeg_decode_proxy_handles_av1(tmp_path):
source = tmp_path / "source.mkv"
proxy = tmp_path / "proxy.mp4"
encode = subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"testsrc=size=64x64:rate=10:duration=1",
"-c:v",
"libaom-av1",
"-cpu-used",
"8",
"-crf",
"40",
str(source),
],
capture_output=True,
)
if encode.returncode != 0:
pytest.skip("ffmpeg does not provide the libaom-av1 encoder")
_create_decode_proxy(
source,
proxy,
start_timestamp=0,
end_timestamp=1,
fps=10,
expected_frames=10,
)
capture = cv2.VideoCapture(str(proxy))
frame_count = int(capture.get(cv2.CAP_PROP_FRAME_COUNT))
capture.release()
assert frame_count == 10
+50
View File
@@ -0,0 +1,50 @@
from types import SimpleNamespace
import numpy as np
import pytest
from lerobot.scripts.lerobot_eval_reward_model import (
_binary_auc,
_compute_advantages,
_correlation,
_held_out_episodes,
_spearman,
)
def test_reward_evaluation_correlations():
target = np.array([-1.0, -0.5, 0.0])
prediction = np.array([-0.9, -0.5, -0.1])
assert _correlation(prediction, target) == pytest.approx(1.0)
assert _spearman(prediction, target) == pytest.approx(1.0)
def test_reward_evaluation_auc():
labels = np.array([False, False, True, True])
scores = np.array([-0.9, -0.7, -0.2, -0.1])
assert _binary_auc(scores, labels) == pytest.approx(1.0)
def test_reward_evaluation_reproduces_training_episode_split():
metadata = SimpleNamespace(
total_episodes=5,
episodes={"tasks": [["a"], ["a"], ["a"], ["b"], ["b"]]},
)
assert _held_out_episodes(metadata, 0.34) == [1, 2, 4]
assert _held_out_episodes(metadata, 0.0) is None
def test_reward_evaluation_n_step_advantages():
target = np.array([-0.9, -0.6, -0.3, -0.8, -0.4], dtype=np.float32)
prediction = np.array([-0.8, -0.5, -0.2, -0.7, -0.3], dtype=np.float32)
episodes = np.array([0, 0, 0, 1, 1])
advantage = _compute_advantages(target, prediction, episodes, n_step=2)
expected_first = target[0] - target[2] + prediction[2] - prediction[0]
assert advantage[0] == pytest.approx(expected_first)
assert advantage[1] == pytest.approx(target[1] - prediction[1])
assert advantage[3] == pytest.approx(target[3] - prediction[3])
+97
View File
@@ -338,6 +338,103 @@ def test_dagger_events_reset():
assert not events.upload_requested.is_set() 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 # Context dataclass
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
Generated
+1243 -1483
View File
File diff suppressed because it is too large Load Diff