diff --git a/src/lerobot/rewards/classifier/configuration_classifier.py b/src/lerobot/rewards/classifier/configuration_classifier.py index 398140022..6a8076f4d 100644 --- a/src/lerobot/rewards/classifier/configuration_classifier.py +++ b/src/lerobot/rewards/classifier/configuration_classifier.py @@ -32,6 +32,7 @@ class RewardClassifierConfig(RewardModelConfig): 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 diff --git a/src/lerobot/rewards/classifier/modeling_classifier.py b/src/lerobot/rewards/classifier/modeling_classifier.py index 49e8d1f40..9b9815dd6 100644 --- a/src/lerobot/rewards/classifier/modeling_classifier.py +++ b/src/lerobot/rewards/classifier/modeling_classifier.py @@ -97,10 +97,7 @@ class SpatialLearnedEmbeddings(nn.Module): class Classifier(PreTrainedRewardModel): - """Image classifier built on top of a pre-trained encoder. - - Binary success/failure classifier from images. Trainable via ``forward()``. - """ + """Image classifier built on top of a pre-trained encoder.""" name = "reward_classifier" config_class = RewardClassifierConfig @@ -108,7 +105,6 @@ class Classifier(PreTrainedRewardModel): def __init__( self, config: RewardClassifierConfig, - **kwargs, ): from transformers import AutoModel @@ -215,6 +211,7 @@ class Classifier(PreTrainedRewardModel): 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] @@ -279,6 +276,11 @@ class Classifier(PreTrainedRewardModel): 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: diff --git a/src/lerobot/rewards/sarm/configuration_sarm.py b/src/lerobot/rewards/sarm/configuration_sarm.py index 2a05b9b2b..7ef402cab 100644 --- a/src/lerobot/rewards/sarm/configuration_sarm.py +++ b/src/lerobot/rewards/sarm/configuration_sarm.py @@ -1,5 +1,3 @@ -#!/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. # @@ -55,6 +53,10 @@ class SARMConfig(RewardModelConfig): 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 @@ -66,7 +68,7 @@ class SARMConfig(RewardModelConfig): batch_size: int = 64 clip_batch_size: int = 64 dropout: float = 0.1 - stage_loss_weight: float = 1.0 + 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 @@ -82,7 +84,8 @@ class SARMConfig(RewardModelConfig): dense_temporal_proportions: list | None = None pretrained_model_path: str | None = None - image_key: str = OBS_IMAGES + ".top" + 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) @@ -114,6 +117,7 @@ class SARMConfig(RewardModelConfig): ) 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] @@ -201,7 +205,11 @@ class SARMConfig(RewardModelConfig): @property def num_frames(self) -> int: - """Total number of frames in sequence.""" + """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 @@ -210,7 +218,14 @@ class SARMConfig(RewardModelConfig): @property def observation_delta_indices(self) -> list[int]: - """Bidirectional frame sampling centered on target frame.""" + """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)] diff --git a/src/lerobot/rewards/sarm/modeling_sarm.py b/src/lerobot/rewards/sarm/modeling_sarm.py index f920607bc..3739d233a 100644 --- a/src/lerobot/rewards/sarm/modeling_sarm.py +++ b/src/lerobot/rewards/sarm/modeling_sarm.py @@ -1,5 +1,3 @@ -#!/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. # @@ -84,6 +82,7 @@ class StageTransformer(nn.Module): 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), @@ -100,48 +99,82 @@ class StageTransformer(nn.Module): ) def _prep_lang(self, lang_emb: torch.Tensor, B: int, T: int, D: int) -> torch.Tensor: # noqa: N803 - """Prepare language embeddings for fusion.""" + """ + 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, - lang_emb: torch.Tensor, - state: torch.Tensor, - lengths: torch.Tensor, - scheme: str = "sparse", + 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 - vis_proj = self.visual_proj(img_seq) - state_proj = self.state_proj(state).unsqueeze(1) - lang_proj = self._prep_lang(lang_emb, B, T, D) + # 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 - base_mask = torch.arange(T, device=device).expand(B, T) >= lengths.unsqueeze(1) + # 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) + fused = self.fusion_backbone(h) # (B, T, D) - logits = self.heads[scheme](fused) + # Scheme-specific logits + logits = self.heads[scheme](fused) # (B, T, num_classes) return logits @@ -150,6 +183,10 @@ 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__( @@ -167,15 +204,20 @@ class SubtaskTransformer(nn.Module): 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), @@ -183,6 +225,7 @@ class SubtaskTransformer(nn.Module): nn.ReLU(), ) + # Scheme-specific regression heads self.heads = nn.ModuleDict( { "sparse": nn.Linear(d_model, 1), @@ -191,12 +234,26 @@ class SubtaskTransformer(nn.Module): ) 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: @@ -209,51 +266,87 @@ class SubtaskTransformer(nn.Module): def forward( self, - img_seq: torch.Tensor, - lang_emb: torch.Tensor, - state: torch.Tensor, - lengths: torch.Tensor, - stage_prior: torch.Tensor, - scheme: str = "sparse", + 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 - vis_proj = self.visual_proj(img_seq) - state_proj = self.state_proj(state).unsqueeze(1) - lang_proj = self._prep_lang(lang_emb, B, T, D) - stage_emb = self._stage_to_dmodel(stage_prior) + # 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) + fused = self.fusion_backbone(h_flat) # (B, T, D) - r = torch.sigmoid(self.heads[scheme](fused)).squeeze(-1) + # 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.""" - idx = targets.long().clamp(min=0, max=num_classes - 1) + """ + 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 - stage_onehot = torch.eye(C, device=targets.device)[idx] - stage_onehot = stage_onehot.unsqueeze(1) + # 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 @@ -271,8 +364,8 @@ class SARMRewardModel(PreTrainedRewardModel): name = "sarm" config_class = SARMConfig - def __init__(self, config: SARMConfig, dataset_stats: dict | None = None, dataset_meta=None, **kwargs): - super().__init__(config) + 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 @@ -295,7 +388,7 @@ class SARMRewardModel(PreTrainedRewardModel): n_layers=config.num_layers, n_heads=config.num_heads, dropout=config.dropout, - num_cameras=1, + 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, ) @@ -314,6 +407,7 @@ class SARMRewardModel(PreTrainedRewardModel): 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: @@ -410,6 +504,20 @@ class SARMRewardModel(PreTrainedRewardModel): 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) @@ -418,6 +526,7 @@ class SARMRewardModel(PreTrainedRewardModel): 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) @@ -432,11 +541,14 @@ class SARMRewardModel(PreTrainedRewardModel): 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 = ( @@ -446,22 +558,29 @@ class SARMRewardModel(PreTrainedRewardModel): ) 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) - stage_idx = stage_probs.argmax(dim=-1) - stage_conf = stage_probs.gather(-1, stage_idx.unsqueeze(-1)).squeeze(-1) + 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) - stage_onehot = F.one_hot(stage_idx, num_classes=num_classes).float() - stage_emb = stage_onehot.unsqueeze(1) + # 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) - raw_reward = stage_idx.float() + tau_pred + # 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, @@ -477,9 +596,11 @@ class SARMRewardModel(PreTrainedRewardModel): 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: @@ -524,34 +645,67 @@ class SARMRewardModel(PreTrainedRewardModel): 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, - lang_emb: torch.Tensor, - state: torch.Tensor, - lengths: torch.Tensor, - targets: torch.Tensor, + 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.""" + """ + 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 - gt_stage = torch.floor(targets).long().clamp(0, num_classes - 1) - gt_tau = torch.remainder(targets, 1.0) + # 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: - stage_emb = gen_stage_emb(num_classes, targets) + # Mode 1: Use ground truth stage -> one-hot + stage_emb = gen_stage_emb(num_classes, targets) # (B, 1, T, C) else: - stage_idx = stage_pred.argmax(dim=-1) - stage_onehot = F.one_hot(stage_idx, num_classes=num_classes).float() - stage_emb = stage_onehot.unsqueeze(1) + # 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") @@ -562,9 +716,30 @@ class SARMRewardModel(PreTrainedRewardModel): } def forward(self, batch): - """Forward pass for SARM reward model training.""" + """ + 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") @@ -574,14 +749,17 @@ class SARMRewardModel(PreTrainedRewardModel): 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: @@ -590,8 +768,10 @@ class SARMRewardModel(PreTrainedRewardModel): 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") @@ -604,6 +784,7 @@ class SARMRewardModel(PreTrainedRewardModel): 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: @@ -623,5 +804,6 @@ def compute_stage_loss(stage_logits: torch.Tensor, target_stages: 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/rewards/sarm/processor_sarm.py b/src/lerobot/rewards/sarm/processor_sarm.py index d60914c4a..e10ac35d9 100644 --- a/src/lerobot/rewards/sarm/processor_sarm.py +++ b/src/lerobot/rewards/sarm/processor_sarm.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python - # Copyright 2025 The HuggingFace Inc. team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -70,14 +68,17 @@ class SARMEncodingProcessorStep(ProcessorStep): 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) @@ -113,6 +114,7 @@ class SARMEncodingProcessorStep(ProcessorStep): 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]) @@ -139,9 +141,11 @@ class SARMEncodingProcessorStep(ProcessorStep): 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 @@ -161,7 +165,15 @@ class SARMEncodingProcessorStep(ProcessorStep): ) def __call__(self, transition: EnvTransition) -> EnvTransition: - """Encode images, text, and normalize states in the transition.""" + """ + 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, {}) @@ -181,17 +193,20 @@ class SARMEncodingProcessorStep(ProcessorStep): 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, ...] + image = image[np.newaxis, ...] # (T, C, H, W) -> (1, T, C, H, W) elif image.ndim == 3: - image = image[np.newaxis, np.newaxis, ...] + image = image[np.newaxis, np.newaxis, ...] # (C, H, W) -> (1, 1, C, H, W) batch_size = image.shape[0] - total_frames = image.shape[1] + 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 + 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 @@ -207,13 +222,17 @@ class SARMEncodingProcessorStep(ProcessorStep): ) rewind_steps[b_idx] = rewind_step - lengths = n_obs_frames + rewind_steps + # 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 + 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 @@ -226,14 +245,15 @@ class SARMEncodingProcessorStep(ProcessorStep): state_tensor = torch.tensor(state_data, dtype=torch.float32) if state_tensor.ndim == 2: - state_tensor = state_tensor.unsqueeze(0) + state_tensor = state_tensor.unsqueeze(0) # (T, D) -> (1, T, D) elif state_tensor.ndim == 1: - state_tensor = state_tensor.unsqueeze(0).unsqueeze(0) + 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 + 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) @@ -241,19 +261,26 @@ class SARMEncodingProcessorStep(ProcessorStep): 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( @@ -261,8 +288,10 @@ class SARMEncodingProcessorStep(ProcessorStep): ) 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( @@ -304,11 +333,13 @@ class SARMEncodingProcessorStep(ProcessorStep): 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( @@ -322,6 +353,7 @@ class SARMEncodingProcessorStep(ProcessorStep): 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( @@ -363,7 +395,15 @@ class SARMEncodingProcessorStep(ProcessorStep): @torch.no_grad() def _encode_images_batch(self, images: np.ndarray) -> torch.Tensor: - """Encode a batch of images using CLIP.""" + """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:]) @@ -371,9 +411,10 @@ class SARMEncodingProcessorStep(ProcessorStep): images_list = [] for i in range(num_frames): img = images[i] - if img.shape[0] in [1, 3]: + 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) @@ -389,21 +430,31 @@ class SARMEncodingProcessorStep(ProcessorStep): 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) - all_embeddings = all_embeddings.reshape(batch_size, seq_length, -1) + 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).""" + """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()} diff --git a/src/lerobot/rewards/sarm/sarm_utils.py b/src/lerobot/rewards/sarm/sarm_utils.py index e7231db2e..d2cd92cff 100644 --- a/src/lerobot/rewards/sarm/sarm_utils.py +++ b/src/lerobot/rewards/sarm/sarm_utils.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python - # Copyright 2025 The HuggingFace Inc. team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -91,9 +89,29 @@ def compute_absolute_indices( n_obs_steps: int, frame_gap: int = 30, ) -> tuple[torch.Tensor, torch.Tensor]: - """Compute absolute frame indices with clamping for bidirectional observation sequence.""" + """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 @@ -103,8 +121,10 @@ def compute_absolute_indices( 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) @@ -118,13 +138,34 @@ def apply_rewind_augmentation( frame_gap: int = 30, rewind_step: int | None = None, ) -> tuple[int, list[int]]: - """Generate rewind frame indices for temporal augmentation.""" + """ + 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, [] + 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)) @@ -132,15 +173,18 @@ def apply_rewind_augmentation( 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) + idx = max(ep_start, idx) # Clamp to episode start rewind_indices.append(idx) return rewind_step, rewind_indices @@ -158,9 +202,10 @@ def pad_state_to_max_dim(state: torch.Tensor, max_state_dim: int) -> torch.Tenso """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] + return state[..., :max_state_dim] # Truncate if larger - padding = (0, max_state_dim - current_dim) + # 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) @@ -201,7 +246,25 @@ def normalize_stage_tau( 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.""" + """ + 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: