mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 19:26:16 +00:00
fix(rewards): restore full Classifier and SARM implementations
This commit is contained in:
@@ -32,6 +32,7 @@ class RewardClassifierConfig(RewardModelConfig):
|
|||||||
image_embedding_pooling_dim: int = 8
|
image_embedding_pooling_dim: int = 8
|
||||||
dropout_rate: float = 0.1
|
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.
|
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"
|
model_type: str = "cnn" # "transformer" or "cnn"
|
||||||
num_cameras: int = 2
|
num_cameras: int = 2
|
||||||
learning_rate: float = 1e-4
|
learning_rate: float = 1e-4
|
||||||
|
|||||||
@@ -97,10 +97,7 @@ class SpatialLearnedEmbeddings(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class Classifier(PreTrainedRewardModel):
|
class Classifier(PreTrainedRewardModel):
|
||||||
"""Image classifier built on top of a pre-trained encoder.
|
"""Image classifier built on top of a pre-trained encoder."""
|
||||||
|
|
||||||
Binary success/failure classifier from images. Trainable via ``forward()``.
|
|
||||||
"""
|
|
||||||
|
|
||||||
name = "reward_classifier"
|
name = "reward_classifier"
|
||||||
config_class = RewardClassifierConfig
|
config_class = RewardClassifierConfig
|
||||||
@@ -108,7 +105,6 @@ class Classifier(PreTrainedRewardModel):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: RewardClassifierConfig,
|
config: RewardClassifierConfig,
|
||||||
**kwargs,
|
|
||||||
):
|
):
|
||||||
from transformers import AutoModel
|
from transformers import AutoModel
|
||||||
|
|
||||||
@@ -215,6 +211,7 @@ class Classifier(PreTrainedRewardModel):
|
|||||||
|
|
||||||
def extract_images_and_labels(self, batch: dict[str, Tensor]) -> tuple[list, Tensor]:
|
def extract_images_and_labels(self, batch: dict[str, Tensor]) -> tuple[list, Tensor]:
|
||||||
"""Extract image tensors and label tensors from batch."""
|
"""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)]
|
images = [batch[key] for key in self.config.input_features if key.startswith(OBS_IMAGE)]
|
||||||
labels = batch[REWARD]
|
labels = batch[REWARD]
|
||||||
|
|
||||||
@@ -279,6 +276,11 @@ class Classifier(PreTrainedRewardModel):
|
|||||||
|
|
||||||
def predict_reward(self, batch, threshold=0.5):
|
def predict_reward(self, batch, threshold=0.5):
|
||||||
"""Eval method. Returns predicted reward with the decision threshold as argument."""
|
"""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)]
|
images = [batch[key] for key in self.config.input_features if key.startswith(OBS_IMAGE)]
|
||||||
|
|
||||||
if self.config.num_classes == 2:
|
if self.config.num_classes == 2:
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu
|
# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu
|
||||||
# and The HuggingFace Inc. team. All rights reserved.
|
# 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)
|
frame_gap: int = 30 # Frame gap between frames (at 30 fps = 1 second)
|
||||||
max_rewind_steps: int = 4 # Maximum rewind steps for temporal augmentation
|
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
|
# Architecture params
|
||||||
image_dim: int = 512
|
image_dim: int = 512
|
||||||
text_dim: int = 512
|
text_dim: int = 512
|
||||||
@@ -66,7 +68,7 @@ class SARMConfig(RewardModelConfig):
|
|||||||
batch_size: int = 64
|
batch_size: int = 64
|
||||||
clip_batch_size: int = 64
|
clip_batch_size: int = 64
|
||||||
dropout: float = 0.1
|
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
|
rewind_probability: float = 0.8
|
||||||
language_perturbation_probability: float = 0.2
|
language_perturbation_probability: float = 0.2
|
||||||
@@ -82,7 +84,8 @@ class SARMConfig(RewardModelConfig):
|
|||||||
dense_temporal_proportions: list | None = None
|
dense_temporal_proportions: list | None = None
|
||||||
|
|
||||||
pretrained_model_path: str | 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
|
state_key: str = OBS_STATE
|
||||||
|
|
||||||
# Populated by the processor (video_features, state_features, text_features)
|
# Populated by the processor (video_features, state_features, text_features)
|
||||||
@@ -114,6 +117,7 @@ class SARMConfig(RewardModelConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.annotation_mode == "single_stage":
|
if self.annotation_mode == "single_stage":
|
||||||
|
# Use task description as stage name, full episode as one stage
|
||||||
self.num_sparse_stages = 1
|
self.num_sparse_stages = 1
|
||||||
self.sparse_subtask_names = ["task"]
|
self.sparse_subtask_names = ["task"]
|
||||||
self.sparse_temporal_proportions = [1.0]
|
self.sparse_temporal_proportions = [1.0]
|
||||||
@@ -201,7 +205,11 @@ class SARMConfig(RewardModelConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def num_frames(self) -> int:
|
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
|
return 1 + self.n_obs_steps + self.max_rewind_steps
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -210,7 +218,14 @@ class SARMConfig(RewardModelConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int]:
|
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
|
half_steps = self.n_obs_steps // 2
|
||||||
|
|
||||||
past_deltas = [-self.frame_gap * i for i in range(half_steps, 0, -1)]
|
past_deltas = [-self.frame_gap * i for i in range(half_steps, 0, -1)]
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu
|
# Copyright 2025 Qianzhong Chen, Justin Yu, Mac Schwager, Pieter Abbeel, Yide Shentu, Philipp Wu
|
||||||
# and The HuggingFace Inc. team. All rights reserved.
|
# 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))
|
self.first_pos = nn.Parameter(torch.zeros(1, d_model))
|
||||||
|
|
||||||
# Shared fusion MLP
|
# Shared fusion MLP
|
||||||
|
# Fuses (num_cameras + 2) streams: cameras + lang + state
|
||||||
fused_in = d_model * (num_cameras + 2)
|
fused_in = d_model * (num_cameras + 2)
|
||||||
self.fusion_backbone = nn.Sequential(
|
self.fusion_backbone = nn.Sequential(
|
||||||
nn.LayerNorm(fused_in),
|
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
|
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:
|
if lang_emb.dim() == 3:
|
||||||
|
# (B, T, E) -> (B, T, D) -> (B, 1, T, D)
|
||||||
lang_proj = self.lang_proj(lang_emb).unsqueeze(1)
|
lang_proj = self.lang_proj(lang_emb).unsqueeze(1)
|
||||||
else:
|
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)
|
lang_proj = self.lang_proj(lang_emb).unsqueeze(1).unsqueeze(2).expand(B, 1, T, D)
|
||||||
return lang_proj
|
return lang_proj
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
img_seq: torch.Tensor,
|
img_seq: torch.Tensor, # (B, N, T, vis_emb_dim)
|
||||||
lang_emb: torch.Tensor,
|
lang_emb: torch.Tensor, # (B, E) or (B, T, E)
|
||||||
state: torch.Tensor,
|
state: torch.Tensor, # (B, T, state_dim)
|
||||||
lengths: torch.Tensor,
|
lengths: torch.Tensor, # (B,) - valid sequence lengths
|
||||||
scheme: str = "sparse",
|
scheme: str = "sparse", # "sparse" or "dense"
|
||||||
) -> torch.Tensor:
|
) -> 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())}."
|
assert scheme in self.heads, f"Unknown scheme '{scheme}'. Use one of {list(self.heads.keys())}."
|
||||||
|
|
||||||
B, N, T, _ = img_seq.shape # noqa: N806
|
B, N, T, _ = img_seq.shape # noqa: N806
|
||||||
D = self.d_model # noqa: N806
|
D = self.d_model # noqa: N806
|
||||||
device = img_seq.device
|
device = img_seq.device
|
||||||
|
|
||||||
vis_proj = self.visual_proj(img_seq)
|
# Project inputs
|
||||||
state_proj = self.state_proj(state).unsqueeze(1)
|
vis_proj = self.visual_proj(img_seq) # (B, N, T, D)
|
||||||
lang_proj = self._prep_lang(lang_emb, B, 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)
|
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
|
x[:, :N, 0, :] = x[:, :N, 0, :] + self.first_pos
|
||||||
|
|
||||||
|
# Flatten to tokens for Transformer
|
||||||
x_tokens = x.view(B, (N + 2) * T, D)
|
x_tokens = x.view(B, (N + 2) * T, D)
|
||||||
L = x_tokens.size(1) # noqa: N806
|
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)
|
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)
|
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)
|
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)
|
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
|
return logits
|
||||||
|
|
||||||
|
|
||||||
@@ -150,6 +183,10 @@ class SubtaskTransformer(nn.Module):
|
|||||||
Subtask progress regression transformer for SARM.
|
Subtask progress regression transformer for SARM.
|
||||||
|
|
||||||
Predicts within-stage normalized progress (tau) conditioned on stage prior.
|
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__(
|
def __init__(
|
||||||
@@ -167,15 +204,20 @@ class SubtaskTransformer(nn.Module):
|
|||||||
self.d_model = d_model
|
self.d_model = d_model
|
||||||
self.num_cameras = num_cameras
|
self.num_cameras = num_cameras
|
||||||
|
|
||||||
|
# Projections
|
||||||
self.lang_proj = nn.Linear(text_emb_dim, d_model)
|
self.lang_proj = nn.Linear(text_emb_dim, d_model)
|
||||||
self.visual_proj = nn.Linear(vis_emb_dim, d_model)
|
self.visual_proj = nn.Linear(vis_emb_dim, d_model)
|
||||||
self.state_proj = nn.Linear(state_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)
|
enc = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, batch_first=True)
|
||||||
self.transformer = nn.TransformerEncoder(enc, n_layers)
|
self.transformer = nn.TransformerEncoder(enc, n_layers)
|
||||||
|
|
||||||
|
# Learned bias on first visual frame
|
||||||
self.first_pos = nn.Parameter(torch.zeros(1, d_model))
|
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)
|
fused_in = d_model * (num_cameras + 3)
|
||||||
self.fusion_backbone = nn.Sequential(
|
self.fusion_backbone = nn.Sequential(
|
||||||
nn.LayerNorm(fused_in),
|
nn.LayerNorm(fused_in),
|
||||||
@@ -183,6 +225,7 @@ class SubtaskTransformer(nn.Module):
|
|||||||
nn.ReLU(),
|
nn.ReLU(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Scheme-specific regression heads
|
||||||
self.heads = nn.ModuleDict(
|
self.heads = nn.ModuleDict(
|
||||||
{
|
{
|
||||||
"sparse": nn.Linear(d_model, 1),
|
"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
|
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:
|
if lang_emb.dim() == 3:
|
||||||
|
# (B, T, E) -> (B, T, D) -> (B, 1, T, D)
|
||||||
return self.lang_proj(lang_emb).unsqueeze(1)
|
return self.lang_proj(lang_emb).unsqueeze(1)
|
||||||
else:
|
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)
|
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:
|
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
|
B, one, T, C = stage_prior.shape # noqa: N806
|
||||||
D = self.d_model # noqa: N806
|
D = self.d_model # noqa: N806
|
||||||
if D == C:
|
if D == C:
|
||||||
@@ -209,51 +266,87 @@ class SubtaskTransformer(nn.Module):
|
|||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
img_seq: torch.Tensor,
|
img_seq: torch.Tensor, # (B, N, T, vis_emb_dim)
|
||||||
lang_emb: torch.Tensor,
|
lang_emb: torch.Tensor, # (B, E) or (B, T, E)
|
||||||
state: torch.Tensor,
|
state: torch.Tensor, # (B, T, state_dim)
|
||||||
lengths: torch.Tensor,
|
lengths: torch.Tensor, # (B,) - valid sequence lengths
|
||||||
stage_prior: torch.Tensor,
|
stage_prior: torch.Tensor, # (B, 1, T, C) one-hot from gen_stage_emb
|
||||||
scheme: str = "sparse",
|
scheme: str = "sparse", # "sparse" or "dense"
|
||||||
) -> torch.Tensor:
|
) -> 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())}."
|
assert scheme in self.heads, f"Unknown scheme '{scheme}'. Use one of {list(self.heads.keys())}."
|
||||||
|
|
||||||
B, N, T, _ = img_seq.shape # noqa: N806
|
B, N, T, _ = img_seq.shape # noqa: N806
|
||||||
D = self.d_model # noqa: N806
|
D = self.d_model # noqa: N806
|
||||||
device = img_seq.device
|
device = img_seq.device
|
||||||
|
|
||||||
vis_proj = self.visual_proj(img_seq)
|
# Project inputs
|
||||||
state_proj = self.state_proj(state).unsqueeze(1)
|
vis_proj = self.visual_proj(img_seq) # (B, N, T, D)
|
||||||
lang_proj = self._prep_lang(lang_emb, B, T, D)
|
state_proj = self.state_proj(state).unsqueeze(1) # (B, 1, T, D)
|
||||||
stage_emb = self._stage_to_dmodel(stage_prior)
|
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)
|
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
|
x[:, :N, 0, :] = x[:, :N, 0, :] + self.first_pos
|
||||||
|
|
||||||
|
# Flatten to tokens
|
||||||
x_tokens = x.view(B, (N + 3) * T, D)
|
x_tokens = x.view(B, (N + 3) * T, D)
|
||||||
L = x_tokens.size(1) # noqa: N806
|
L = x_tokens.size(1) # noqa: N806
|
||||||
|
|
||||||
|
# Create padding mask
|
||||||
base_mask = torch.arange(T, device=device).expand(B, T) >= lengths.unsqueeze(1)
|
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)
|
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)
|
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)
|
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 = h.view(B, N + 3, T, D)
|
||||||
h_flat = h.permute(0, 2, 1, 3).reshape(B, T, (N + 3) * 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
|
return r
|
||||||
|
|
||||||
|
|
||||||
def gen_stage_emb(num_classes: int, targets: torch.Tensor) -> torch.Tensor:
|
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
|
C = num_classes # noqa: N806
|
||||||
stage_onehot = torch.eye(C, device=targets.device)[idx]
|
# Identity-lookup one-hot
|
||||||
stage_onehot = stage_onehot.unsqueeze(1)
|
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
|
return stage_onehot
|
||||||
|
|
||||||
|
|
||||||
@@ -271,8 +364,8 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
name = "sarm"
|
name = "sarm"
|
||||||
config_class = SARMConfig
|
config_class = SARMConfig
|
||||||
|
|
||||||
def __init__(self, config: SARMConfig, dataset_stats: dict | None = None, dataset_meta=None, **kwargs):
|
def __init__(self, config: SARMConfig, dataset_stats: dict | None = None, dataset_meta=None):
|
||||||
super().__init__(config)
|
super().__init__(config, dataset_stats)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
self.config = config
|
self.config = config
|
||||||
self.dataset_stats = dataset_stats
|
self.dataset_stats = dataset_stats
|
||||||
@@ -295,7 +388,7 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
n_layers=config.num_layers,
|
n_layers=config.num_layers,
|
||||||
n_heads=config.num_heads,
|
n_heads=config.num_heads,
|
||||||
dropout=config.dropout,
|
dropout=config.dropout,
|
||||||
num_cameras=1,
|
num_cameras=1, # Single camera for now
|
||||||
num_classes_sparse=config.num_sparse_stages,
|
num_classes_sparse=config.num_sparse_stages,
|
||||||
num_classes_dense=config.num_dense_stages or 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.stage_model.to(self.device)
|
||||||
self.subtask_model.to(self.device)
|
self.subtask_model.to(self.device)
|
||||||
|
|
||||||
|
# GT/predicted stage ratio for teacher forcing
|
||||||
self.gt_stage_ratio = 0.75
|
self.gt_stage_ratio = 0.75
|
||||||
|
|
||||||
if config.uses_dual_heads:
|
if config.uses_dual_heads:
|
||||||
@@ -410,6 +504,20 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
This is the canonical method for SARM reward computation, used for:
|
This is the canonical method for SARM reward computation, used for:
|
||||||
- Inference/visualization
|
- Inference/visualization
|
||||||
- RA-BC weight computation
|
- 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):
|
if isinstance(text_embeddings, np.ndarray):
|
||||||
text_embeddings = torch.tensor(text_embeddings, dtype=torch.float32)
|
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):
|
if state_features is not None and isinstance(state_features, np.ndarray):
|
||||||
state_features = torch.tensor(state_features, dtype=torch.float32)
|
state_features = torch.tensor(state_features, dtype=torch.float32)
|
||||||
|
|
||||||
|
# Handle single sample case
|
||||||
if text_embeddings.dim() == 1:
|
if text_embeddings.dim() == 1:
|
||||||
text_embeddings = text_embeddings.unsqueeze(0)
|
text_embeddings = text_embeddings.unsqueeze(0)
|
||||||
video_embeddings = video_embeddings.unsqueeze(0)
|
video_embeddings = video_embeddings.unsqueeze(0)
|
||||||
@@ -432,11 +541,14 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
|
|
||||||
scheme = head_mode
|
scheme = head_mode
|
||||||
|
|
||||||
|
# Default lengths if not provided
|
||||||
if lengths is None:
|
if lengths is None:
|
||||||
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32)
|
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32)
|
||||||
elif isinstance(lengths, np.ndarray):
|
elif isinstance(lengths, np.ndarray):
|
||||||
lengths = torch.tensor(lengths, dtype=torch.int32)
|
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)
|
img_seq = video_embeddings.unsqueeze(1).to(self.device)
|
||||||
lang_emb = text_embeddings.to(self.device)
|
lang_emb = text_embeddings.to(self.device)
|
||||||
state = (
|
state = (
|
||||||
@@ -446,22 +558,29 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
)
|
)
|
||||||
lens = lengths.to(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)
|
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
|
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_logits = self.stage_model(img_seq, lang_emb, state, lens, scheme=scheme)
|
||||||
stage_probs = F.softmax(stage_logits, dim=-1)
|
stage_probs = F.softmax(stage_logits, dim=-1) # (B, T, num_classes)
|
||||||
stage_idx = stage_probs.argmax(dim=-1)
|
stage_idx = stage_probs.argmax(dim=-1) # (B, T)
|
||||||
stage_conf = stage_probs.gather(-1, stage_idx.unsqueeze(-1)).squeeze(-1)
|
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()
|
# Create one-hot stage prior
|
||||||
stage_emb = stage_onehot.unsqueeze(1)
|
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)
|
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":
|
if scheme == "sparse":
|
||||||
normalized_reward = normalize_stage_tau(
|
normalized_reward = normalize_stage_tau(
|
||||||
raw_reward,
|
raw_reward,
|
||||||
@@ -477,9 +596,11 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
subtask_names=self.config.dense_subtask_names,
|
subtask_names=self.config.dense_subtask_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Default frame index is n_obs_steps (last observation frame)
|
||||||
if frame_index is None:
|
if frame_index is None:
|
||||||
frame_index = self.config.n_obs_steps
|
frame_index = self.config.n_obs_steps
|
||||||
|
|
||||||
|
# Prepare outputs (batch mode or no smoothing)
|
||||||
if return_all_frames:
|
if return_all_frames:
|
||||||
rewards = normalized_reward.cpu().numpy()
|
rewards = normalized_reward.cpu().numpy()
|
||||||
else:
|
else:
|
||||||
@@ -524,34 +645,67 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
|
"""Required by PreTrainedPolicy but not used for reward models."""
|
||||||
pass
|
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(
|
def _train_step(
|
||||||
self,
|
self,
|
||||||
img_emb: torch.Tensor,
|
img_emb: torch.Tensor, # (B, N, T, D)
|
||||||
lang_emb: torch.Tensor,
|
lang_emb: torch.Tensor, # (B, E) or (B, T, E)
|
||||||
state: torch.Tensor,
|
state: torch.Tensor, # (B, T, state_dim)
|
||||||
lengths: torch.Tensor,
|
lengths: torch.Tensor, # (B,)
|
||||||
targets: torch.Tensor,
|
targets: torch.Tensor, # (B, T) - format: stage.tau
|
||||||
scheme: str,
|
scheme: str,
|
||||||
) -> dict[str, torch.Tensor]:
|
) -> 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
|
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)
|
# Ground truth: stage (integer) and tau (fractional)
|
||||||
gt_tau = torch.remainder(targets, 1.0)
|
# 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)
|
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:
|
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:
|
else:
|
||||||
stage_idx = stage_pred.argmax(dim=-1)
|
# Mode 2: Use predicted stage argmax -> one-hot
|
||||||
stage_onehot = F.one_hot(stage_idx, num_classes=num_classes).float()
|
stage_idx = stage_pred.argmax(dim=-1) # (B, T)
|
||||||
stage_emb = stage_onehot.unsqueeze(1)
|
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)
|
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")
|
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")
|
subtask_loss = F.mse_loss(tau_pred, gt_tau, reduction="mean")
|
||||||
|
|
||||||
@@ -562,9 +716,30 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def forward(self, batch):
|
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)
|
observation = batch.get(OBS_STR, batch)
|
||||||
|
|
||||||
|
# Extract features
|
||||||
video_features = observation["video_features"].to(self.device)
|
video_features = observation["video_features"].to(self.device)
|
||||||
text_features = observation["text_features"].to(self.device)
|
text_features = observation["text_features"].to(self.device)
|
||||||
state_features = observation.get("state_features")
|
state_features = observation.get("state_features")
|
||||||
@@ -574,14 +749,17 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
batch_size = video_features.shape[0]
|
batch_size = video_features.shape[0]
|
||||||
seq_len = video_features.shape[1]
|
seq_len = video_features.shape[1]
|
||||||
|
|
||||||
|
# Get lengths (default to full sequence)
|
||||||
lengths = observation.get("lengths")
|
lengths = observation.get("lengths")
|
||||||
if lengths is None:
|
if lengths is None:
|
||||||
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32, device=self.device)
|
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32, device=self.device)
|
||||||
else:
|
else:
|
||||||
lengths = lengths.to(self.device)
|
lengths = lengths.to(self.device)
|
||||||
|
|
||||||
|
# Reshape video to (B, N, T, D) - single camera
|
||||||
img_emb = video_features.unsqueeze(1)
|
img_emb = video_features.unsqueeze(1)
|
||||||
|
|
||||||
|
# Pad state to max_state_dim
|
||||||
if state_features is None:
|
if state_features is None:
|
||||||
state_features = torch.zeros(batch_size, seq_len, self.config.max_state_dim, device=self.device)
|
state_features = torch.zeros(batch_size, seq_len, self.config.max_state_dim, device=self.device)
|
||||||
else:
|
else:
|
||||||
@@ -590,8 +768,10 @@ class SARMRewardModel(PreTrainedRewardModel):
|
|||||||
output_dict = {}
|
output_dict = {}
|
||||||
total_loss = torch.tensor(0.0, device=self.device)
|
total_loss = torch.tensor(0.0, device=self.device)
|
||||||
|
|
||||||
|
# Sparse training (always)
|
||||||
sparse_targets = observation.get("sparse_targets")
|
sparse_targets = observation.get("sparse_targets")
|
||||||
if sparse_targets is None:
|
if sparse_targets is None:
|
||||||
|
# Try legacy format
|
||||||
sparse_targets = observation.get("targets")
|
sparse_targets = observation.get("targets")
|
||||||
if sparse_targets is None:
|
if sparse_targets is None:
|
||||||
raise ValueError("sparse_targets (or targets) is required for SARM training")
|
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()
|
output_dict["sparse_subtask_loss"] = sparse_result["subtask_loss"].item()
|
||||||
total_loss = total_loss + sparse_result["total_loss"]
|
total_loss = total_loss + sparse_result["total_loss"]
|
||||||
|
|
||||||
|
# Dense training (if dual mode)
|
||||||
if self.config.uses_dual_heads:
|
if self.config.uses_dual_heads:
|
||||||
dense_targets = observation.get("dense_targets")
|
dense_targets = observation.get("dense_targets")
|
||||||
if dense_targets is not None:
|
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."""
|
"""Compute cross-entropy loss for stage classification."""
|
||||||
_, _, num_stages = stage_logits.shape
|
_, _, num_stages = stage_logits.shape
|
||||||
stage_logits_flat = stage_logits.reshape(-1, num_stages)
|
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)
|
target_stages_flat = target_stages.reshape(-1).clamp(0, num_stages - 1)
|
||||||
return F.cross_entropy(stage_logits_flat, target_stages_flat)
|
return F.cross_entropy(stage_logits_flat, target_stages_flat)
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
||||||
#
|
#
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
@@ -70,14 +68,17 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
self.dataset_stats = dataset_stats
|
self.dataset_stats = dataset_stats
|
||||||
self.annotation_mode = config.annotation_mode
|
self.annotation_mode = config.annotation_mode
|
||||||
|
|
||||||
|
# Helper to create temporal proportions dict
|
||||||
def make_props_dict(names, props):
|
def make_props_dict(names, props):
|
||||||
return dict(zip(names, props, strict=True)) if names and props else None
|
return dict(zip(names, props, strict=True)) if names and props else None
|
||||||
|
|
||||||
|
# Sparse annotations (always needed)
|
||||||
self.sparse_temporal_proportions = make_props_dict(
|
self.sparse_temporal_proportions = make_props_dict(
|
||||||
config.sparse_subtask_names, config.sparse_temporal_proportions
|
config.sparse_subtask_names, config.sparse_temporal_proportions
|
||||||
)
|
)
|
||||||
self.sparse_subtask_names = config.sparse_subtask_names
|
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_subtask_names = config.dense_subtask_names if config.uses_dual_heads else None
|
||||||
self.dense_temporal_proportions = (
|
self.dense_temporal_proportions = (
|
||||||
make_props_dict(config.dense_subtask_names, config.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)))
|
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:
|
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 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],
|
global_names: list[str],
|
||||||
) -> tuple[list | None, list | None, list | None]:
|
) -> tuple[list | None, list | None, list | None]:
|
||||||
"""Load subtask annotations for an episode from DataFrame."""
|
"""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:
|
if episodes_df is None or len(global_names) == 1:
|
||||||
return None, None, None
|
return None, None, None
|
||||||
|
|
||||||
|
# Resolve column name with fallback
|
||||||
def col(suffix):
|
def col(suffix):
|
||||||
prefixed = f"{annotation_type}_{suffix}"
|
prefixed = f"{annotation_type}_{suffix}"
|
||||||
return prefixed if prefixed in episodes_df.columns else suffix
|
return prefixed if prefixed in episodes_df.columns else suffix
|
||||||
@@ -161,7 +165,15 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
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)
|
new_transition = transition.copy() if hasattr(transition, "copy") else dict(transition)
|
||||||
observation = new_transition.get(TransitionKey.OBSERVATION)
|
observation = new_transition.get(TransitionKey.OBSERVATION)
|
||||||
comp_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
comp_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||||
@@ -181,17 +193,20 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
if isinstance(image, torch.Tensor):
|
if isinstance(image, torch.Tensor):
|
||||||
image = image.cpu().numpy()
|
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:
|
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:
|
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]
|
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
|
n_obs_steps = self.config.n_obs_steps
|
||||||
max_rewind_steps = self.config.max_rewind_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)
|
rewind_steps = torch.zeros(batch_size, dtype=torch.int32)
|
||||||
apply_rewind = self.training and random.random() < self.config.rewind_probability
|
apply_rewind = self.training and random.random() < self.config.rewind_probability
|
||||||
|
|
||||||
@@ -207,13 +222,17 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
)
|
)
|
||||||
rewind_steps[b_idx] = rewind_step
|
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):
|
for b_idx in range(batch_size):
|
||||||
valid_len = lengths[b_idx].item()
|
valid_len = lengths[b_idx].item()
|
||||||
if valid_len < total_frames:
|
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)
|
video_features = self._encode_images_batch(image)
|
||||||
observation["video_features"] = video_features
|
observation["video_features"] = video_features
|
||||||
|
|
||||||
@@ -226,14 +245,15 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
state_tensor = torch.tensor(state_data, dtype=torch.float32)
|
state_tensor = torch.tensor(state_data, dtype=torch.float32)
|
||||||
|
|
||||||
if state_tensor.ndim == 2:
|
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:
|
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):
|
for b_idx in range(batch_size):
|
||||||
valid_len = lengths[b_idx].item()
|
valid_len = lengths[b_idx].item()
|
||||||
if valid_len < state_tensor.shape[1]:
|
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)
|
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):
|
if isinstance(task, list):
|
||||||
task = task[0] if task else ""
|
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
|
apply_perturbation = self.training and random.random() < self.config.language_perturbation_probability
|
||||||
if apply_perturbation:
|
if apply_perturbation:
|
||||||
task = self._generate_perturbed_task()
|
task = self._generate_perturbed_task()
|
||||||
|
|
||||||
|
# Encode text with CLIP
|
||||||
observation["text_features"] = self._encode_text_clip(task, batch_size)
|
observation["text_features"] = self._encode_text_clip(task, batch_size)
|
||||||
|
|
||||||
|
# Store lengths for model
|
||||||
observation["lengths"] = lengths
|
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:
|
if self.dataset_meta is not None:
|
||||||
episodes_df = self.dataset_meta.episodes.to_pandas()
|
episodes_df = self.dataset_meta.episodes.to_pandas()
|
||||||
|
|
||||||
|
# Generate sparse targets
|
||||||
if self.sparse_temporal_proportions is not None:
|
if self.sparse_temporal_proportions is not None:
|
||||||
if apply_perturbation:
|
if apply_perturbation:
|
||||||
|
# Zero targets when language is perturbed
|
||||||
sparse_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32)
|
sparse_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32)
|
||||||
else:
|
else:
|
||||||
sparse_targets = self._compute_batch_targets(
|
sparse_targets = self._compute_batch_targets(
|
||||||
@@ -261,8 +288,10 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
)
|
)
|
||||||
observation["sparse_targets"] = sparse_targets
|
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 self.config.uses_dual_heads and self.dense_temporal_proportions is not None:
|
||||||
if apply_perturbation:
|
if apply_perturbation:
|
||||||
|
# Zero targets when language is perturbed
|
||||||
dense_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32)
|
dense_targets = torch.zeros(batch_size, total_frames, dtype=torch.float32)
|
||||||
else:
|
else:
|
||||||
dense_targets = self._compute_batch_targets(
|
dense_targets = self._compute_batch_targets(
|
||||||
@@ -304,11 +333,13 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
ep_idx, episodes_df, annotation_type, global_names
|
ep_idx, episodes_df, annotation_type, global_names
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Compute observation frame indices
|
||||||
obs_indices, _ = compute_absolute_indices(
|
obs_indices, _ = compute_absolute_indices(
|
||||||
frame_idx, ep_start, ep_end, n_obs_steps, frame_gap=frame_gap
|
frame_idx, ep_start, ep_end, n_obs_steps, frame_gap=frame_gap
|
||||||
)
|
)
|
||||||
obs_indices = obs_indices.tolist()
|
obs_indices = obs_indices.tolist()
|
||||||
|
|
||||||
|
# Compute targets for observation frames
|
||||||
for t_idx, abs_idx in enumerate(obs_indices):
|
for t_idx, abs_idx in enumerate(obs_indices):
|
||||||
rel_frame = abs_idx - ep_start
|
rel_frame = abs_idx - ep_start
|
||||||
targets[b_idx, t_idx] = find_stage_and_tau(
|
targets[b_idx, t_idx] = find_stage_and_tau(
|
||||||
@@ -322,6 +353,7 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
return_combined=True,
|
return_combined=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Compute targets for rewind frames (if any)
|
||||||
rewind_step = rewind_steps[b_idx].item()
|
rewind_step = rewind_steps[b_idx].item()
|
||||||
if rewind_step > 0:
|
if rewind_step > 0:
|
||||||
_, rewind_indices = apply_rewind_augmentation(
|
_, rewind_indices = apply_rewind_augmentation(
|
||||||
@@ -363,7 +395,15 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def _encode_images_batch(self, images: np.ndarray) -> torch.Tensor:
|
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]
|
batch_size, seq_length = images.shape[0], images.shape[1]
|
||||||
images = images.reshape(batch_size * seq_length, *images.shape[2:])
|
images = images.reshape(batch_size * seq_length, *images.shape[2:])
|
||||||
|
|
||||||
@@ -371,9 +411,10 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
|||||||
images_list = []
|
images_list = []
|
||||||
for i in range(num_frames):
|
for i in range(num_frames):
|
||||||
img = images[i]
|
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)
|
img = img.transpose(1, 2, 0)
|
||||||
|
|
||||||
|
# Handle single channel
|
||||||
if img.shape[-1] == 1:
|
if img.shape[-1] == 1:
|
||||||
img = np.repeat(img, 3, axis=-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 = self.clip_processor(images=batch_imgs, return_tensors="pt")
|
||||||
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
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()
|
embeddings = self.clip_model.get_image_features(**inputs).detach().cpu()
|
||||||
|
|
||||||
|
# Handle single frame case
|
||||||
if embeddings.dim() == 1:
|
if embeddings.dim() == 1:
|
||||||
embeddings = embeddings.unsqueeze(0)
|
embeddings = embeddings.unsqueeze(0)
|
||||||
|
|
||||||
all_embeddings.append(embeddings)
|
all_embeddings.append(embeddings)
|
||||||
|
|
||||||
all_embeddings = torch.cat(all_embeddings)
|
all_embeddings = torch.cat(all_embeddings) # (B*T, 512)
|
||||||
all_embeddings = all_embeddings.reshape(batch_size, seq_length, -1)
|
all_embeddings = all_embeddings.reshape(batch_size, seq_length, -1) # (B, T, 512)
|
||||||
|
|
||||||
return all_embeddings
|
return all_embeddings
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def _encode_text_clip(self, text: str, batch_size: int) -> torch.Tensor:
|
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 = self.clip_processor.tokenizer([text], return_tensors="pt", padding=True, truncation=True)
|
||||||
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
||||||
#
|
#
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
@@ -91,9 +89,29 @@ def compute_absolute_indices(
|
|||||||
n_obs_steps: int,
|
n_obs_steps: int,
|
||||||
frame_gap: int = 30,
|
frame_gap: int = 30,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> 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
|
half_steps = n_obs_steps // 2
|
||||||
|
|
||||||
|
# Bidirectional deltas: past + current + future
|
||||||
past_deltas = [-frame_gap * i for i in range(half_steps, 0, -1)]
|
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)]
|
future_deltas = [frame_gap * i for i in range(1, half_steps + 1)]
|
||||||
delta_indices = past_deltas + [0] + future_deltas
|
delta_indices = past_deltas + [0] + future_deltas
|
||||||
@@ -103,8 +121,10 @@ def compute_absolute_indices(
|
|||||||
|
|
||||||
for delta in delta_indices:
|
for delta in delta_indices:
|
||||||
target_idx = frame_idx + delta
|
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))
|
clamped_idx = max(ep_start, min(ep_end - 1, target_idx))
|
||||||
frames.append(clamped_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)
|
out_of_bounds.append(1 if target_idx != clamped_idx else 0)
|
||||||
|
|
||||||
return torch.tensor(frames), torch.tensor(out_of_bounds)
|
return torch.tensor(frames), torch.tensor(out_of_bounds)
|
||||||
@@ -118,13 +138,34 @@ def apply_rewind_augmentation(
|
|||||||
frame_gap: int = 30,
|
frame_gap: int = 30,
|
||||||
rewind_step: int | None = None,
|
rewind_step: int | None = None,
|
||||||
) -> tuple[int, list[int]]:
|
) -> 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
|
half_steps = n_obs_steps // 2
|
||||||
earliest_obs_frame = frame_idx - half_steps * frame_gap
|
earliest_obs_frame = frame_idx - half_steps * frame_gap
|
||||||
|
|
||||||
|
# Required history: frames before earliest observation frame
|
||||||
if earliest_obs_frame <= ep_start:
|
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
|
available_history = earliest_obs_frame - ep_start
|
||||||
max_valid_step = available_history // frame_gap
|
max_valid_step = available_history // frame_gap
|
||||||
max_rewind = min(max_rewind_steps, max(0, max_valid_step))
|
max_rewind = min(max_rewind_steps, max(0, max_valid_step))
|
||||||
@@ -132,15 +173,18 @@ def apply_rewind_augmentation(
|
|||||||
if max_rewind <= 0:
|
if max_rewind <= 0:
|
||||||
return 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)
|
rewind_step = random.randint(1, max_rewind) if rewind_step is None else min(rewind_step, max_rewind)
|
||||||
|
|
||||||
if rewind_step == 0:
|
if rewind_step == 0:
|
||||||
return 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 = []
|
rewind_indices = []
|
||||||
for i in range(1, rewind_step + 1):
|
for i in range(1, rewind_step + 1):
|
||||||
idx = earliest_obs_frame - i * frame_gap
|
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)
|
rewind_indices.append(idx)
|
||||||
|
|
||||||
return rewind_step, rewind_indices
|
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."""
|
"""Pad the state tensor's last dimension to max_state_dim with zeros."""
|
||||||
current_dim = state.shape[-1]
|
current_dim = state.shape[-1]
|
||||||
if current_dim >= max_state_dim:
|
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)
|
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,
|
temporal_proportions: dict[str, float] | list[float] | None = None,
|
||||||
subtask_names: list[str] | None = None,
|
subtask_names: list[str] | None = None,
|
||||||
) -> float | torch.Tensor:
|
) -> 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:
|
if breakpoints is not None:
|
||||||
num_stages = len(breakpoints) - 1
|
num_stages = len(breakpoints) - 1
|
||||||
elif temporal_proportions is not None:
|
elif temporal_proportions is not None:
|
||||||
|
|||||||
Reference in New Issue
Block a user