From 445ffae5ced6043415ce9f39a8d533c4f3d308f5 Mon Sep 17 00:00:00 2001 From: Khalil Meftah Date: Thu, 12 Mar 2026 15:23:42 +0100 Subject: [PATCH] refactor(rewards): update imports and delete old reward model locations --- docs/source/sarm.mdx | 10 +- examples/dataset/slurm_compute_rabc.py | 2 +- examples/tutorial/rl/hilserl_example.py | 2 +- .../tutorial/rl/reward_classifier_example.py | 7 +- .../reward_model/configuration_classifier.py | 77 -- .../sac/reward_model/modeling_classifier.py | 308 ------- .../sac/reward_model/processor_classifier.py | 82 -- .../policies/sarm/compute_rabc_weights.py | 870 ------------------ .../policies/sarm/configuration_sarm.py | 249 ----- src/lerobot/policies/sarm/modeling_sarm.py | 794 ---------------- src/lerobot/policies/sarm/processor_sarm.py | 516 ----------- src/lerobot/policies/sarm/sarm_utils.py | 295 ------ src/lerobot/processor/hil_processor.py | 2 +- src/lerobot/scripts/lerobot_train.py | 26 +- 14 files changed, 31 insertions(+), 3209 deletions(-) delete mode 100644 src/lerobot/policies/sac/reward_model/configuration_classifier.py delete mode 100644 src/lerobot/policies/sac/reward_model/modeling_classifier.py delete mode 100644 src/lerobot/policies/sac/reward_model/processor_classifier.py delete mode 100644 src/lerobot/policies/sarm/compute_rabc_weights.py delete mode 100644 src/lerobot/policies/sarm/configuration_sarm.py delete mode 100644 src/lerobot/policies/sarm/modeling_sarm.py delete mode 100644 src/lerobot/policies/sarm/processor_sarm.py delete mode 100644 src/lerobot/policies/sarm/sarm_utils.py diff --git a/docs/source/sarm.mdx b/docs/source/sarm.mdx index cd488fe1f..20f1709e1 100644 --- a/docs/source/sarm.mdx +++ b/docs/source/sarm.mdx @@ -46,7 +46,7 @@ This ensures identical task states map to consistent progress values, even acros ## Inputs and Targets (What the new code expects) -SARM is trained through its processor (`src/lerobot/policies/sarm/processor_sarm.py`), which: +SARM is trained through its processor (`src/lerobot/rewards/sarm/processor_sarm.py`), which: - **Encodes** images and task text with CLIP (ViT-B/32) into `video_features` and `text_features` - **Pads/truncates** robot state into `state_features` (up to `max_state_dim`) @@ -347,7 +347,7 @@ Use `compute_rabc_weights.py` with `--visualize-only` to visualize model predict ```bash -python src/lerobot/policies/sarm/compute_rabc_weights.py \ +python -m lerobot.rewards.sarm.compute_rabc_weights \ --dataset-repo-id your-username/your-dataset \ --reward-model-path your-username/sarm-model \ --visualize-only \ @@ -360,7 +360,7 @@ python src/lerobot/policies/sarm/compute_rabc_weights.py \ ```bash -python src/lerobot/policies/sarm/compute_rabc_weights.py \ +python -m lerobot.rewards.sarm.compute_rabc_weights \ --dataset-repo-id your-username/your-dataset \ --reward-model-path your-username/sarm-model \ --visualize-only \ @@ -373,7 +373,7 @@ python src/lerobot/policies/sarm/compute_rabc_weights.py \ ```bash -python src/lerobot/policies/sarm/compute_rabc_weights.py \ +python -m lerobot.rewards.sarm.compute_rabc_weights \ --dataset-repo-id your-username/your-dataset \ --reward-model-path your-username/sarm-model \ --visualize-only \ @@ -429,7 +429,7 @@ The weighting follows **Equations 8-9** from the paper: First, run the SARM model on all frames in your dataset to compute progress values: ```bash -python src/lerobot/policies/sarm/compute_rabc_weights.py \ +python -m lerobot.rewards.sarm.compute_rabc_weights \ --dataset-repo-id your-username/your-dataset \ --reward-model-path your-username/sarm-model \ --head-mode sparse \ diff --git a/examples/dataset/slurm_compute_rabc.py b/examples/dataset/slurm_compute_rabc.py index 2ddf84d07..5367e929a 100644 --- a/examples/dataset/slurm_compute_rabc.py +++ b/examples/dataset/slurm_compute_rabc.py @@ -69,7 +69,7 @@ class ComputeProgressShards(PipelineStep): import torch from tqdm import tqdm - from lerobot.policies.sarm.compute_rabc_weights import ( + from lerobot.rewards.sarm.compute_rabc_weights import ( generate_all_frame_indices, interpolate_progress, load_sarm_resources, diff --git a/examples/tutorial/rl/hilserl_example.py b/examples/tutorial/rl/hilserl_example.py index 980ac7985..ad914b4c6 100644 --- a/examples/tutorial/rl/hilserl_example.py +++ b/examples/tutorial/rl/hilserl_example.py @@ -11,7 +11,7 @@ from lerobot.datasets.utils import hw_to_dataset_features from lerobot.envs.configs import HILSerlProcessorConfig, HILSerlRobotEnvConfig from lerobot.policies.sac.configuration_sac import SACConfig from lerobot.policies.sac.modeling_sac import SACPolicy -from lerobot.policies.sac.reward_model.modeling_classifier import Classifier +from lerobot.rewards.classifier.modeling_classifier import Classifier from lerobot.rl.buffer import ReplayBuffer from lerobot.rl.gym_manipulator import make_robot_env from lerobot.robots.so_follower import SO100FollowerConfig diff --git a/examples/tutorial/rl/reward_classifier_example.py b/examples/tutorial/rl/reward_classifier_example.py index 4af6b899c..691b4eb19 100644 --- a/examples/tutorial/rl/reward_classifier_example.py +++ b/examples/tutorial/rl/reward_classifier_example.py @@ -1,8 +1,9 @@ import torch from lerobot.datasets.lerobot_dataset import LeRobotDataset -from lerobot.policies.factory import make_policy, make_pre_post_processors -from lerobot.policies.sac.reward_model.configuration_classifier import RewardClassifierConfig +from lerobot.policies.factory import make_policy +from lerobot.rewards.classifier.configuration_classifier import RewardClassifierConfig +from lerobot.rewards.factory import make_reward_pre_post_processors def main(): @@ -26,7 +27,7 @@ def main(): # Make policy, preprocessor, and optimizer policy = make_policy(config, ds_meta=dataset.meta) optimizer = config.get_optimizer_preset().build(policy.parameters()) - preprocessor, _ = make_pre_post_processors(policy_cfg=config, dataset_stats=dataset.meta.stats) + preprocessor, _ = make_reward_pre_post_processors(config, dataset_stats=dataset.meta.stats) classifier_id = "/reward_classifier_hil_serl_example" diff --git a/src/lerobot/policies/sac/reward_model/configuration_classifier.py b/src/lerobot/policies/sac/reward_model/configuration_classifier.py deleted file mode 100644 index 879e3c1af..000000000 --- a/src/lerobot/policies/sac/reward_model/configuration_classifier.py +++ /dev/null @@ -1,77 +0,0 @@ -# !/usr/bin/env python - -# Copyright 2025 The HuggingFace Inc. team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -from dataclasses import dataclass, field - -from lerobot.configs.policies import PreTrainedConfig -from lerobot.configs.types import NormalizationMode -from lerobot.optim.optimizers import AdamWConfig, OptimizerConfig -from lerobot.optim.schedulers import LRSchedulerConfig -from lerobot.utils.constants import OBS_IMAGE - - -@PreTrainedConfig.register_subclass(name="reward_classifier") -@dataclass -class RewardClassifierConfig(PreTrainedConfig): - """Configuration for the Reward Classifier model.""" - - name: str = "reward_classifier" - num_classes: int = 2 - hidden_dim: int = 256 - latent_dim: int = 256 - image_embedding_pooling_dim: int = 8 - dropout_rate: float = 0.1 - model_name: str = "helper2424/resnet10" # TODO: This needs to be updated. The model on the Hub doesn't call self.post_init() in its __init__, which is required by transformers v5 to set all_tied_weights_keys. The from_pretrained call fails when it tries to access this attribute during _finalize_model_loading. - device: str = "cpu" - model_type: str = "cnn" # "transformer" or "cnn" - num_cameras: int = 2 - learning_rate: float = 1e-4 - weight_decay: float = 0.01 - grad_clip_norm: float = 1.0 - normalization_mapping: dict[str, NormalizationMode] = field( - default_factory=lambda: { - "VISUAL": NormalizationMode.MEAN_STD, - } - ) - - @property - def observation_delta_indices(self) -> list | None: - return None - - @property - def action_delta_indices(self) -> list | None: - return None - - @property - def reward_delta_indices(self) -> list | None: - return None - - def get_optimizer_preset(self) -> OptimizerConfig: - return AdamWConfig( - lr=self.learning_rate, - weight_decay=self.weight_decay, - grad_clip_norm=self.grad_clip_norm, - ) - - def get_scheduler_preset(self) -> LRSchedulerConfig | None: - return None - - def validate_features(self) -> None: - """Validate feature configurations.""" - has_image = any(key.startswith(OBS_IMAGE) for key in self.input_features) - if not has_image: - raise ValueError( - "You must provide an image observation (key starting with 'observation.image') in the input features" - ) diff --git a/src/lerobot/policies/sac/reward_model/modeling_classifier.py b/src/lerobot/policies/sac/reward_model/modeling_classifier.py deleted file mode 100644 index dba6a174b..000000000 --- a/src/lerobot/policies/sac/reward_model/modeling_classifier.py +++ /dev/null @@ -1,308 +0,0 @@ -# !/usr/bin/env python - -# Copyright 2025 The HuggingFace Inc. team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import logging - -import torch -from torch import Tensor, nn - -from lerobot.policies.pretrained import PreTrainedPolicy -from lerobot.policies.sac.reward_model.configuration_classifier import RewardClassifierConfig -from lerobot.utils.constants import OBS_IMAGE, REWARD - - -class ClassifierOutput: - """Wrapper for classifier outputs with additional metadata.""" - - def __init__( - self, - logits: Tensor, - probabilities: Tensor | None = None, - hidden_states: Tensor | None = None, - ): - self.logits = logits - self.probabilities = probabilities - self.hidden_states = hidden_states - - def __repr__(self): - return ( - f"ClassifierOutput(logits={self.logits}, " - f"probabilities={self.probabilities}, " - f"hidden_states={self.hidden_states})" - ) - - -class SpatialLearnedEmbeddings(nn.Module): - def __init__(self, height, width, channel, num_features=8): - """ - PyTorch implementation of learned spatial embeddings - - Args: - height: Spatial height of input features - width: Spatial width of input features - channel: Number of input channels - num_features: Number of output embedding dimensions - """ - super().__init__() - self.height = height - self.width = width - self.channel = channel - self.num_features = num_features - - self.kernel = nn.Parameter(torch.empty(channel, height, width, num_features)) - - nn.init.kaiming_normal_(self.kernel, mode="fan_in", nonlinearity="linear") - - def forward(self, features): - """ - Forward pass for spatial embedding - - Args: - features: Input tensor of shape [B, H, W, C] or [H, W, C] if no batch - Returns: - Output tensor of shape [B, C*F] or [C*F] if no batch - """ - - features = features.last_hidden_state - - original_shape = features.shape - if features.dim() == 3: - features = features.unsqueeze(0) # Add batch dim - - features_expanded = features.unsqueeze(-1) # [B, H, W, C, 1] - kernel_expanded = self.kernel.unsqueeze(0) # [1, H, W, C, F] - - # Element-wise multiplication and spatial reduction - output = (features_expanded * kernel_expanded).sum(dim=(2, 3)) # Sum H,W - - # Reshape to combine channel and feature dimensions - output = output.view(output.size(0), -1) # [B, C*F] - - # Remove batch dim - if len(original_shape) == 3: - output = output.squeeze(0) - - return output - - -class Classifier(PreTrainedPolicy): - """Image classifier built on top of a pre-trained encoder.""" - - name = "reward_classifier" - config_class = RewardClassifierConfig - - def __init__( - self, - config: RewardClassifierConfig, - ): - from transformers import AutoModel - - super().__init__(config) - self.config = config - - # Set up encoder - encoder = AutoModel.from_pretrained(self.config.model_name, trust_remote_code=True) - # Extract vision model if we're given a multimodal model - if hasattr(encoder, "vision_model"): - logging.info("Multimodal model detected - using vision encoder only") - self.encoder = encoder.vision_model - self.vision_config = encoder.config.vision_config - else: - self.encoder = encoder - self.vision_config = getattr(encoder, "config", None) - - # Model type from config - self.is_cnn = self.config.model_type == "cnn" - - # For CNNs, initialize backbone - if self.is_cnn: - self._setup_cnn_backbone() - - self._freeze_encoder() - - # Extract image keys from input_features - self.image_keys = [ - key.replace(".", "_") for key in config.input_features if key.startswith(OBS_IMAGE) - ] - - if self.is_cnn: - self.encoders = nn.ModuleDict() - for image_key in self.image_keys: - encoder = self._create_single_encoder() - self.encoders[image_key] = encoder - - self._build_classifier_head() - - def _setup_cnn_backbone(self): - """Set up CNN encoder""" - if hasattr(self.encoder, "fc"): - self.feature_dim = self.encoder.fc.in_features - self.encoder = nn.Sequential(*list(self.encoder.children())[:-1]) - elif hasattr(self.encoder.config, "hidden_sizes"): - self.feature_dim = self.encoder.config.hidden_sizes[-1] # Last channel dimension - else: - raise ValueError("Unsupported CNN architecture") - - def _freeze_encoder(self) -> None: - """Freeze the encoder parameters.""" - for param in self.encoder.parameters(): - param.requires_grad = False - - def _create_single_encoder(self): - encoder = nn.Sequential( - self.encoder, - SpatialLearnedEmbeddings( - height=4, - width=4, - channel=self.feature_dim, - num_features=self.config.image_embedding_pooling_dim, - ), - nn.Dropout(self.config.dropout_rate), - nn.Linear(self.feature_dim * self.config.image_embedding_pooling_dim, self.config.latent_dim), - nn.LayerNorm(self.config.latent_dim), - nn.Tanh(), - ) - - return encoder - - def _build_classifier_head(self) -> None: - """Initialize the classifier head architecture.""" - # Get input dimension based on model type - if self.is_cnn: - input_dim = self.config.latent_dim - else: # Transformer models - if hasattr(self.encoder.config, "hidden_size"): - input_dim = self.encoder.config.hidden_size - else: - raise ValueError("Unsupported transformer architecture since hidden_size is not found") - - self.classifier_head = nn.Sequential( - nn.Linear(input_dim * self.config.num_cameras, self.config.hidden_dim), - nn.Dropout(self.config.dropout_rate), - nn.LayerNorm(self.config.hidden_dim), - nn.ReLU(), - nn.Linear( - self.config.hidden_dim, - 1 if self.config.num_classes == 2 else self.config.num_classes, - ), - ) - - def _get_encoder_output(self, x: torch.Tensor, image_key: str) -> torch.Tensor: - """Extract the appropriate output from the encoder.""" - with torch.no_grad(): - if self.is_cnn: - # The HF ResNet applies pooling internally - outputs = self.encoders[image_key](x) - return outputs - else: # Transformer models - outputs = self.encoder(x) - return outputs.last_hidden_state[:, 0, :] - - def extract_images_and_labels(self, batch: dict[str, Tensor]) -> tuple[list, Tensor]: - """Extract image tensors and label tensors from batch.""" - # Check for both OBS_IMAGE and OBS_IMAGES prefixes - images = [batch[key] for key in self.config.input_features if key.startswith(OBS_IMAGE)] - labels = batch[REWARD] - - return images, labels - - def predict(self, xs: list) -> ClassifierOutput: - """Forward pass of the classifier for inference.""" - encoder_outputs = torch.hstack( - [self._get_encoder_output(x, img_key) for x, img_key in zip(xs, self.image_keys, strict=True)] - ) - logits = self.classifier_head(encoder_outputs) - - if self.config.num_classes == 2: - logits = logits.squeeze(-1) - probabilities = torch.sigmoid(logits) - else: - probabilities = torch.softmax(logits, dim=-1) - - return ClassifierOutput(logits=logits, probabilities=probabilities, hidden_states=encoder_outputs) - - def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Tensor]]: - """Standard forward pass for training compatible with train.py.""" - # Extract images and labels - images, labels = self.extract_images_and_labels(batch) - - # Get predictions - outputs = self.predict(images) - - # Calculate loss - if self.config.num_classes == 2: - # Binary classification - loss = nn.functional.binary_cross_entropy_with_logits(outputs.logits, labels) - predictions = (torch.sigmoid(outputs.logits) > 0.5).float() - else: - # Multi-class classification - loss = nn.functional.cross_entropy(outputs.logits, labels.long()) - predictions = torch.argmax(outputs.logits, dim=1) - - # Calculate accuracy for logging - correct = (predictions == labels).sum().item() - total = labels.size(0) - accuracy = 100 * correct / total - - # Return loss and metrics for logging - output_dict = { - "accuracy": accuracy, - "correct": correct, - "total": total, - } - - return loss, output_dict - - def predict_reward(self, batch, threshold=0.5): - """Eval method. Returns predicted reward with the decision threshold as argument.""" - # Check for both OBS_IMAGE and OBS_IMAGES prefixes - batch = self.normalize_inputs(batch) - batch = self.normalize_targets(batch) - - # Extract images from batch dict - images = [batch[key] for key in self.config.input_features if key.startswith(OBS_IMAGE)] - - if self.config.num_classes == 2: - probs = self.predict(images).probabilities - logging.debug(f"Predicted reward images: {probs}") - return (probs > threshold).float() - else: - return torch.argmax(self.predict(images).probabilities, dim=1) - - def get_optim_params(self): - """Return optimizer parameters for the policy.""" - return self.parameters() - - def select_action(self, batch: dict[str, Tensor]) -> Tensor: - """ - This method is required by PreTrainedPolicy but not used for reward classifiers. - The reward classifier is not an actor and does not select actions. - """ - raise NotImplementedError("Reward classifiers do not select actions") - - def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor: - """ - This method is required by PreTrainedPolicy but not used for reward classifiers. - The reward classifier is not an actor and does not produce action chunks. - """ - raise NotImplementedError("Reward classifiers do not predict action chunks") - - def reset(self): - """ - This method is required by PreTrainedPolicy but not used for reward classifiers. - The reward classifier is not an actor and does not select actions. - """ - pass diff --git a/src/lerobot/policies/sac/reward_model/processor_classifier.py b/src/lerobot/policies/sac/reward_model/processor_classifier.py deleted file mode 100644 index c2a34eab2..000000000 --- a/src/lerobot/policies/sac/reward_model/processor_classifier.py +++ /dev/null @@ -1,82 +0,0 @@ -# !/usr/bin/env python - -# Copyright 2025 The HuggingFace Inc. team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from typing import Any - -import torch - -from lerobot.policies.sac.reward_model.configuration_classifier import RewardClassifierConfig -from lerobot.processor import ( - DeviceProcessorStep, - IdentityProcessorStep, - NormalizerProcessorStep, - PolicyAction, - PolicyProcessorPipeline, -) -from lerobot.processor.converters import policy_action_to_transition, transition_to_policy_action - - -def make_classifier_processor( - config: RewardClassifierConfig, - dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, -) -> tuple[ - PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], - PolicyProcessorPipeline[PolicyAction, PolicyAction], -]: - """ - Constructs pre-processor and post-processor pipelines for the reward classifier. - - The pre-processing pipeline prepares input data for the classifier by: - 1. Normalizing both input and output features based on dataset statistics. - 2. Moving the data to the specified device. - - The post-processing pipeline handles the classifier's output by: - 1. Moving the data to the CPU. - 2. Applying an identity step, as no unnormalization is needed for the output logits. - - Args: - config: The configuration object for the RewardClassifier. - dataset_stats: A dictionary of statistics for normalization. - preprocessor_kwargs: Additional arguments for the pre-processor pipeline. - postprocessor_kwargs: Additional arguments for the post-processor pipeline. - - Returns: - A tuple containing the configured pre-processor and post-processor pipelines. - """ - - input_steps = [ - NormalizerProcessorStep( - features=config.input_features, norm_map=config.normalization_mapping, stats=dataset_stats - ), - NormalizerProcessorStep( - features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats - ), - DeviceProcessorStep(device=config.device), - ] - output_steps = [DeviceProcessorStep(device="cpu"), IdentityProcessorStep()] - - return ( - PolicyProcessorPipeline( - steps=input_steps, - name="classifier_preprocessor", - ), - PolicyProcessorPipeline( - steps=output_steps, - name="classifier_postprocessor", - to_transition=policy_action_to_transition, - to_output=transition_to_policy_action, - ), - ) diff --git a/src/lerobot/policies/sarm/compute_rabc_weights.py b/src/lerobot/policies/sarm/compute_rabc_weights.py deleted file mode 100644 index 485c1096b..000000000 --- a/src/lerobot/policies/sarm/compute_rabc_weights.py +++ /dev/null @@ -1,870 +0,0 @@ -#!/usr/bin/env python - -# Copyright 2024 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 SARM progress values for RA-BC (Reward-Aware Behavior Cloning) weighting. - -This script processes all frames in a dataset with SARM to compute progress values [0, 1]. -The results are saved as a parquet file that can be loaded during training for RA-BC weighting. - -Uses multi-output extraction: each SARM query returns progress for 9 frames, so we only -need ~num_frames/30 queries instead of one per frame (~30x speedup). - -Usage: - # Full RA-BC computation with visualizations - python src/lerobot/policies/sarm/compute_rabc_weights.py \\ - --dataset-repo-id lerobot/aloha_sim_insertion_human \\ - --reward-model-path /sarm_single_uni4 - - # Faster computation with stride (compute every 5 frames, interpolate the rest) - python src/lerobot/policies/sarm/compute_rabc_weights.py \\ - --dataset-repo-id lerobot/aloha_sim_insertion_human \\ - --reward-model-path /sarm_single_uni4 \\ - --stride 5 - - # Visualize predictions only (no RA-BC computation) - python src/lerobot/policies/sarm/compute_rabc_weights.py \\ - --dataset-repo-id lerobot/aloha_sim_insertion_human \\ - --reward-model-path /sarm_single_uni4 \\ - --visualize-only \\ - --num-visualizations 5 - -The output is saved to the dataset's local cache directory as 'sarm_progress.parquet'. -""" - -import argparse -import logging -from pathlib import Path - -import matplotlib.gridspec as gridspec -import matplotlib.pyplot as plt -import numpy as np -import pyarrow as pa -import pyarrow.parquet as pq -import torch -from tqdm import tqdm - -from lerobot.datasets.lerobot_dataset import LeRobotDataset -from lerobot.policies.sarm.modeling_sarm import SARMRewardModel -from lerobot.policies.sarm.processor_sarm import make_sarm_pre_post_processors -from lerobot.policies.sarm.sarm_utils import normalize_stage_tau - - -def get_reward_model_path_from_parquet(parquet_path: Path) -> str | None: - """Read reward_model_path from parquet metadata if available.""" - if not parquet_path.exists(): - return None - try: - metadata = pq.read_metadata(parquet_path).schema.to_arrow_schema().metadata - if metadata and b"reward_model_path" in metadata: - return metadata[b"reward_model_path"].decode() - except Exception: # nosec B110 - return None - return None - - -def load_sarm_resources( - dataset_repo_id: str, - reward_model_path: str, - device: str = "cuda", -) -> tuple[LeRobotDataset, SARMRewardModel, any]: - """ - Load SARM model, dataset, and preprocessor. - - Returns: - Tuple of (dataset, reward_model, preprocessor) - """ - logging.info(f"Loading model: {reward_model_path}") - reward_model = SARMRewardModel.from_pretrained(reward_model_path) - reward_model.config.device = device - reward_model.to(device).eval() - - image_key = reward_model.config.image_key - state_key = reward_model.config.state_key - delta_indices = reward_model.config.observation_delta_indices - - logging.info(f"Loading dataset: {dataset_repo_id}") - temp_dataset = LeRobotDataset(dataset_repo_id, download_videos=True) - fps = temp_dataset.fps - - delta_timestamps = { - image_key: [idx / fps for idx in delta_indices], - state_key: [idx / fps for idx in delta_indices], - } - dataset = LeRobotDataset(dataset_repo_id, delta_timestamps=delta_timestamps) - logging.info(f"Dataset: {dataset.num_episodes} episodes, {dataset.num_frames} frames") - - preprocess, _ = make_sarm_pre_post_processors( - config=reward_model.config, - dataset_stats=dataset.meta.stats, - dataset_meta=dataset.meta, - ) - - return dataset, reward_model, preprocess - - -def to_numpy_image(img) -> np.ndarray: - """Convert image tensor to numpy uint8 (H, W, C).""" - if isinstance(img, torch.Tensor): - img = img.cpu().numpy() - if img.ndim == 4: - # Take center frame for bidirectional sampling - img = img[img.shape[0] // 2] - if img.shape[0] in [1, 3]: - img = np.transpose(img, (1, 2, 0)) - if img.dtype != np.uint8: - # Handle normalized images (may have negative values or values > 1) - img = img.astype(np.float32) - img = (img - img.min()) / (img.max() - img.min() + 1e-8) # Normalize to [0, 1] - img = (img * 255).astype(np.uint8) - return img - - -def visualize_episode( - frames, progress_preds, stage_preds, title, output_path, stage_labels, gt_progress=None, gt_stages=None -): - """Create visualization with progress plot, stage probabilities, and sample frames. - - Same as sarm_inference_visualization.py - """ - num_stages = stage_preds.shape[1] - colors = plt.cm.tab10(np.linspace(0, 1, num_stages)) - frame_indices = np.arange(len(progress_preds)) - - fig = plt.figure(figsize=(14, 12)) - gs = gridspec.GridSpec(3, 1, height_ratios=[2, 1, 1], hspace=0.3) - ax_progress, ax_stages, ax_frames = fig.add_subplot(gs[0]), fig.add_subplot(gs[1]), fig.add_subplot(gs[2]) - - # Progress plot - ax_progress.plot(frame_indices, progress_preds, linewidth=2, color="#2E86AB", label="Predicted") - ax_progress.fill_between(frame_indices, 0, progress_preds, alpha=0.3, color="#2E86AB") - if gt_progress is not None: - ax_progress.plot( - frame_indices, gt_progress, linewidth=2, color="#28A745", linestyle="--", label="Ground Truth" - ) - ax_progress.axhline(y=1.0, color="gray", linestyle="--", alpha=0.5) - ax_progress.set_ylabel("Progress") - ax_progress.set_title(f'Task: "{title}"', fontweight="bold") - ax_progress.set_ylim(-0.05, 1.1) - ax_progress.legend(loc="upper left") - ax_progress.grid(True, alpha=0.3) - - # Stage predictions - ax_stages.stackplot( - frame_indices, - *[stage_preds[:, i] for i in range(num_stages)], - colors=colors, - alpha=0.8, - labels=stage_labels, - ) - if gt_stages is not None: - for change_idx in np.where(np.diff(gt_stages) != 0)[0] + 1: - ax_stages.axvline(x=change_idx, color="black", linestyle="-", alpha=0.7, linewidth=1.5) - ax_stages.set_xlabel("Frame") - ax_stages.set_ylabel("Stage Probability") - ax_stages.set_ylim(0, 1) - ax_stages.legend(loc="upper left", ncol=min(num_stages, 5), fontsize=8) - ax_stages.grid(True, alpha=0.3) - - # Sample frames - ax_frames.axis("off") - num_sample = 8 - sample_indices = np.linspace(0, len(frames) - 1, num_sample, dtype=int) - h, w = frames[0].shape[:2] - combined = np.zeros((h, w * num_sample, 3), dtype=np.uint8) - for i, idx in enumerate(sample_indices): - frame = frames[idx] - if frame.shape[-1] == 1: - frame = np.repeat(frame, 3, axis=-1) - combined[:, i * w : (i + 1) * w] = frame - stage_name = stage_labels[np.argmax(stage_preds[idx])][:12] - ax_frames.text( - i * w + w / 2, - -10, - f"Frame {idx}\n{progress_preds[idx]:.2f}\n{stage_name}", - ha="center", - va="top", - fontsize=7, - ) - ax_frames.imshow(combined) - ax_frames.set_title("Sample Frames", pad=20) - - output_path.parent.mkdir(parents=True, exist_ok=True) - plt.savefig(output_path, dpi=150, bbox_inches="tight") - plt.close() - print(f"Saved: {output_path}") - - -def visualize_sarm_predictions( - dataset: LeRobotDataset, - reward_model: SARMRewardModel, - preprocess, - episode_indices: list[int], - head_mode: str, - output_dir: Path, - num_display_frames: int = 5, - stride: int = 1, -): - """ - Visualize SARM predictions for multiple episodes. - - Computes predictions for every frame by default. With stride > 1, computes predictions - every N frames and interpolates (progress + stage probabilities) for visualization. - - Args: - dataset: LeRobotDataset with delta_timestamps configured - reward_model: Loaded SARM model - preprocess: Preprocessor from make_sarm_pre_post_processors - episode_indices: List of episode indices to visualize - head_mode: "sparse", "dense", or "both" - output_dir: Directory to save visualizations - num_display_frames: Number of frames to display in thumbnail strip (default: 5) - stride: Compute predictions every N frames, interpolate the rest (default: 1) - """ - output_dir = Path(output_dir) - output_dir.mkdir(parents=True, exist_ok=True) - - image_key = reward_model.config.image_key - state_key = reward_model.config.state_key - dual_mode = reward_model.config.uses_dual_heads - device = reward_model.device - - # Center frame index for bidirectional sampling - target_idx = reward_model.config.n_obs_steps // 2 - - # Determine which heads to visualize - schemes_to_viz = [] - if head_mode in ("sparse", "both") or not dual_mode: - schemes_to_viz.append("sparse") - if head_mode in ("dense", "both") and dual_mode: - schemes_to_viz.append("dense") - - # Set preprocessor to eval mode to disable augmentations - if hasattr(preprocess, "eval"): - preprocess.eval() - for step in preprocess.steps: - if hasattr(step, "eval"): - step.eval() - - for episode_idx in episode_indices: - ep = dataset.meta.episodes[episode_idx] - ep_start = ep["dataset_from_index"] - ep_end = ep["dataset_to_index"] - task = dataset[ep_start].get("task", "perform the task") - num_frames = ep_end - ep_start - - # Select frames for display thumbnails (evenly sampled from begin to end) - display_indices = set( - [ - ep_start + int(i * (num_frames - 1) / (num_display_frames - 1)) - for i in range(num_display_frames) - ] - if num_frames >= num_display_frames - else list(range(ep_start, ep_end)) - ) - viz_frames = {} - - # Load display frames up-front (stride mode might skip them otherwise). - for frame_idx in display_indices: - sample = dataset[frame_idx] - viz_frames[frame_idx] = to_numpy_image(sample[image_key]) - - # Initialize storage for each scheme - scheme_data = {} - for scheme in schemes_to_viz: - num_stages = getattr(reward_model.config, f"num_{scheme}_stages") - scheme_data[scheme] = { - "viz_progress": np.full(num_frames, np.nan), - "viz_stages": np.full((num_frames, num_stages), np.nan), - "viz_gt_progress": np.full(num_frames, np.nan), - "viz_gt_stages": np.full(num_frames, np.nan), - "target_key": f"{scheme}_targets", - "num_stages": num_stages, - "temporal_props": getattr(reward_model.config, f"{scheme}_temporal_proportions"), - "subtask_names": getattr(reward_model.config, f"{scheme}_subtask_names"), - } - - if stride > 1: - logging.info(f"Visualization stride={stride}: inferring every {stride} frames and interpolating") - - # Process frames one at a time to avoid memory buildup - frame_indices = list(range(ep_start, ep_end, stride)) - if (ep_end - 1) not in frame_indices: - frame_indices.append(ep_end - 1) - frame_indices = sorted(set(frame_indices)) - - for frame_idx in tqdm(frame_indices, desc=f"Episode {episode_idx}", leave=False): - local_idx = frame_idx - ep_start - sample = dataset[frame_idx] - - batch = { - image_key: sample[image_key], - "task": task, - "index": frame_idx, - "episode_index": episode_idx, - } - if state_key in sample: - batch[state_key] = sample[state_key] - - with torch.no_grad(): - processed = preprocess(batch) - video_features = processed["video_features"].to(device) - text_features = processed["text_features"].to(device) - state_features = processed.get("state_features") - if state_features is not None: - state_features = state_features.to(device) - lengths = processed.get("lengths") - - for scheme in schemes_to_viz: - sd = scheme_data[scheme] - - # Ground truth - # In stride visualization mode, ground-truth plots can be misleading - # (only sparse points are available), so we skip GT. - if stride == 1 and sd["target_key"] in processed: - gt_target = processed[sd["target_key"]][0, target_idx].cpu().item() - sd["viz_gt_stages"][local_idx] = int(gt_target) - sd["viz_gt_progress"][local_idx] = normalize_stage_tau( - gt_target, - num_stages=sd["num_stages"], - temporal_proportions=sd["temporal_props"], - subtask_names=sd["subtask_names"], - ) - - # Predictions - reward, stage_probs = reward_model.calculate_rewards( - text_embeddings=text_features, - video_embeddings=video_features, - state_features=state_features, - lengths=lengths, - return_all_frames=True, - return_stages=True, - head_mode=scheme, - ) - - # Handle both tensor and numpy outputs - if isinstance(reward, torch.Tensor): - reward = reward.cpu().numpy() - stage_probs = stage_probs.cpu().numpy() - - if reward.ndim == 2: - sd["viz_progress"][local_idx] = reward[0, target_idx] - sd["viz_stages"][local_idx] = stage_probs[0, target_idx, :] - else: - sd["viz_progress"][local_idx] = reward[target_idx] - sd["viz_stages"][local_idx] = stage_probs[target_idx, :] - - # Clear GPU memory after each frame - del processed, video_features, text_features - if state_features is not None: - del state_features - - torch.cuda.empty_cache() - - # Interpolate predictions back to per-frame arrays for smooth visualization. - if stride > 1: - all_local = np.arange(num_frames) - for scheme in schemes_to_viz: - sd = scheme_data[scheme] - - valid = np.isfinite(sd["viz_progress"]) - valid_idx = np.where(valid)[0] - if valid_idx.size >= 1: - sd["viz_progress"] = interpolate_progress( - valid_idx, sd["viz_progress"][valid_idx], all_local - ) - - stage_interp = np.zeros_like(sd["viz_stages"], dtype=np.float32) - for s in range(sd["num_stages"]): - stage_interp[:, s] = interpolate_progress( - valid_idx, sd["viz_stages"][valid_idx, s], all_local - ) - - stage_interp = np.clip(stage_interp, 0.0, 1.0) - row_sums = stage_interp.sum(axis=1, keepdims=True) - nz = row_sums.squeeze(-1) > 0 - stage_interp[nz] = stage_interp[nz] / row_sums[nz] - sd["viz_stages"] = stage_interp - else: - # No valid points: keep NaNs/zeros; visualization will be empty. - sd["viz_stages"] = np.nan_to_num(sd["viz_stages"], nan=0.0) - - # Generate visualization for each head - ordered_viz_frames = [viz_frames[idx] for idx in sorted(display_indices)] - for scheme in schemes_to_viz: - sd = scheme_data[scheme] - stage_labels = sd["subtask_names"] or [f"Stage {i + 1}" for i in range(sd["num_stages"])] - viz_path = output_dir / f"sarm_prediction_ep{episode_idx}_{scheme}.png" - - visualize_episode( - frames=np.array(ordered_viz_frames), - progress_preds=sd["viz_progress"], - stage_preds=sd["viz_stages"], - title=f"{task} (Episode {episode_idx})", - output_path=viz_path, - stage_labels=stage_labels, - gt_progress=sd["viz_gt_progress"] if not np.all(np.isnan(sd["viz_gt_progress"])) else None, - gt_stages=sd["viz_gt_stages"] if not np.all(np.isnan(sd["viz_gt_stages"])) else None, - ) - - # Clear memory between episodes - torch.cuda.empty_cache() - - logging.info(f"Visualizations saved to: {output_dir.absolute()}") - - -def generate_all_frame_indices(ep_start: int, ep_end: int, frame_gap: int = 30) -> list[int]: - """Generate all frame indices, ordered by offset for cache-friendly access. - - Orders frames as: [0, 30, 60...], [1, 31, 61...], ..., [29, 59, 89...] - This groups frames that share similar temporal windows together. - """ - num_frames = ep_end - ep_start - indices = [] - for offset in range(frame_gap): - for frame_rel in range(offset, num_frames, frame_gap): - indices.append(ep_start + frame_rel) - return indices - - -def interpolate_progress( - computed_indices: np.ndarray, - computed_values: np.ndarray, - all_indices: np.ndarray, -) -> np.ndarray: - """Linearly interpolate values to fill in gaps (robust to NaNs / edge cases).""" - computed_indices = np.asarray(computed_indices) - computed_values = np.asarray(computed_values) - all_indices = np.asarray(all_indices) - - mask = np.isfinite(computed_values) - if mask.sum() == 0: - return np.full(all_indices.shape, np.nan, dtype=np.float32) - if mask.sum() == 1: - return np.full(all_indices.shape, float(computed_values[mask][0]), dtype=np.float32) - - out = np.interp(all_indices, computed_indices[mask], computed_values[mask]) - return out.astype(np.float32) - - -def compute_sarm_progress( - dataset_repo_id: str, - reward_model_path: str, - output_path: str | None = None, - head_mode: str = "sparse", - device: str = "cuda", - num_visualizations: int = 5, - output_dir: str = "./sarm_viz", - stride: int = 1, -): - """ - Compute SARM progress predictions for all frames in a dataset. - - Args: - dataset_repo_id: HuggingFace dataset repo ID or local path - reward_model_path: Path to pretrained SARM model - output_path: Path to save results. If None, saves to dataset's cache directory - head_mode: SARM head to use ("sparse", "dense", or "both") - device: Device to use for inference - num_visualizations: Number of episodes to visualize (0 to skip) - output_dir: Directory to save visualizations - stride: Compute progress every N frames, interpolate the rest (default: 1 = every frame) - """ - dataset, reward_model, preprocess = load_sarm_resources(dataset_repo_id, reward_model_path, device) - - # Set preprocessor to eval mode to disable augmentations - if hasattr(preprocess, "eval"): - preprocess.eval() - for step in preprocess.steps: - if hasattr(step, "eval"): - step.eval() - - image_key = reward_model.config.image_key - state_key = reward_model.config.state_key - frame_gap = reward_model.config.frame_gap - num_episodes = dataset.num_episodes - total_frames = dataset.num_frames - logging.info(f"Processing {total_frames} frames across {num_episodes} episodes") - - # Determine which heads to compute - dual_mode = reward_model.config.uses_dual_heads - compute_sparse = head_mode in ("sparse", "both") or not dual_mode - compute_dense = head_mode in ("dense", "both") and dual_mode - - # Storage arrays - all_indices = [] - all_episode_indices = [] - all_frame_indices = [] - all_progress_sparse = [] if compute_sparse else None - all_progress_dense = [] if compute_dense else None - - if stride > 1: - logging.info(f"Using stride={stride}: computing every {stride} frames, interpolating the rest") - - # Process all episodes - for episode_idx in tqdm(range(num_episodes), desc="Episodes"): - ep = dataset.meta.episodes[episode_idx] - ep_start = ep["dataset_from_index"] - ep_end = ep["dataset_to_index"] - - # Get task description - task = dataset[ep_start].get("task", "perform the task") - - # Generate frames to compute (with stride applied) - all_ep_indices = generate_all_frame_indices(ep_start, ep_end, frame_gap) - if stride > 1: - # Only compute every stride-th frame (relative to episode start) - compute_indices = [idx for idx in all_ep_indices if (idx - ep_start) % stride == 0] - # Always include last frame for better interpolation at episode end - last_frame = ep_end - 1 - if last_frame not in compute_indices: - compute_indices.append(last_frame) - compute_indices = sorted(set(compute_indices)) - else: - compute_indices = all_ep_indices - - center_idx = reward_model.config.n_obs_steps // 2 # Center of bidirectional window - - # Dictionary to collect results - frame_results = {} - - for query_idx in tqdm(compute_indices, desc=f" Ep {episode_idx}", leave=False): - try: - sample = dataset[query_idx] - - batch = { - image_key: sample[image_key], - "task": task, - "index": query_idx, - "episode_index": episode_idx, - } - if state_key in sample: - batch[state_key] = sample[state_key] - - with torch.no_grad(): - processed = preprocess(batch) - video_features = processed["video_features"].to(device) - text_features = processed["text_features"].to(device) - state_features = processed.get("state_features") - if state_features is not None: - state_features = state_features.to(device) - lengths = processed.get("lengths") - - sparse_val = np.nan - dense_val = np.nan - - # Compute sparse prediction for center frame - if compute_sparse: - sparse_progress = reward_model.calculate_rewards( - text_embeddings=text_features, - video_embeddings=video_features, - state_features=state_features, - lengths=lengths, - return_all_frames=True, - head_mode="sparse", - ) - sparse_val = float( - sparse_progress[0, center_idx] - if sparse_progress.ndim == 2 - else sparse_progress[center_idx] - ) - - # Compute dense prediction for center frame - if compute_dense: - dense_progress = reward_model.calculate_rewards( - text_embeddings=text_features, - video_embeddings=video_features, - state_features=state_features, - lengths=lengths, - return_all_frames=True, - head_mode="dense", - ) - dense_val = float( - dense_progress[0, center_idx] - if dense_progress.ndim == 2 - else dense_progress[center_idx] - ) - - frame_results[query_idx] = (sparse_val, dense_val) - - except Exception as e: - logging.warning(f"Failed to process frame {query_idx}: {e}") - - # Interpolate to get values for all frames - computed_indices = np.array(sorted(frame_results.keys())) - computed_sparse = ( - np.array([frame_results[i][0] for i in computed_indices]) if compute_sparse else None - ) - computed_dense = np.array([frame_results[i][1] for i in computed_indices]) if compute_dense else None - - # All frame indices for this episode - all_frame_idx_array = np.arange(ep_start, ep_end) - - if stride > 1 and len(computed_indices) > 1: - # Interpolate progress values - if compute_sparse: - interp_sparse = interpolate_progress(computed_indices, computed_sparse, all_frame_idx_array) - if compute_dense: - interp_dense = interpolate_progress(computed_indices, computed_dense, all_frame_idx_array) - else: - # No interpolation needed - interp_sparse = computed_sparse if compute_sparse else None - interp_dense = computed_dense if compute_dense else None - - # Store results for all frames - for i, frame_idx in enumerate(all_frame_idx_array): - local_idx = frame_idx - ep_start - all_indices.append(frame_idx) - all_episode_indices.append(episode_idx) - all_frame_indices.append(local_idx) - if compute_sparse: - if stride > 1 and len(computed_indices) > 1: - all_progress_sparse.append(float(interp_sparse[i])) - elif frame_idx in frame_results: - all_progress_sparse.append(frame_results[frame_idx][0]) - else: - all_progress_sparse.append(np.nan) - if compute_dense: - if stride > 1 and len(computed_indices) > 1: - all_progress_dense.append(float(interp_dense[i])) - elif frame_idx in frame_results: - all_progress_dense.append(frame_results[frame_idx][1]) - else: - all_progress_dense.append(np.nan) - - # Create output table - table_data = { - "index": np.array(all_indices, dtype=np.int64), - "episode_index": np.array(all_episode_indices, dtype=np.int64), - "frame_index": np.array(all_frame_indices, dtype=np.int64), - } - if compute_sparse: - table_data["progress_sparse"] = np.array(all_progress_sparse, dtype=np.float32) - if compute_dense: - table_data["progress_dense"] = np.array(all_progress_dense, dtype=np.float32) - - # Sort by index - df = pa.table(table_data).to_pandas() - df = df.sort_values("index").reset_index(drop=True) - final_table = pa.Table.from_pandas(df, preserve_index=False) - - # Add metadata with reward model path - metadata = {b"reward_model_path": reward_model_path.encode()} - final_table = final_table.replace_schema_metadata(metadata) - - # Determine output path - output_path = Path(dataset.root) / "sarm_progress.parquet" if output_path is None else Path(output_path) - - # Save - output_path.parent.mkdir(parents=True, exist_ok=True) - pq.write_table(final_table, output_path) - logging.info(f"Saved {len(final_table)} frame progress values to {output_path}") - - # Print statistics - if "progress_sparse" in df.columns: - valid = df["progress_sparse"].dropna() - logging.info( - f"Sparse progress: mean={valid.mean():.4f}, std={valid.std():.4f}, " - f"min={valid.min():.4f}, max={valid.max():.4f}" - ) - - if "progress_dense" in df.columns: - valid = df["progress_dense"].dropna() - logging.info( - f"Dense progress: mean={valid.mean():.4f}, std={valid.std():.4f}, " - f"min={valid.min():.4f}, max={valid.max():.4f}" - ) - - # Visualize episodes after processing - if num_visualizations > 0: - viz_episodes = list(range(min(num_visualizations, num_episodes))) - logging.info(f"Generating {len(viz_episodes)} visualizations...") - visualize_sarm_predictions( - dataset=dataset, - reward_model=reward_model, - preprocess=preprocess, - episode_indices=viz_episodes, - head_mode=head_mode, - output_dir=Path(output_dir), - stride=stride, - ) - - return output_path - - -def main(): - parser = argparse.ArgumentParser( - description="Compute SARM progress values for RA-BC weighting or visualize SARM predictions", - formatter_class=argparse.RawDescriptionHelpFormatter, - epilog=""" -Examples: - # Full RA-BC computation with visualizations - python src/lerobot/policies/sarm/compute_rabc_weights.py \\ - --dataset-repo-id lerobot/aloha_sim_insertion_human \\ - --reward-model-path /sarm_single_uni4 - - # Visualize predictions only (no RA-BC computation) - python src/lerobot/policies/sarm/compute_rabc_weights.py \\ - --dataset-repo-id lerobot/aloha_sim_insertion_human \\ - --reward-model-path /sarm_single_uni4 \\ - --visualize-only \\ - --num-visualizations 10 - """, - ) - parser.add_argument( - "--dataset-repo-id", - type=str, - required=True, - help="HuggingFace dataset repo ID or local path", - ) - parser.add_argument( - "--reward-model-path", - type=str, - default=None, - help="Path to pretrained SARM model (reads from existing parquet metadata if not provided)", - ) - parser.add_argument( - "--output-path", - type=str, - default=None, - help="Output path for parquet. If not set, saves to dataset's cache directory", - ) - parser.add_argument( - "--head-mode", - type=str, - default="sparse", - choices=["sparse", "dense", "both"], - help="SARM head to use (default: sparse)", - ) - parser.add_argument( - "--device", - type=str, - default="cuda", - help="Device to use (default: cuda)", - ) - # Visualization options - parser.add_argument( - "--visualize-only", - action="store_true", - help="Only visualize SARM predictions (no RA-BC computation)", - ) - parser.add_argument( - "--num-visualizations", - type=int, - default=5, - help="Number of episodes to visualize (default: 5, set to 0 to skip)", - ) - parser.add_argument( - "--output-dir", - type=str, - default="./sarm_viz", - help="Output directory for visualizations (default: ./sarm_viz)", - ) - parser.add_argument( - "--push-to-hub", - action="store_true", - help="Upload progress file to the dataset repo on HuggingFace Hub", - default=True, - ) - parser.add_argument( - "--stride", - type=int, - default=1, - help="Compute progress every N frames, interpolate the rest (default: 1 = every frame)", - ) - - args = parser.parse_args() - - logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") - - # Try to get reward_model_path from parquet metadata if not provided - reward_model_path = args.reward_model_path - if reward_model_path is None: - # Load dataset to find parquet path - temp_dataset = LeRobotDataset(args.dataset_repo_id, download_videos=False) - parquet_path = Path(temp_dataset.root) / "sarm_progress.parquet" - reward_model_path = get_reward_model_path_from_parquet(parquet_path) - if reward_model_path: - logging.info(f"Using reward model from parquet metadata: {reward_model_path}") - else: - raise ValueError( - "--reward-model-path is required (no existing parquet with model metadata found)" - ) - - # Handle visualize-only mode - if args.visualize_only: - dataset, reward_model, preprocess = load_sarm_resources( - args.dataset_repo_id, reward_model_path, args.device - ) - logging.info(f"Visualization-only mode: visualizing {args.num_visualizations} episodes") - viz_episodes = list(range(min(args.num_visualizations, dataset.num_episodes))) - visualize_sarm_predictions( - dataset=dataset, - reward_model=reward_model, - preprocess=preprocess, - episode_indices=viz_episodes, - head_mode=args.head_mode, - output_dir=Path(args.output_dir), - stride=args.stride, - ) - print(f"\nVisualizations saved to: {Path(args.output_dir).absolute()}") - return - - # Full RABC computation (compute_sarm_progress loads model/dataset itself) - output_path = compute_sarm_progress( - dataset_repo_id=args.dataset_repo_id, - reward_model_path=reward_model_path, - output_path=args.output_path, - head_mode=args.head_mode, - device=args.device, - num_visualizations=args.num_visualizations, - output_dir=args.output_dir, - stride=args.stride, - ) - - print(f"\nSARM progress values saved to: {output_path}") - - # Upload to Hub if requested - if args.push_to_hub: - from huggingface_hub import HfApi - - api = HfApi() - hub_path = "sarm_progress.parquet" - - print(f"\nUploading to Hub: {args.dataset_repo_id}/{hub_path}") - api.upload_file( - path_or_fileobj=str(output_path), - path_in_repo=hub_path, - repo_id=args.dataset_repo_id, - repo_type="dataset", - ) - print( - f"Successfully uploaded to: https://huggingface.co/datasets/{args.dataset_repo_id}/blob/main/{hub_path}" - ) - - print("\nTo use in training, add to your config:") - print(" use_rabc: true") - print(f" rabc_progress_path: hf://datasets/{args.dataset_repo_id}/{hub_path}") - print(" rabc_head_mode: sparse # or dense") - else: - print("\nTo use in training, add to your config:") - print(" use_rabc: true") - print(f" rabc_progress_path: {output_path}") - print(" rabc_head_mode: sparse # or dense") - - -if __name__ == "__main__": - main() diff --git a/src/lerobot/policies/sarm/configuration_sarm.py b/src/lerobot/policies/sarm/configuration_sarm.py deleted file mode 100644 index 673422fe2..000000000 --- a/src/lerobot/policies/sarm/configuration_sarm.py +++ /dev/null @@ -1,249 +0,0 @@ -#!/usr/bin/env python - -# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu -# and 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. - -""" -SARM: Stage-Aware Reward Modeling for Long Horizon Robot Manipulation. -Paper: https://arxiv.org/abs/2509.25358 -""" - -from dataclasses import dataclass, field - -from lerobot.configs.policies import PreTrainedConfig -from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature -from lerobot.optim.optimizers import AdamWConfig -from lerobot.optim.schedulers import CosineDecayWithWarmupSchedulerConfig -from lerobot.utils.constants import OBS_IMAGES, OBS_STATE - - -@PreTrainedConfig.register_subclass("sarm") -@dataclass -class SARMConfig(PreTrainedConfig): - """Configuration class for SARM (Stage-Aware Reward Modeling). - - Supports three annotation modes: - - 1. single_stage (default): No annotations needed. Uses the episode's task description - as a single stage covering the entire episode. - - 2. dense_only: Uses dense (fine-grained) annotations from VLM, with an auto-generated - single sparse "task" stage covering the full episode. The dense head learns detailed - subtask progression while sparse provides overall task completion. - - 3. dual: Full dual-head mode with both sparse (high-level) and dense (fine-grained) - annotations from VLM. Both heads are trained on their respective annotations. - - The annotation_mode determines how sparse_temporal_proportions and dense_temporal_proportions - are loaded/generated during model initialization. - """ - - annotation_mode: str = "single_stage" # "single_stage", "dense_only", or "dual" - n_obs_steps: int = 8 # Number of observation history steps - frame_gap: int = 30 # Frame gap between frames (at 30 fps = 1 second) - max_rewind_steps: int = 4 # Maximum rewind steps for temporal augmentation - - # Total frames = 1 + n_obs_steps + max_rewind_steps (computed in property) - # During training with rewind: [obs_frames] + [rewind_frames] - # During inference: [obs_frames] only - - # Architecture params - image_dim: int = 512 - text_dim: int = 512 - hidden_dim: int = 768 - num_heads: int = 12 - num_layers: int = 8 - max_state_dim: int = 32 - drop_n_last_frames: int = 1 - batch_size: int = 64 - clip_batch_size: int = 64 - dropout: float = 0.1 - stage_loss_weight: float = 1.0 # Weight for stage classification loss when using subtask annotations - - rewind_probability: float = 0.8 - language_perturbation_probability: float = 0.2 - - # Sparse annotations (high-level stages) - num_sparse_stages: int = 1 - sparse_subtask_names: list | None = None - sparse_temporal_proportions: list | None = None - - # Dense annotations (fine-grained stages) - num_dense_stages: int | None = None - dense_subtask_names: list | None = None - dense_temporal_proportions: list | None = None - - pretrained_model_path: str | None = None - device: str | None = None - image_key: str = OBS_IMAGES + ".top" # Key for image used from the dataset - state_key: str = OBS_STATE - - # Populated by the processor (video_features, state_features, text_features) - input_features: dict = field(default_factory=lambda: {}) - - # Output features (updated in __post_init__) - output_features: dict = field( - default_factory=lambda: { - "stage": PolicyFeature(shape=(9, 5), type=FeatureType.REWARD), - "progress": PolicyFeature(shape=(9, 1), type=FeatureType.REWARD), - } - ) - - normalization_mapping: dict[str, NormalizationMode] = field( - default_factory=lambda: { - "VISUAL": NormalizationMode.IDENTITY, - "STATE": NormalizationMode.MEAN_STD, - "LANGUAGE": NormalizationMode.IDENTITY, - "REWARD": NormalizationMode.IDENTITY, - } - ) - - def __post_init__(self): - super().__post_init__() - - if self.annotation_mode not in ["single_stage", "dense_only", "dual"]: - raise ValueError( - f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}" - ) - - if self.annotation_mode == "single_stage": - # Use task description as stage name, full episode as one stage - self.num_sparse_stages = 1 - self.sparse_subtask_names = ["task"] - self.sparse_temporal_proportions = [1.0] - self.num_dense_stages = None - self.dense_subtask_names = None - self.dense_temporal_proportions = None - - elif self.annotation_mode == "dense_only": - self.num_sparse_stages = 1 - self.sparse_subtask_names = ["task"] - self.sparse_temporal_proportions = [1.0] - - self.input_features = {} - self.output_features = {} - - if self.image_key: - self.input_features[self.image_key] = PolicyFeature(shape=(480, 640, 3), type=FeatureType.VISUAL) - - self.input_features[self.state_key] = PolicyFeature( - shape=(self.max_state_dim,), - type=FeatureType.STATE, - ) - - # Update output features based on annotation_mode - if self.annotation_mode in ["dense_only", "dual"]: - self.output_features["sparse_stage"] = PolicyFeature( - shape=(self.num_frames, self.num_sparse_stages), type=FeatureType.REWARD - ) - self.output_features["sparse_progress"] = PolicyFeature( - shape=(self.num_frames, 1), type=FeatureType.REWARD - ) - dense_stages = self.num_dense_stages or self.num_sparse_stages - self.output_features["dense_stage"] = PolicyFeature( - shape=(self.num_frames, dense_stages), type=FeatureType.REWARD - ) - self.output_features["dense_progress"] = PolicyFeature( - shape=(self.num_frames, 1), type=FeatureType.REWARD - ) - else: - self.output_features["sparse_stage"] = PolicyFeature( - shape=(self.num_frames, self.num_sparse_stages), type=FeatureType.REWARD - ) - self.output_features["sparse_progress"] = PolicyFeature( - shape=(self.num_frames, 1), type=FeatureType.REWARD - ) - - if self.max_rewind_steps >= self.n_obs_steps: - raise ValueError( - f"max_rewind_steps ({self.max_rewind_steps}) must be less than n_obs_steps ({self.n_obs_steps})" - ) - if self.num_sparse_stages < 1: - raise ValueError(f"num_sparse_stages must be at least 1, got {self.num_sparse_stages}") - if ( - self.annotation_mode in ["dense_only", "dual"] - and self.num_dense_stages is not None - and self.num_dense_stages < 2 - ): - raise ValueError(f"num_dense_stages must be at least 2, got {self.num_dense_stages}") - - def get_optimizer_preset(self) -> AdamWConfig: - """Get default optimizer configuration for SARM training.""" - return AdamWConfig( - lr=5e-5, - weight_decay=1e-3, - betas=(0.9, 0.999), - eps=1e-8, - ) - - def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig: - """Get default learning rate scheduler configuration.""" - return CosineDecayWithWarmupSchedulerConfig( - peak_lr=5e-5, - decay_lr=5e-6, - num_warmup_steps=500, - num_decay_steps=50000, - ) - - def validate_features(self) -> None: - pass - - @property - def uses_dual_heads(self) -> bool: - """Whether the model uses dual heads (dense_only or dual annotation modes).""" - return self.annotation_mode in ["dense_only", "dual"] - - @property - def num_frames(self) -> int: - """Total number of frames in sequence. - - For training: 1 + n_obs_steps + max_rewind_steps - The sequence is: [obs_frames (n_obs_steps + 1)] + [rewind_frames (max_rewind_steps)] - """ - return 1 + self.n_obs_steps + self.max_rewind_steps - - @property - def max_length(self) -> int: - return self.num_frames - - @property - def observation_delta_indices(self) -> list[int]: - """Bidirectional frame sampling centered on target frame. - - Example with n_obs_steps=8, gap=30: - Before: [-120, -90, -60, -30] (4 frames) - Current: [0] (1 frame) - After: [30, 60, 90, 120] (4 frames) - Total: 9 frames - """ - half_steps = self.n_obs_steps // 2 - - past_deltas = [-self.frame_gap * i for i in range(half_steps, 0, -1)] - future_deltas = [self.frame_gap * i for i in range(1, half_steps + 1)] - obs_deltas = past_deltas + [0] + future_deltas - - # Rewind placeholders - rewind_deltas = [-self.frame_gap * (i + 1) for i in range(self.max_rewind_steps)] - - return obs_deltas + rewind_deltas - - @property - def action_delta_indices(self) -> None: - """SARM is a reward model, not an action policy.""" - return None - - @property - def reward_delta_indices(self) -> None: - return None diff --git a/src/lerobot/policies/sarm/modeling_sarm.py b/src/lerobot/policies/sarm/modeling_sarm.py deleted file mode 100644 index 6051d90f8..000000000 --- a/src/lerobot/policies/sarm/modeling_sarm.py +++ /dev/null @@ -1,794 +0,0 @@ -#!/usr/bin/env python - -# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu -# and 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. - -""" -SARM: Stage-Aware Reward Modeling for Long Horizon Robot Manipulation. - -Paper: https://arxiv.org/abs/2509.25358 - -- StageTransformer: Predicts stage classification (sparse/dense) -- SubtaskTransformer: Predicts within-stage progress (tau) conditioned on stage -""" - -import json -import logging -import random - -import numpy as np -import torch -import torch.nn as nn -import torch.nn.functional as F # noqa: N812 -from torch import Tensor - -from lerobot.policies.pretrained import PreTrainedPolicy -from lerobot.policies.sarm.configuration_sarm import SARMConfig -from lerobot.policies.sarm.sarm_utils import ( - normalize_stage_tau, - pad_state_to_max_dim, -) -from lerobot.utils.constants import OBS_STR - - -class StageTransformer(nn.Module): - """ - Stage classification transformer for SARM. - - Predicts which stage/subtask the current frame belongs to. - Supports both sparse (high-level) and dense (fine-grained) annotation schemes. - - Input streams: [vis_proj, lang_proj, state_proj] concatenated -> (B, N+2, T, D) - Output: stage logits (B, T, num_classes) - """ - - def __init__( - self, - d_model: int = 512, - vis_emb_dim: int = 512, - text_emb_dim: int = 512, - state_dim: int = 32, - n_layers: int = 6, - n_heads: int = 8, - dropout: float = 0.1, - num_cameras: int = 1, - num_classes_sparse: int = 4, - num_classes_dense: int = 8, - ): - super().__init__() - self.d_model = d_model - self.num_cameras = num_cameras - - # Projections - self.lang_proj = nn.Linear(text_emb_dim, d_model) - self.visual_proj = nn.Linear(vis_emb_dim, d_model) - self.state_proj = nn.Linear(state_dim, d_model) - - # Encoder - enc_layer = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, batch_first=True) - self.transformer = nn.TransformerEncoder(enc_layer, n_layers) - - # Positional bias on first visual frame - self.first_pos = nn.Parameter(torch.zeros(1, d_model)) - - # Shared fusion MLP - # Fuses (num_cameras + 2) streams: cameras + lang + state - fused_in = d_model * (num_cameras + 2) - self.fusion_backbone = nn.Sequential( - nn.LayerNorm(fused_in), - nn.Linear(fused_in, d_model), - nn.ReLU(), - ) - - # Scheme-specific heads - self.heads = nn.ModuleDict( - { - "sparse": nn.Linear(d_model, num_classes_sparse), - "dense": nn.Linear(d_model, num_classes_dense), - } - ) - - def _prep_lang(self, lang_emb: torch.Tensor, B: int, T: int, D: int) -> torch.Tensor: # noqa: N803 - """ - Prepare language embeddings for fusion. - - Accepts lang_emb of shape: - - (B, text_emb_dim) -> broadcast across time - - (B, T, text_emb_dim) -> per-timestep (dense annotation mode) - - Returns: (B, 1, T, D) - """ - if lang_emb.dim() == 3: - # (B, T, E) -> (B, T, D) -> (B, 1, T, D) - lang_proj = self.lang_proj(lang_emb).unsqueeze(1) - else: - # (B, E) -> (B, 1, 1, D) -> expand to (B, 1, T, D) - lang_proj = self.lang_proj(lang_emb).unsqueeze(1).unsqueeze(2).expand(B, 1, T, D) - return lang_proj - - def forward( - self, - img_seq: torch.Tensor, # (B, N, T, vis_emb_dim) - lang_emb: torch.Tensor, # (B, E) or (B, T, E) - state: torch.Tensor, # (B, T, state_dim) - lengths: torch.Tensor, # (B,) - valid sequence lengths - scheme: str = "sparse", # "sparse" or "dense" - ) -> torch.Tensor: - """ - Forward pass for stage classification. - - Args: - img_seq: Image embeddings (B, N, T, vis_emb_dim) where N=num_cameras - lang_emb: Language embeddings (B, E) or (B, T, E) for dense - state: State features (B, T, state_dim) - lengths: Valid sequence lengths (B,) for masking - scheme: "sparse" or "dense" for head selection - - Returns: - Stage logits (B, T, num_classes) - """ - assert scheme in self.heads, f"Unknown scheme '{scheme}'. Use one of {list(self.heads.keys())}." - - B, N, T, _ = img_seq.shape # noqa: N806 - D = self.d_model # noqa: N806 - device = img_seq.device - - # Project inputs - vis_proj = self.visual_proj(img_seq) # (B, N, T, D) - state_proj = self.state_proj(state).unsqueeze(1) # (B, 1, T, D) - lang_proj = self._prep_lang(lang_emb, B, T, D) # (B, 1, T, D) - - # Concatenate streams - # cameras + lang + state -> (B, N+2, T, D) - x = torch.cat([vis_proj, lang_proj, state_proj], dim=1) - - # Add positional bias to first visual frame - x[:, :N, 0, :] = x[:, :N, 0, :] + self.first_pos - - # Flatten to tokens for Transformer - x_tokens = x.view(B, (N + 2) * T, D) - L = x_tokens.size(1) # noqa: N806 - - # Create padding mask - base_mask = torch.arange(T, device=device).expand(B, T) >= lengths.unsqueeze(1) # (B, T) - mask = base_mask.unsqueeze(1).expand(B, N + 2, T).reshape(B, (N + 2) * T) - - # Create causal mask - causal_mask = torch.triu(torch.ones(L, L, device=device, dtype=torch.bool), diagonal=1) - - # Encode - h = self.transformer(x_tokens, mask=causal_mask, src_key_padding_mask=mask, is_causal=True) - - # Reshape and fuse - h = h.view(B, N + 2, T, D).permute(0, 2, 1, 3).reshape(B, T, (N + 2) * D) - fused = self.fusion_backbone(h) # (B, T, D) - - # Scheme-specific logits - logits = self.heads[scheme](fused) # (B, T, num_classes) - return logits - - -class SubtaskTransformer(nn.Module): - """ - Subtask progress regression transformer for SARM. - - Predicts within-stage normalized progress (tau) conditioned on stage prior. - The stage prior is a one-hot encoding passed from StageTransformer predictions. - - Input streams: [vis_proj, lang_proj, state_proj, stage_emb] -> (B, N+3, T, D) - Output: tau predictions (B, T) in [0, 1] - """ - - def __init__( - self, - d_model: int = 512, - vis_emb_dim: int = 512, - text_emb_dim: int = 512, - state_dim: int = 32, - n_layers: int = 6, - n_heads: int = 8, - dropout: float = 0.1, - num_cameras: int = 1, - ): - super().__init__() - self.d_model = d_model - self.num_cameras = num_cameras - - # Projections - self.lang_proj = nn.Linear(text_emb_dim, d_model) - self.visual_proj = nn.Linear(vis_emb_dim, d_model) - self.state_proj = nn.Linear(state_dim, d_model) - - # Encoder - enc = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, batch_first=True) - self.transformer = nn.TransformerEncoder(enc, n_layers) - - # Learned bias on first visual frame - self.first_pos = nn.Parameter(torch.zeros(1, d_model)) - - # Shared fusion backbone - # Fuses (num_cameras + 3) streams: cameras + lang + state + stage_emb - fused_in = d_model * (num_cameras + 3) - self.fusion_backbone = nn.Sequential( - nn.LayerNorm(fused_in), - nn.Linear(fused_in, d_model), - nn.ReLU(), - ) - - # Scheme-specific regression heads - self.heads = nn.ModuleDict( - { - "sparse": nn.Linear(d_model, 1), - "dense": nn.Linear(d_model, 1), - } - ) - - def _prep_lang(self, lang_emb: torch.Tensor, B: int, T: int, D: int) -> torch.Tensor: # noqa: N803 - """ - Prepare language embeddings for fusion. - """ - if lang_emb.dim() == 3: - # (B, T, E) -> (B, T, D) -> (B, 1, T, D) - return self.lang_proj(lang_emb).unsqueeze(1) - else: - # (B, E) -> (B, 1, 1, D) -> (B, 1, T, D) - return self.lang_proj(lang_emb).unsqueeze(1).unsqueeze(2).expand(B, 1, T, D) - - def _stage_to_dmodel(self, stage_prior: torch.Tensor) -> torch.Tensor: - """ - Deterministic projection of one-hot stage to d_model by pad/truncate. - - Args: - stage_prior: One-hot stage embedding (B, 1, T, C) - - Returns: - Projected stage embedding (B, 1, T, d_model) - """ - B, one, T, C = stage_prior.shape # noqa: N806 - D = self.d_model # noqa: N806 - if D == C: - return stage_prior - elif D > C: - pad = torch.zeros(B, one, T, D - C, device=stage_prior.device, dtype=stage_prior.dtype) - return torch.cat([stage_prior, pad], dim=-1) - else: - return stage_prior[..., :D] - - def forward( - self, - img_seq: torch.Tensor, # (B, N, T, vis_emb_dim) - lang_emb: torch.Tensor, # (B, E) or (B, T, E) - state: torch.Tensor, # (B, T, state_dim) - lengths: torch.Tensor, # (B,) - valid sequence lengths - stage_prior: torch.Tensor, # (B, 1, T, C) one-hot from gen_stage_emb - scheme: str = "sparse", # "sparse" or "dense" - ) -> torch.Tensor: - """ - Forward pass for subtask progress regression. - - Args: - img_seq: Image embeddings (B, N, T, vis_emb_dim) - lang_emb: Language embeddings (B, E) or (B, T, E) - state: State features (B, T, state_dim) - lengths: Valid sequence lengths (B,) for masking - stage_prior: One-hot stage prior (B, 1, T, num_classes) - scheme: "sparse" or "dense" for head selection - - Returns: - Tau predictions (B, T) in [0, 1] via sigmoid - """ - assert scheme in self.heads, f"Unknown scheme '{scheme}'. Use one of {list(self.heads.keys())}." - - B, N, T, _ = img_seq.shape # noqa: N806 - D = self.d_model # noqa: N806 - device = img_seq.device - - # Project inputs - vis_proj = self.visual_proj(img_seq) # (B, N, T, D) - state_proj = self.state_proj(state).unsqueeze(1) # (B, 1, T, D) - lang_proj = self._prep_lang(lang_emb, B, T, D) # (B, 1, T, D) - stage_emb = self._stage_to_dmodel(stage_prior) # (B, 1, T, D) - - # Concatenate all streams - # cameras + lang + state + stage_emb -> (B, N+3, T, D) - x = torch.cat([vis_proj, lang_proj, state_proj, stage_emb], dim=1) - - # Add positional bias to first visual frame - x[:, :N, 0, :] = x[:, :N, 0, :] + self.first_pos - - # Flatten to tokens - x_tokens = x.view(B, (N + 3) * T, D) - L = x_tokens.size(1) # noqa: N806 - - # Create padding mask - base_mask = torch.arange(T, device=device).expand(B, T) >= lengths.unsqueeze(1) - mask = base_mask.unsqueeze(1).expand(B, N + 3, T).reshape(B, (N + 3) * T) - - # Create causal mask - causal_mask = torch.triu(torch.ones(L, L, device=device, dtype=torch.bool), diagonal=1) - - # Encode - h = self.transformer(x_tokens, mask=causal_mask, src_key_padding_mask=mask, is_causal=True) - - # Reshape and fuse - h = h.view(B, N + 3, T, D) - h_flat = h.permute(0, 2, 1, 3).reshape(B, T, (N + 3) * D) - fused = self.fusion_backbone(h_flat) # (B, T, D) - - # Scheme-specific regression head -> sigmoid - r = torch.sigmoid(self.heads[scheme](fused)).squeeze(-1) # (B, T) - return r - - -def gen_stage_emb(num_classes: int, targets: torch.Tensor) -> torch.Tensor: - """ - Generate one-hot stage embeddings from targets. - - Args: - num_classes: Number of stage classes - targets: Target values (B, T) where integer part is stage index - - Returns: - One-hot stage embedding (B, 1, T, num_classes) - """ - # Integer part of float targets -> [0, C-1] - idx = targets.long().clamp(min=0, max=num_classes - 1) # (B, T) - C = num_classes # noqa: N806 - # Identity-lookup one-hot - stage_onehot = torch.eye(C, device=targets.device)[idx] # (B, T, C) - stage_onehot = stage_onehot.unsqueeze(1) # (B, 1, T, C) - return stage_onehot - - -class SARMRewardModel(PreTrainedPolicy): - """ - SARM Reward Model for stage-aware task completion rewards. - - Uses two separate transformer models: - - StageTransformer: Classifies which stage/subtask - - SubtaskTransformer: Predicts within-stage progress (tau) - - Training uses 75%/25% GT/predicted stage conditioning (teacher forcing). - """ - - name = "sarm" - config_class = SARMConfig - - def __init__(self, config: SARMConfig, dataset_stats: dict | None = None, dataset_meta=None): - super().__init__(config, dataset_stats) - config.validate_features() - self.config = config - self.dataset_stats = dataset_stats - self.device = torch.device( - config.device if config.device else "cuda" if torch.cuda.is_available() else "cpu" - ) - - # Load temporal proportions based on annotation_mode - if config.annotation_mode == "single_stage": - logging.info(f"Using single_stage mode: sparse_subtask_names={config.sparse_subtask_names}") - elif dataset_meta is not None: - self._load_temporal_proportions(dataset_meta) - - # Create two separate models - self.stage_model = StageTransformer( - d_model=config.hidden_dim, - vis_emb_dim=config.image_dim, - text_emb_dim=config.text_dim, - state_dim=config.max_state_dim, - n_layers=config.num_layers, - n_heads=config.num_heads, - dropout=config.dropout, - num_cameras=1, # Single camera for now - num_classes_sparse=config.num_sparse_stages, - num_classes_dense=config.num_dense_stages or config.num_sparse_stages, - ) - - self.subtask_model = SubtaskTransformer( - d_model=config.hidden_dim, - vis_emb_dim=config.image_dim, - text_emb_dim=config.text_dim, - state_dim=config.max_state_dim, - n_layers=config.num_layers, - n_heads=config.num_heads, - dropout=config.dropout, - num_cameras=1, - ) - - self.stage_model.to(self.device) - self.subtask_model.to(self.device) - - # GT/predicted stage ratio for teacher forcing - self.gt_stage_ratio = 0.75 - - if config.uses_dual_heads: - logging.info( - f"SARM initialized with dual heads: {config.num_sparse_stages} sparse stages, " - f"{config.num_dense_stages} dense stages" - ) - else: - logging.info(f"SARM initialized with sparse head only: {config.num_sparse_stages} stages") - - logging.info(f"SARM initialized on {self.device}") - - def _load_proportions_from_json(self, path, annotation_type: str) -> tuple[list[str], list[float]]: - """Load temporal proportions from a JSON file (preserving order).""" - if not path.exists(): - raise ValueError( - f"{annotation_type.capitalize()} temporal proportions not found at {path}. " - f"Run the subtask annotation tool with --{annotation_type}-subtasks to generate annotations." - ) - with open(path) as f: - proportions_dict = json.load(f) - names = list(proportions_dict.keys()) - logging.info(f"Loaded {len(names)} {annotation_type} subtasks: {names}") - logging.info(f"{annotation_type.capitalize()} temporal proportions: {proportions_dict}") - return names, [proportions_dict[name] for name in names] - - def _load_temporal_proportions(self, dataset_meta) -> None: - """Load temporal proportions based on annotation_mode.""" - meta_path = dataset_meta.root / "meta" - - if self.config.annotation_mode == "dual": - names, props = self._load_proportions_from_json( - meta_path / "temporal_proportions_sparse.json", "sparse" - ) - ( - self.config.num_sparse_stages, - self.config.sparse_subtask_names, - self.config.sparse_temporal_proportions, - ) = len(names), names, props - - if self.config.annotation_mode in ["dense_only", "dual"]: - names, props = self._load_proportions_from_json( - meta_path / "temporal_proportions_dense.json", "dense" - ) - ( - self.config.num_dense_stages, - self.config.dense_subtask_names, - self.config.dense_temporal_proportions, - ) = len(names), names, props - if self.config.annotation_mode == "dense_only": - logging.info(f"Using auto-generated sparse 'task' stage: {self.config.sparse_subtask_names}") - - def to(self, device): - """Override to method to ensure all components move together.""" - super().to(device) - self.device = device if isinstance(device, torch.device) else torch.device(device) - self.stage_model.to(device) - self.subtask_model.to(device) - return self - - @torch.no_grad() - def calculate_rewards( - self, - text_embeddings: np.ndarray | torch.Tensor, - video_embeddings: np.ndarray | torch.Tensor, - state_features: np.ndarray | torch.Tensor | None = None, - lengths: np.ndarray | torch.Tensor | None = None, - return_all_frames: bool = False, - return_stages: bool = False, - return_confidence: bool = False, - head_mode: str | None = "sparse", - frame_index: int | None = None, - ) -> np.ndarray | tuple: - """ - Calculate rewards for given text, video, and state representations. - - This is the canonical method for SARM reward computation, used for: - - Inference/visualization - - RA-BC weight computation - - Args: - text_embeddings: Encoded text representations (batch_size, 512) - video_embeddings: Encoded video representations (batch_size, num_frames, 512) - state_features: Joint state features (batch_size, num_frames, state_dim) - lengths: Valid sequence lengths (batch_size,) - return_all_frames: If True, return rewards for all frames - return_stages: If True, also return stage predictions - return_confidence: If True, also return stage confidence - head_mode: Which head to use ("sparse" or "dense") - frame_index: Index of the target frame to extract (default: n_obs_steps). - - Returns: - Rewards and optionally stage probs/confidence. - """ - if isinstance(text_embeddings, np.ndarray): - text_embeddings = torch.tensor(text_embeddings, dtype=torch.float32) - if isinstance(video_embeddings, np.ndarray): - video_embeddings = torch.tensor(video_embeddings, dtype=torch.float32) - if state_features is not None and isinstance(state_features, np.ndarray): - state_features = torch.tensor(state_features, dtype=torch.float32) - - # Handle single sample case - if text_embeddings.dim() == 1: - text_embeddings = text_embeddings.unsqueeze(0) - video_embeddings = video_embeddings.unsqueeze(0) - if state_features is not None: - state_features = state_features.unsqueeze(0) - single_sample = True - else: - single_sample = False - - batch_size = video_embeddings.shape[0] - seq_len = video_embeddings.shape[1] - - scheme = head_mode - - # Default lengths if not provided - if lengths is None: - lengths = torch.full((batch_size,), seq_len, dtype=torch.int32) - elif isinstance(lengths, np.ndarray): - lengths = torch.tensor(lengths, dtype=torch.int32) - - # Reshape video to (B, N, T, D) for multi-camera format - # Currently single camera: (B, T, D) -> (B, 1, T, D) - img_seq = video_embeddings.unsqueeze(1).to(self.device) - lang_emb = text_embeddings.to(self.device) - state = ( - state_features.to(self.device) - if state_features is not None - else torch.zeros(batch_size, seq_len, self.config.max_state_dim, device=self.device) - ) - lens = lengths.to(self.device) - - # Pad state to max_state_dim - state = pad_state_to_max_dim(state, self.config.max_state_dim) - - # Get num_classes for this scheme - num_classes = self.config.num_sparse_stages if scheme == "sparse" else self.config.num_dense_stages - - # Run stage model - stage_logits = self.stage_model(img_seq, lang_emb, state, lens, scheme=scheme) - stage_probs = F.softmax(stage_logits, dim=-1) # (B, T, num_classes) - stage_idx = stage_probs.argmax(dim=-1) # (B, T) - stage_conf = stage_probs.gather(-1, stage_idx.unsqueeze(-1)).squeeze(-1) # (B, T) - - # Create one-hot stage prior - stage_onehot = F.one_hot(stage_idx, num_classes=num_classes).float() # (B, T, C) - stage_emb = stage_onehot.unsqueeze(1) # (B, 1, T, C) - - # Run subtask model - tau_pred = self.subtask_model(img_seq, lang_emb, state, lens, stage_emb, scheme=scheme) - - # Compute final reward: stage + tau - raw_reward = stage_idx.float() + tau_pred # (B, T) - - # Normalize to [0, 1] using temporal proportions for proper weighting - if scheme == "sparse": - normalized_reward = normalize_stage_tau( - raw_reward, - num_stages=num_classes, - temporal_proportions=self.config.sparse_temporal_proportions, - subtask_names=self.config.sparse_subtask_names, - ) - else: - normalized_reward = normalize_stage_tau( - raw_reward, - num_stages=num_classes, - temporal_proportions=self.config.dense_temporal_proportions, - subtask_names=self.config.dense_subtask_names, - ) - - # Default frame index is n_obs_steps (last observation frame) - if frame_index is None: - frame_index = self.config.n_obs_steps - - # Prepare outputs (batch mode or no smoothing) - if return_all_frames: - rewards = normalized_reward.cpu().numpy() - else: - rewards = normalized_reward[:, frame_index].cpu().numpy() - - if single_sample: - rewards = rewards[0] if not return_all_frames else rewards[0] - - outputs = [rewards] - if return_stages: - probs = stage_probs.cpu().numpy() - if single_sample: - probs = probs[0] - outputs.append(probs) - if return_confidence: - conf = stage_conf.cpu().numpy() - if single_sample: - conf = conf[0] - outputs.append(conf) - - return outputs[0] if len(outputs) == 1 else tuple(outputs) - - def train(self, mode: bool = True): - """Set training mode for both models.""" - super().train(mode) - self.stage_model.train(mode) - self.subtask_model.train(mode) - return self - - def eval(self): - """Set evaluation mode for both models.""" - return self.train(False) - - def parameters(self): - """Override to return trainable parameters from both models.""" - from itertools import chain - - return chain(self.stage_model.parameters(), self.subtask_model.parameters()) - - def get_optim_params(self): - """Override to return optimizer parameters from both models.""" - return self.parameters() - - def reset(self): - """Required by PreTrainedPolicy but not used for reward models.""" - pass - - def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor: - """Required by PreTrainedPolicy but not used for reward models.""" - raise NotImplementedError("SARM model does not predict action chunks") - - def select_action(self, batch: dict[str, Tensor]) -> Tensor: - """Required by PreTrainedPolicy but not used for SARM.""" - raise NotImplementedError("SARM model does not select actions") - - def _train_step( - self, - img_emb: torch.Tensor, # (B, N, T, D) - lang_emb: torch.Tensor, # (B, E) or (B, T, E) - state: torch.Tensor, # (B, T, state_dim) - lengths: torch.Tensor, # (B,) - targets: torch.Tensor, # (B, T) - format: stage.tau - scheme: str, - ) -> dict[str, torch.Tensor]: - """ - Single training step for one annotation scheme. - - Implements 75%/25% GT/predicted stage conditioning. - - Args: - img_emb: Image embeddings (B, N, T, D) - lang_emb: Language embeddings - state: State features - lengths: Valid sequence lengths - targets: Target values where floor=stage, remainder=tau - scheme: "sparse" or "dense" - - Returns: - Dict with stage_loss, subtask_loss, total_loss - """ - num_classes = self.config.num_sparse_stages if scheme == "sparse" else self.config.num_dense_stages - - # Ground truth: stage (integer) and tau (fractional) - # Clamp stage indices to valid range [0, num_classes-1] to handle edge cases - # where targets may exceed expected range (e.g., frames between subtasks) - gt_stage = torch.floor(targets).long().clamp(0, num_classes - 1) # (B, T) - gt_tau = torch.remainder(targets, 1.0) # (B, T) - - # Run stage model - stage_pred = self.stage_model(img_emb, lang_emb, state, lengths, scheme=scheme) - - # 75%/25% GT/predicted stage conditioning - if random.random() < self.gt_stage_ratio: - # Mode 1: Use ground truth stage -> one-hot - stage_emb = gen_stage_emb(num_classes, targets) # (B, 1, T, C) - else: - # Mode 2: Use predicted stage argmax -> one-hot - stage_idx = stage_pred.argmax(dim=-1) # (B, T) - stage_onehot = F.one_hot(stage_idx, num_classes=num_classes).float() # (B, T, C) - stage_emb = stage_onehot.unsqueeze(1) # (B, 1, T, C) - - # Run subtask model with stage prior - tau_pred = self.subtask_model(img_emb, lang_emb, state, lengths, stage_emb, scheme=scheme) - - # Compute losses - stage_loss = F.cross_entropy(stage_pred.view(-1, num_classes), gt_stage.view(-1), reduction="mean") - subtask_loss = F.mse_loss(tau_pred, gt_tau, reduction="mean") - - return { - "stage_loss": stage_loss, - "subtask_loss": subtask_loss, - "total_loss": stage_loss + subtask_loss, - } - - def forward(self, batch): - """ - Forward pass for SARM reward model training. - - Uses stage+tau target format where: - - Integer part = stage index - - Fractional part = within-stage progress (tau) - - Training uses 75%/25% GT/predicted stage conditioning. - - Args: - batch: Dictionary with 'observation' containing: - - 'video_features': (B, T, 512) pre-encoded video features - - 'text_features': (B, 512) or (B, T, 512) text features - - 'state_features': (B, T, state_dim) joint state features - - 'lengths': (B,) valid sequence lengths - - 'sparse_targets': (B, T) sparse targets (stage.tau format) - - 'dense_targets': (B, T) dense targets (optional, for dual mode) - - Returns: - Tuple of (total_loss, output_dict with loss components) - """ - observation = batch.get(OBS_STR, batch) - - # Extract features - video_features = observation["video_features"].to(self.device) - text_features = observation["text_features"].to(self.device) - state_features = observation.get("state_features") - if state_features is not None: - state_features = state_features.to(self.device) - - batch_size = video_features.shape[0] - seq_len = video_features.shape[1] - - # Get lengths (default to full sequence) - lengths = observation.get("lengths") - if lengths is None: - lengths = torch.full((batch_size,), seq_len, dtype=torch.int32, device=self.device) - else: - lengths = lengths.to(self.device) - - # Reshape video to (B, N, T, D) - single camera - img_emb = video_features.unsqueeze(1) - - # Pad state to max_state_dim - if state_features is None: - state_features = torch.zeros(batch_size, seq_len, self.config.max_state_dim, device=self.device) - else: - state_features = pad_state_to_max_dim(state_features, self.config.max_state_dim) - - output_dict = {} - total_loss = torch.tensor(0.0, device=self.device) - - # Sparse training (always) - sparse_targets = observation.get("sparse_targets") - if sparse_targets is None: - # Try legacy format - sparse_targets = observation.get("targets") - if sparse_targets is None: - raise ValueError("sparse_targets (or targets) is required for SARM training") - sparse_targets = sparse_targets.to(self.device) - - sparse_result = self._train_step( - img_emb, text_features, state_features, lengths, sparse_targets, scheme="sparse" - ) - output_dict["sparse_stage_loss"] = sparse_result["stage_loss"].item() - output_dict["sparse_subtask_loss"] = sparse_result["subtask_loss"].item() - total_loss = total_loss + sparse_result["total_loss"] - - # Dense training (if dual mode) - if self.config.uses_dual_heads: - dense_targets = observation.get("dense_targets") - if dense_targets is not None: - dense_targets = dense_targets.to(self.device) - dense_result = self._train_step( - img_emb, text_features, state_features, lengths, dense_targets, scheme="dense" - ) - output_dict["dense_stage_loss"] = dense_result["stage_loss"].item() - output_dict["dense_subtask_loss"] = dense_result["subtask_loss"].item() - total_loss = total_loss + dense_result["total_loss"] - - output_dict["total_loss"] = total_loss.item() - return total_loss, output_dict - - -def compute_stage_loss(stage_logits: torch.Tensor, target_stages: torch.Tensor) -> torch.Tensor: - """Compute cross-entropy loss for stage classification.""" - _, _, num_stages = stage_logits.shape - stage_logits_flat = stage_logits.reshape(-1, num_stages) - # Clamp target stage indices to valid range [0, num_stages-1] - target_stages_flat = target_stages.reshape(-1).clamp(0, num_stages - 1) - return F.cross_entropy(stage_logits_flat, target_stages_flat) diff --git a/src/lerobot/policies/sarm/processor_sarm.py b/src/lerobot/policies/sarm/processor_sarm.py deleted file mode 100644 index 8f2bc23db..000000000 --- a/src/lerobot/policies/sarm/processor_sarm.py +++ /dev/null @@ -1,516 +0,0 @@ -#!/usr/bin/env python - -# Copyright 2025 The HuggingFace Inc. team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""SARM Processor for encoding images/text and generating stage+tau targets.""" - -import random -from typing import Any - -import numpy as np -import pandas as pd -import torch -from faker import Faker -from PIL import Image -from transformers import CLIPModel, CLIPProcessor - -from lerobot.configs.types import FeatureType, PolicyFeature -from lerobot.policies.sarm.configuration_sarm import SARMConfig -from lerobot.policies.sarm.sarm_utils import ( - apply_rewind_augmentation, - compute_absolute_indices, - find_stage_and_tau, - pad_state_to_max_dim, -) -from lerobot.processor import ( - AddBatchDimensionProcessorStep, - DeviceProcessorStep, - NormalizerProcessorStep, - PolicyAction, - PolicyProcessorPipeline, - ProcessorStep, - RenameObservationsProcessorStep, -) -from lerobot.processor.converters import ( - from_tensor_to_numpy, - policy_action_to_transition, - transition_to_policy_action, -) -from lerobot.processor.core import EnvTransition, TransitionKey -from lerobot.processor.pipeline import PipelineFeatureType -from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME - - -class SARMEncodingProcessorStep(ProcessorStep): - """ProcessorStep that encodes images and text with CLIP and generates stage and progress labels for SARM.""" - - def __init__( - self, - config: SARMConfig, - image_key: str | None = None, - dataset_meta=None, - dataset_stats: dict | None = None, - ): - super().__init__() - self.config = config - self.image_key = image_key or config.image_key - self.dataset_meta = dataset_meta - self.dataset_stats = dataset_stats - self.annotation_mode = config.annotation_mode - - # Helper to create temporal proportions dict - def make_props_dict(names, props): - return dict(zip(names, props, strict=True)) if names and props else None - - # Sparse annotations (always needed) - self.sparse_temporal_proportions = make_props_dict( - config.sparse_subtask_names, config.sparse_temporal_proportions - ) - self.sparse_subtask_names = config.sparse_subtask_names - - # Dense annotations (only for dual mode) - self.dense_subtask_names = config.dense_subtask_names if config.uses_dual_heads else None - self.dense_temporal_proportions = ( - make_props_dict(config.dense_subtask_names, config.dense_temporal_proportions) - if config.uses_dual_heads - else None - ) - - self.device = torch.device( - self.config.device if self.config.device else "cuda" if torch.cuda.is_available() else "cpu" - ) - - self.clip_model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32") - self.clip_processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32", use_fast=True) - self.clip_model.to(self.device) - self.clip_model.eval() - - self.verbs = ["move", "grasp", "rotate", "push", "pull", "slide", "lift", "place"] - self.fake = Faker() - - def _find_episode_for_frame(self, frame_idx: int) -> int: - """Find the episode index for a given frame index.""" - for ep_idx in range(len(self.dataset_meta.episodes)): - ep_start = self.dataset_meta.episodes[ep_idx]["dataset_from_index"] - ep_end = self.dataset_meta.episodes[ep_idx]["dataset_to_index"] - if ep_start <= frame_idx < ep_end: - return ep_idx - return 0 - - def _get_episode_indices(self, frame_indices: np.ndarray, episode_index) -> np.ndarray: - """Get episode indices for each frame index.""" - if episode_index is None: - return np.array([self._find_episode_for_frame(int(f)) for f in frame_indices]) - - episode_indices = np.atleast_1d(np.asarray(from_tensor_to_numpy(episode_index))) - - # If single episode but multiple frames, compute episode for each frame - if len(episode_indices) == 1 and len(frame_indices) > 1: - return np.array([self._find_episode_for_frame(int(f)) for f in frame_indices]) - - return episode_indices - - def _generate_perturbed_task(self) -> str: - """Generate a random perturbed task string for language perturbation.""" - num_words = random.randint(1, 5) - verb = random.choice(self.verbs) - phrase = " ".join([verb] + self.fake.words(nb=num_words)) - return phrase - - def _get_annotation_config(self, annotation_type: str) -> tuple[list[str], dict[str, float] | None]: - """Get global subtask names and temporal proportions for an annotation type.""" - if annotation_type == "dense": - return self.dense_subtask_names, self.dense_temporal_proportions - return self.sparse_subtask_names, self.sparse_temporal_proportions - - def _load_episode_annotations( - self, - ep_idx: int, - episodes_df: pd.DataFrame | None, - annotation_type: str, - global_names: list[str], - ) -> tuple[list | None, list | None, list | None]: - """Load subtask annotations for an episode from DataFrame.""" - # Single-stage mode: (linear progress 0→1) - if episodes_df is None or len(global_names) == 1: - return None, None, None - - # Resolve column name with fallback - def col(suffix): - prefixed = f"{annotation_type}_{suffix}" - return prefixed if prefixed in episodes_df.columns else suffix - - col_names = col("subtask_names") - if col_names not in episodes_df.columns or ep_idx >= len(episodes_df): - return None, None, None - - subtask_names = episodes_df.loc[ep_idx, col_names] - if subtask_names is None or (isinstance(subtask_names, float) and pd.isna(subtask_names)): - return None, None, None - - return ( - subtask_names, - episodes_df.loc[ep_idx, col("subtask_start_frames")], - episodes_df.loc[ep_idx, col("subtask_end_frames")], - ) - - def __call__(self, transition: EnvTransition) -> EnvTransition: - """ - Encode images, text, and normalize states in the transition. - - Implements SARM training data preparation: - - Applies language perturbation (20% probability) - - Applies rewind augmentation (80% probability) - - Generates stage+tau targets for all frames - - Outputs lengths tensor for valid sequence masking - """ - new_transition = transition.copy() if hasattr(transition, "copy") else dict(transition) - observation = new_transition.get(TransitionKey.OBSERVATION) - comp_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) - - frame_index = comp_data.get("index") - episode_index = comp_data.get("episode_index") - - if frame_index is None: - raise ValueError("Frame index ('index') not found in COMPLEMENTARY_DATA") - if episode_index is None: - raise ValueError("Episode index ('episode_index') not found in COMPLEMENTARY_DATA") - - frame_indices = np.atleast_1d(np.asarray(from_tensor_to_numpy(frame_index))) - episode_indices = self._get_episode_indices(frame_indices, episode_index) - - image = observation.get(self.image_key) - if isinstance(image, torch.Tensor): - image = image.cpu().numpy() - - # If 4D (T, C, H, W) from delta_timestamps, add batch dim - # If 3D (C, H, W) single frame, add batch and time dims - if image.ndim == 4: - image = image[np.newaxis, ...] # (T, C, H, W) -> (1, T, C, H, W) - elif image.ndim == 3: - image = image[np.newaxis, np.newaxis, ...] # (C, H, W) -> (1, 1, C, H, W) - - batch_size = image.shape[0] - total_frames = image.shape[1] # Should be 13: 9 obs + 4 rewind placeholders - n_obs_steps = self.config.n_obs_steps - max_rewind_steps = self.config.max_rewind_steps - n_obs_frames = 1 + n_obs_steps # 9 observation frames (including current) - - # Rewind augmentation - rewind_steps = torch.zeros(batch_size, dtype=torch.int32) - apply_rewind = self.training and random.random() < self.config.rewind_probability - - if apply_rewind and self.dataset_meta is not None: - for b_idx, (ep_idx, frame_idx) in enumerate( - zip(episode_indices.tolist(), frame_indices.tolist(), strict=True) - ): - ep_idx, frame_idx = int(ep_idx), int(frame_idx) - ep_start = self.dataset_meta.episodes[ep_idx]["dataset_from_index"] - - rewind_step, _ = apply_rewind_augmentation( - frame_idx, ep_start, n_obs_steps, max_rewind_steps, frame_gap=self.config.frame_gap - ) - rewind_steps[b_idx] = rewind_step - - # Compute valid lengths: n_obs_frames + rewind_steps - lengths = n_obs_frames + rewind_steps # (B,) - - # Apply rewind masking to images - # For frames beyond valid length, we mask with zeros (or copy last valid frame) - for b_idx in range(batch_size): - valid_len = lengths[b_idx].item() - if valid_len < total_frames: - image[b_idx, valid_len:] = 0 # Zero out frames beyond valid length - - # Encode images with CLIP - video_features = self._encode_images_batch(image) - observation["video_features"] = video_features - - state_key = self.config.state_key - state_data = observation.get(state_key) - - if isinstance(state_data, torch.Tensor): - state_tensor = state_data.float() - else: - state_tensor = torch.tensor(state_data, dtype=torch.float32) - - if state_tensor.ndim == 2: - state_tensor = state_tensor.unsqueeze(0) # (T, D) -> (1, T, D) - elif state_tensor.ndim == 1: - state_tensor = state_tensor.unsqueeze(0).unsqueeze(0) # (D,) -> (1, 1, D) - - # Apply same rewind masking to state - for b_idx in range(batch_size): - valid_len = lengths[b_idx].item() - if valid_len < state_tensor.shape[1]: - state_tensor[b_idx, valid_len:] = 0 # Zero out frames beyond valid length - - observation["state_features"] = pad_state_to_max_dim(state_tensor, self.config.max_state_dim) - - task = comp_data.get("task") - if isinstance(task, list): - task = task[0] if task else "" - - # Apply language perturbation during training (20% probability) - # When perturbed, targets will be zeroed to train model to output low values for irrelevant text - apply_perturbation = self.training and random.random() < self.config.language_perturbation_probability - if apply_perturbation: - task = self._generate_perturbed_task() - - # Encode text with CLIP - observation["text_features"] = self._encode_text_clip(task, batch_size) - - # Store lengths for model - observation["lengths"] = lengths - - # When language is perturbed, targets are zero so perturbed samples don't contribute to progress loss - if self.dataset_meta is not None: - episodes_df = self.dataset_meta.episodes.to_pandas() - - # Generate sparse targets - if self.sparse_temporal_proportions is not None: - if apply_perturbation: - # Zero targets when language is perturbed - sparse_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32) - else: - sparse_targets = self._compute_batch_targets( - frame_indices, episode_indices, lengths, rewind_steps, episodes_df, "sparse" - ) - observation["sparse_targets"] = sparse_targets - - # Generate dense targets (for dual mode) - if self.config.uses_dual_heads and self.dense_temporal_proportions is not None: - if apply_perturbation: - # Zero targets when language is perturbed - dense_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32) - else: - dense_targets = self._compute_batch_targets( - frame_indices, episode_indices, lengths, rewind_steps, episodes_df, "dense" - ) - observation["dense_targets"] = dense_targets - - new_transition[TransitionKey.OBSERVATION] = observation - return new_transition - - def _compute_batch_targets( - self, - frame_indices: np.ndarray, - episode_indices: np.ndarray, - lengths: torch.Tensor, - rewind_steps: torch.Tensor, - episodes_df: pd.DataFrame | None, - annotation_type: str, - ) -> torch.Tensor: - """Compute stage+tau targets for a batch of samples.""" - batch_size = len(frame_indices) - n_obs_steps = self.config.n_obs_steps - max_rewind_steps = self.config.max_rewind_steps - total_frames = 1 + n_obs_steps + max_rewind_steps - frame_gap = self.config.frame_gap - - global_names, temporal_props = self._get_annotation_config(annotation_type) - targets = torch.zeros(batch_size, total_frames, dtype=torch.float32) - - for b_idx in range(batch_size): - ep_idx = int(episode_indices[b_idx]) - frame_idx = int(frame_indices[b_idx]) - - ep_start = self.dataset_meta.episodes[ep_idx]["dataset_from_index"] - ep_end = self.dataset_meta.episodes[ep_idx]["dataset_to_index"] - ep_length = ep_end - ep_start - - subtask_names, subtask_start_frames, subtask_end_frames = self._load_episode_annotations( - ep_idx, episodes_df, annotation_type, global_names - ) - - # Compute observation frame indices - obs_indices, _ = compute_absolute_indices( - frame_idx, ep_start, ep_end, n_obs_steps, frame_gap=frame_gap - ) - obs_indices = obs_indices.tolist() - - # Compute targets for observation frames - for t_idx, abs_idx in enumerate(obs_indices): - rel_frame = abs_idx - ep_start - targets[b_idx, t_idx] = find_stage_and_tau( - rel_frame, - ep_length, - subtask_names, - subtask_start_frames, - subtask_end_frames, - global_names, - temporal_props, - return_combined=True, - ) - - # Compute targets for rewind frames (if any) - rewind_step = rewind_steps[b_idx].item() - if rewind_step > 0: - _, rewind_indices = apply_rewind_augmentation( - frame_idx, - ep_start, - n_obs_steps, - max_rewind_steps, - frame_gap=frame_gap, - rewind_step=rewind_step, - ) - - for r_idx, abs_idx in enumerate(rewind_indices[:rewind_step]): - rel_frame = max(0, abs_idx - ep_start) - targets[b_idx, n_obs_steps + 1 + r_idx] = find_stage_and_tau( - rel_frame, - ep_length, - subtask_names, - subtask_start_frames, - subtask_end_frames, - global_names, - temporal_props, - return_combined=True, - ) - - return targets - - @property - def training(self) -> bool: - return getattr(self, "_training_mode", True) - - def train(self, mode: bool = True): - """Set training mode for augmentation decisions.""" - self._training_mode = mode - return self - - def eval(self): - """Set evaluation mode (disable augmentations).""" - return self.train(False) - - @torch.no_grad() - def _encode_images_batch(self, images: np.ndarray) -> torch.Tensor: - """Encode a batch of images using CLIP. - - Args: - images: Batched images with shape: (B, T, C, H, W) - - Returns: - Encoded feature vectors with shape (B, T, 512) - """ - - batch_size, seq_length = images.shape[0], images.shape[1] - images = images.reshape(batch_size * seq_length, *images.shape[2:]) - - num_frames = images.shape[0] - images_list = [] - for i in range(num_frames): - img = images[i] - if img.shape[0] in [1, 3]: # Channel first (C, H, W) - img = img.transpose(1, 2, 0) - - # Handle single channel - if img.shape[-1] == 1: - img = np.repeat(img, 3, axis=-1) - - if img.dtype != np.uint8: - img = (img * 255).astype(np.uint8) if img.max() <= 1.0 else img.astype(np.uint8) - - images_list.append(Image.fromarray(img)) - - all_embeddings = [] - for i in range(0, num_frames, self.config.clip_batch_size): - batch_imgs = images_list[i : i + self.config.clip_batch_size] - - inputs = self.clip_processor(images=batch_imgs, return_tensors="pt") - inputs = {k: v.to(self.device) for k, v in inputs.items()} - - # Get image embeddings - embeddings = self.clip_model.get_image_features(**inputs).detach().cpu() - - # Handle single frame case - if embeddings.dim() == 1: - embeddings = embeddings.unsqueeze(0) - - all_embeddings.append(embeddings) - - all_embeddings = torch.cat(all_embeddings) # (B*T, 512) - all_embeddings = all_embeddings.reshape(batch_size, seq_length, -1) # (B, T, 512) - - return all_embeddings - - @torch.no_grad() - def _encode_text_clip(self, text: str, batch_size: int) -> torch.Tensor: - """Encode text using CLIP text encoder (per SARM paper A.4). - - Args: - text: Task description text to encode - batch_size: Batch size to replicate for - - Returns: - Encoded text features with shape (B, 512) - """ - inputs = self.clip_processor.tokenizer([text], return_tensors="pt", padding=True, truncation=True) - inputs = {k: v.to(self.device) for k, v in inputs.items()} - - text_embedding = self.clip_model.get_text_features(**inputs).detach().cpu() - text_embedding = text_embedding.expand(batch_size, -1) - - return text_embedding - - def transform_features( - self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] - ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: - """Add encoded features to the observation features.""" - features[PipelineFeatureType.OBSERVATION]["video_features"] = PolicyFeature( - type=FeatureType.VISUAL, shape=(self.config.num_frames, self.config.image_dim) - ) - features[PipelineFeatureType.OBSERVATION]["text_features"] = PolicyFeature( - type=FeatureType.LANGUAGE, shape=(self.config.text_dim,) - ) - features[PipelineFeatureType.OBSERVATION]["state_features"] = PolicyFeature( - type=FeatureType.STATE, shape=(self.config.num_frames, self.config.max_state_dim) - ) - return features - - -def make_sarm_pre_post_processors( - config: SARMConfig, - dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, - dataset_meta=None, -) -> tuple[ - PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], - PolicyProcessorPipeline[PolicyAction, PolicyAction], -]: - """Create pre-processor and post-processor pipelines for SARM.""" - return ( - PolicyProcessorPipeline[dict[str, Any], dict[str, Any]]( - steps=[ - AddBatchDimensionProcessorStep(), - RenameObservationsProcessorStep(rename_map={}), - NormalizerProcessorStep( - features={**config.input_features, **config.output_features}, - norm_map=config.normalization_mapping, - stats=dataset_stats, - ), - SARMEncodingProcessorStep( - config=config, dataset_meta=dataset_meta, dataset_stats=dataset_stats - ), - DeviceProcessorStep(device=config.device), - ], - name=POLICY_PREPROCESSOR_DEFAULT_NAME, - ), - PolicyProcessorPipeline[PolicyAction, PolicyAction]( - steps=[DeviceProcessorStep(device="cpu")], - name=POLICY_POSTPROCESSOR_DEFAULT_NAME, - to_transition=policy_action_to_transition, - to_output=transition_to_policy_action, - ), - ) diff --git a/src/lerobot/policies/sarm/sarm_utils.py b/src/lerobot/policies/sarm/sarm_utils.py deleted file mode 100644 index 5b6955d38..000000000 --- a/src/lerobot/policies/sarm/sarm_utils.py +++ /dev/null @@ -1,295 +0,0 @@ -#!/usr/bin/env python - -# Copyright 2025 The HuggingFace Inc. team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import random - -import numpy as np -import torch -import torch.nn.functional as F # noqa: N812 - - -def find_stage_and_tau( - current_frame: int, - episode_length: int, - subtask_names: list | None, - subtask_start_frames: list | None, - subtask_end_frames: list | None, - global_subtask_names: list, - temporal_proportions: dict, - return_combined: bool = False, -) -> tuple[int, float] | float: - """Find stage and within-stage progress (tau) for a frame. - - Args: - current_frame: Frame index relative to episode start - episode_length: Total frames in episode - subtask_names: Subtask names for this episode (None for single_stage) - subtask_start_frames: Subtask start frames - subtask_end_frames: Subtask end frames - global_subtask_names: Global list of all subtask names - temporal_proportions: Dict of temporal proportions - return_combined: If True, return stage+tau as float; else (stage_idx, tau) tuple - - Returns: - Float (stage.tau) if return_combined, else (stage_idx, tau) tuple - """ - stage_idx, tau = 0, 0.0 - num_stages = len(global_subtask_names) - - # Single-stage mode: linear progress from 0 to 1 - if num_stages == 1: - tau = min(1.0, max(0.0, current_frame / max(episode_length - 1, 1))) - elif subtask_names is None: - pass # stage_idx=0, tau=0.0 - elif current_frame < subtask_start_frames[0]: - pass # Before first subtask: stage_idx=0, tau=0.0 - elif current_frame > subtask_end_frames[-1]: - stage_idx, tau = num_stages - 1, 0.999 # After last subtask - else: - # Find which subtask this frame belongs to - found = False - for name, start, end in zip(subtask_names, subtask_start_frames, subtask_end_frames, strict=True): - if start <= current_frame <= end: - stage_idx = global_subtask_names.index(name) if name in global_subtask_names else 0 - tau = compute_tau(current_frame, start, end) - found = True - break - # Frame between subtasks - use previous subtask's end state - if not found: - for j in range(len(subtask_names) - 1): - if subtask_end_frames[j] < current_frame < subtask_start_frames[j + 1]: - name = subtask_names[j] - stage_idx = global_subtask_names.index(name) if name in global_subtask_names else j - tau = 1.0 - break - - if return_combined: - # Clamp to avoid overflow at end - if stage_idx >= num_stages - 1 and tau >= 1.0: - return num_stages - 1 + 0.999 - return stage_idx + tau - return stage_idx, tau - - -def compute_absolute_indices( - frame_idx: int, - ep_start: int, - ep_end: int, - n_obs_steps: int, - frame_gap: int = 30, -) -> tuple[torch.Tensor, torch.Tensor]: - """Compute absolute frame indices with clamping for bidirectional observation sequence. - - Bidirectional sampling centered on target frame: - - Before: [-frame_gap * half_steps, ..., -frame_gap] (half_steps frames) - - Current: [0] (1 frame) - - After: [frame_gap, ..., frame_gap * half_steps] (half_steps frames) - - Total: n_obs_steps + 1 frames - - Out-of-bounds frames are clamped (duplicated from boundary). - - Args: - frame_idx: Target frame index (center frame of sequence) - ep_start: Episode start index - ep_end: Episode end index (exclusive) - n_obs_steps: Number of observation steps (must be even for symmetric sampling) - frame_gap: Gap between observation frames - - Returns: - Tuple of (indices, out_of_bounds_flags) - """ - half_steps = n_obs_steps // 2 - - # Bidirectional deltas: past + current + future - past_deltas = [-frame_gap * i for i in range(half_steps, 0, -1)] - future_deltas = [frame_gap * i for i in range(1, half_steps + 1)] - delta_indices = past_deltas + [0] + future_deltas - - frames = [] - out_of_bounds = [] - - for delta in delta_indices: - target_idx = frame_idx + delta - # Clamp to episode bounds (duplicate boundary frames for out-of-bounds) - clamped_idx = max(ep_start, min(ep_end - 1, target_idx)) - frames.append(clamped_idx) - # Flag as out-of-bounds if clamping occurred - out_of_bounds.append(1 if target_idx != clamped_idx else 0) - - return torch.tensor(frames), torch.tensor(out_of_bounds) - - -def apply_rewind_augmentation( - frame_idx: int, - ep_start: int, - n_obs_steps: int, - max_rewind_steps: int, - frame_gap: int = 30, - rewind_step: int | None = None, -) -> tuple[int, list[int]]: - """ - Generate rewind frame indices for temporal augmentation. - - Rewind simulates going backwards through previously seen frames, - starting from before the earliest observation frame (for bidirectional sampling). - Appends reversed frames after the observation sequence. - - Args: - frame_idx: Target frame index (center of bidirectional observation window) - ep_start: Episode start index - n_obs_steps: Number of observation steps - max_rewind_steps: Maximum rewind steps - frame_gap: Gap between frames - rewind_step: If provided, use this exact rewind step (for deterministic behavior). - If None, sample randomly. - - Returns: - Tuple of (rewind_step, rewind_indices) - """ - # For bidirectional sampling, earliest obs frame is at frame_idx - half_steps * frame_gap - half_steps = n_obs_steps // 2 - earliest_obs_frame = frame_idx - half_steps * frame_gap - - # Required history: frames before earliest observation frame - if earliest_obs_frame <= ep_start: - return 0, [] # No history before observation window - - # Max valid rewind steps based on available history before earliest obs frame - available_history = earliest_obs_frame - ep_start - max_valid_step = available_history // frame_gap - max_rewind = min(max_rewind_steps, max(0, max_valid_step)) - - if max_rewind <= 0: - return 0, [] - - # Sample rewind steps if not provided - rewind_step = random.randint(1, max_rewind) if rewind_step is None else min(rewind_step, max_rewind) - - if rewind_step == 0: - return 0, [] - - # Generate rewind indices going backwards from earliest obs frame - # rewind_indices[0] is closest to obs window, rewind_indices[-1] is furthest back - rewind_indices = [] - for i in range(1, rewind_step + 1): - idx = earliest_obs_frame - i * frame_gap - idx = max(ep_start, idx) # Clamp to episode start - rewind_indices.append(idx) - - return rewind_step, rewind_indices - - -def compute_tau(current_frame: int | float, subtask_start: int | float, subtask_end: int | float) -> float: - """Compute τ_t = (t - s_k) / (e_k - s_k) ∈ [0, 1]. Returns 1.0 for zero-duration subtasks.""" - duration = subtask_end - subtask_start - if duration <= 0: - return 1.0 - return float(np.clip((current_frame - subtask_start) / duration, 0.0, 1.0)) - - -def pad_state_to_max_dim(state: torch.Tensor, max_state_dim: int) -> torch.Tensor: - """Pad the state tensor's last dimension to max_state_dim with zeros.""" - current_dim = state.shape[-1] - if current_dim >= max_state_dim: - return state[..., :max_state_dim] # Truncate if larger - - # Pad with zeros on the right - padding = (0, max_state_dim - current_dim) # (left, right) for last dim - return F.pad(state, padding, mode="constant", value=0) - - -def temporal_proportions_to_breakpoints( - temporal_proportions: dict[str, float] | list[float] | None, - subtask_names: list[str] | None = None, -) -> list[float] | None: - """Convert temporal proportions to cumulative breakpoints for normalization.""" - if temporal_proportions is None: - return None - - if isinstance(temporal_proportions, dict): - if subtask_names is not None: - proportions = [temporal_proportions.get(name, 0.0) for name in subtask_names] - else: - proportions = list(temporal_proportions.values()) - else: - proportions = list(temporal_proportions) - - total = sum(proportions) - if total > 0 and abs(total - 1.0) > 1e-6: - proportions = [p / total for p in proportions] - - breakpoints = [0.0] - cumsum = 0.0 - for prop in proportions: - cumsum += prop - breakpoints.append(cumsum) - breakpoints[-1] = 1.0 - - return breakpoints - - -def normalize_stage_tau( - x: float | torch.Tensor, - num_stages: int | None = None, - breakpoints: list[float] | None = None, - temporal_proportions: dict[str, float] | list[float] | None = None, - subtask_names: list[str] | None = None, -) -> float | torch.Tensor: - """ - Normalize stage+tau reward to [0, 1] with custom breakpoints. - - Maps stage index + within-stage tau to normalized progress [0, 1]. - The breakpoints are designed to give appropriate weight to each stage - based on their importance in the task (using temporal proportions). - - Priority: breakpoints > temporal_proportions > linear fallback - - Args: - x: Raw reward value (stage index + tau) where stage ∈ [0, num_stages-1] and tau ∈ [0, 1) - num_stages: Number of stages (required if breakpoints/proportions not provided) - breakpoints: Optional custom breakpoints list of length num_stages + 1. - temporal_proportions: Optional temporal proportions dict/list to compute breakpoints. - subtask_names: Optional ordered list of subtask names (for dict proportions) - - Returns: - Normalized progress value ∈ [0, 1] - """ - if breakpoints is not None: - num_stages = len(breakpoints) - 1 - elif temporal_proportions is not None: - breakpoints = temporal_proportions_to_breakpoints(temporal_proportions, subtask_names) - num_stages = len(breakpoints) - 1 - elif num_stages is not None: - breakpoints = [i / num_stages for i in range(num_stages + 1)] - else: - raise ValueError("Either num_stages, breakpoints, or temporal_proportions must be provided") - - if isinstance(x, torch.Tensor): - result = torch.zeros_like(x) - for i in range(num_stages): - mask = (x >= i) & (x < i + 1) - tau_in_stage = x - i - result[mask] = breakpoints[i] + tau_in_stage[mask] * (breakpoints[i + 1] - breakpoints[i]) - result[x >= num_stages] = 1.0 - return result.clamp(0.0, 1.0) - else: - if x < 0: - return 0.0 - if x >= num_stages: - return 1.0 - stage = int(x) - tau = x - stage - return breakpoints[stage] + tau * (breakpoints[stage + 1] - breakpoints[stage]) diff --git a/src/lerobot/processor/hil_processor.py b/src/lerobot/processor/hil_processor.py index 34eaeed51..324426f3a 100644 --- a/src/lerobot/processor/hil_processor.py +++ b/src/lerobot/processor/hil_processor.py @@ -556,7 +556,7 @@ class RewardClassifierProcessorStep(ProcessorStep): def __post_init__(self): """Initializes the reward classifier model after the dataclass is created.""" if self.pretrained_path is not None: - from lerobot.policies.sac.reward_model.modeling_classifier import Classifier + from lerobot.rewards.classifier.modeling_classifier import Classifier self.reward_classifier = Classifier.from_pretrained(self.pretrained_path) self.reward_classifier.to(self.device) diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index 1fed3bee4..885008312 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -36,6 +36,7 @@ from lerobot.envs.utils import close_envs from lerobot.optim.factory import make_optimizer_and_scheduler from lerobot.policies.factory import make_policy, make_pre_post_processors from lerobot.policies.pretrained import PreTrainedPolicy +from lerobot.rewards.factory import make_reward_pre_post_processors from lerobot.rl.wandb_utils import WandBLogger from lerobot.scripts.lerobot_eval import eval_policy_all from lerobot.utils.import_utils import register_third_party_plugins @@ -283,19 +284,30 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None): }, } - preprocessor, postprocessor = make_pre_post_processors( - policy_cfg=cfg.policy, - pretrained_path=cfg.policy.pretrained_path, - **processor_kwargs, - **postprocessor_kwargs, - ) + # TODO (kmeftah): Add `reward_model` field to TrainPipelineConfig so policy/reward training + # paths are separated at the config level instead of relying on isinstance checks. + # HACK: Reward model configs are currently passed via cfg.policy; detect and route accordingly. + from lerobot.configs.rewards import RewardModelConfig + + if isinstance(cfg.policy, RewardModelConfig): + preprocessor, postprocessor = make_reward_pre_post_processors( + cfg.policy, + **processor_kwargs, + ) + else: + preprocessor, postprocessor = make_pre_post_processors( + policy_cfg=cfg.policy, + pretrained_path=cfg.policy.pretrained_path, + **processor_kwargs, + **postprocessor_kwargs, + ) if is_main_process: logging.info("Creating optimizer and scheduler") optimizer, lr_scheduler = make_optimizer_and_scheduler(cfg, policy) # Load precomputed SARM progress for RA-BC if enabled - # Generate progress using: src/lerobot/policies/sarm/compute_rabc_weights.py + # Generate progress using: src/lerobot/rewards/sarm/compute_rabc_weights.py rabc_weights = None if cfg.use_rabc: from lerobot.utils.rabc import RABCWeights