fix(sarm): fail fast on unusable episode annotations (#4306)

* fix(sarm): warn when dense/sparse targets silently collapse to all-zero

In dense_only/dual modes, if meta/episodes/*.parquet has no usable
subtask columns (column absent or NaN), _load_episode_annotations
returns None and find_stage_and_tau yields stage 0 / tau 0 for every
frame. Training "succeeds" but the head silently learns to predict 0
everywhere, with no warning. This complements #2880 (which restored
loading of episodes_df): there the DataFrame is loaded but the
*_subtask_names column is missing/NaN.

Add a one-time validation at processor construction that logs a clear
warning (all episodes missing -> predict-all-zero; some missing ->
partial). Purely additive logging, no change to training math.

Closes #3842

* fix(sarm): fail fast on unusable episode annotations

---------

Co-authored-by: 1thanShih <Smartshithan1620.en12@nycu.edu.tw>
This commit is contained in:
Khalil Meftah
2026-08-03 13:30:00 +02:00
committed by GitHub
parent 3e37269dc6
commit f1efa588b8
2 changed files with 222 additions and 14 deletions
+85 -14
View File
@@ -16,6 +16,7 @@
from __future__ import annotations from __future__ import annotations
import logging
import random import random
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -69,6 +70,8 @@ from .sarm_utils import (
pad_state_to_max_dim, pad_state_to_max_dim,
) )
logger = logging.getLogger(__name__)
class SARMEncodingProcessorStep(ProcessorStep): class SARMEncodingProcessorStep(ProcessorStep):
"""ProcessorStep that encodes images and text with CLIP and generates stage and progress labels for SARM.""" """ProcessorStep that encodes images and text with CLIP and generates stage and progress labels for SARM."""
@@ -108,6 +111,8 @@ class SARMEncodingProcessorStep(ProcessorStep):
else None else None
) )
self._validate_annotation_columns()
self.device = torch.device( self.device = torch.device(
self.config.device if self.config.device else "cuda" if torch.cuda.is_available() else "cpu" self.config.device if self.config.device else "cuda" if torch.cuda.is_available() else "cpu"
) )
@@ -120,6 +125,78 @@ class SARMEncodingProcessorStep(ProcessorStep):
self.verbs = ["move", "grasp", "rotate", "push", "pull", "slide", "lift", "place"] self.verbs = ["move", "grasp", "rotate", "push", "pull", "slide", "lift", "place"]
self.fake = Faker() self.fake = Faker()
@staticmethod
def _resolve_annotation_column(episodes_df: pd.DataFrame, annotation_type: str, suffix: str) -> str:
"""Resolve a mode-specific annotation column, falling back to the legacy unprefixed name."""
prefixed = f"{annotation_type}_{suffix}"
return prefixed if prefixed in episodes_df.columns else suffix
@staticmethod
def _annotations_are_usable(names: Any, starts: Any, ends: Any) -> bool:
"""Return whether an episode has non-empty, aligned annotation arrays."""
values = (names, starts, ends)
if not all(isinstance(value, (list, tuple, np.ndarray)) for value in values):
return False
lengths = {len(value) for value in values}
return len(lengths) == 1 and next(iter(lengths)) > 0
def _validate_annotation_columns(self) -> None:
"""Validate annotation coverage before loading models or generating training targets.
A multi-stage head with no usable episode annotations would otherwise train entirely
against all-zero targets. Reject that configuration and warn when only part of the
dataset is usable.
"""
if self.dataset_meta is None:
return
episodes_df = self.dataset_meta.episodes.to_pandas()
num_episodes = len(episodes_df)
modes = []
if self.dense_subtask_names and len(self.dense_subtask_names) > 1:
modes.append(("dense", self.dense_subtask_names))
if self.sparse_subtask_names and len(self.sparse_subtask_names) > 1:
modes.append(("sparse", self.sparse_subtask_names))
for annotation_type, names in modes:
columns = [
self._resolve_annotation_column(episodes_df, annotation_type, suffix)
for suffix in ("subtask_names", "subtask_start_frames", "subtask_end_frames")
]
missing_columns = [column for column in columns if column not in episodes_df.columns]
if missing_columns:
num_usable = 0
else:
num_usable = sum(
self._annotations_are_usable(*(episodes_df.loc[ep_idx, column] for column in columns))
for ep_idx in episodes_df.index
)
if num_usable == 0:
missing_columns_message = (
f" Missing required columns: {', '.join(missing_columns)}." if missing_columns else ""
)
raise ValueError(
f"SARM {annotation_type} head is configured with {len(names)} stages, but none of "
f"the {num_episodes} episodes have usable annotations in meta/episodes/*.parquet. "
f"Required columns: {', '.join(columns)}.{missing_columns_message} "
"Training would produce all-zero "
"targets. Materialize the annotations into the episodes metadata before training."
)
num_unusable = num_episodes - num_usable
if num_unusable:
logger.warning(
"SARM %s head: %d/%d episodes have unusable annotations in columns %s; "
"their targets will be 0 and only the %d annotated episodes will train the head.",
annotation_type,
num_unusable,
num_episodes,
", ".join(columns),
num_usable,
)
def _find_episode_for_frame(self, frame_idx: int) -> int: def _find_episode_for_frame(self, frame_idx: int) -> int:
"""Find the episode index for a given frame index.""" """Find the episode index for a given frame index."""
for ep_idx in range(len(self.dataset_meta.episodes)): for ep_idx in range(len(self.dataset_meta.episodes)):
@@ -167,24 +244,18 @@ class SARMEncodingProcessorStep(ProcessorStep):
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 columns = [
def col(suffix): self._resolve_annotation_column(episodes_df, annotation_type, suffix)
prefixed = f"{annotation_type}_{suffix}" for suffix in ("subtask_names", "subtask_start_frames", "subtask_end_frames")
return prefixed if prefixed in episodes_df.columns else suffix ]
if any(column not in episodes_df.columns for column in columns) or ep_idx >= len(episodes_df):
col_names = col("subtask_names")
if col_names not in episodes_df.columns or ep_idx >= len(episodes_df):
return None, None, None return None, None, None
subtask_names = episodes_df.loc[ep_idx, col_names] annotations = tuple(episodes_df.loc[ep_idx, column] for column in columns)
if subtask_names is None or (isinstance(subtask_names, float) and pd.isna(subtask_names)): if not self._annotations_are_usable(*annotations):
return None, None, None return None, None, None
return ( return annotations
subtask_names,
episodes_df.loc[ep_idx, col("subtask_start_frames")],
episodes_df.loc[ep_idx, col("subtask_end_frames")],
)
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """
+137
View File
@@ -692,3 +692,140 @@ class TestSARMEncodingProcessorStepEndToEnd:
assert abs(actual_dense - expected_dense) < 0.01, ( assert abs(actual_dense - expected_dense) < 0.01, (
f"Frame {frame}: dense mismatch {actual_dense:.3f} vs expected {expected_dense:.3f}" f"Frame {frame}: dense mismatch {actual_dense:.3f} vs expected {expected_dense:.3f}"
) )
def test_rejects_missing_dense_annotation_columns(self, mock_clip_model):
"""A multi-stage dense head must reject metadata with no annotation columns."""
from lerobot.rewards.sarm.processor_sarm import SARMEncodingProcessorStep
config = MockConfig(
annotation_mode="dense_only",
dense_subtask_names=["d1", "d2", "d3", "d4"],
dense_temporal_proportions=[0.25, 0.25, 0.25, 0.25],
)
# episodes metadata WITHOUT any dense_subtask_* columns
episodes = [
{"dataset_from_index": 0, "dataset_to_index": 100, "task": "t"},
{"dataset_from_index": 100, "dataset_to_index": 200, "task": "t"},
]
dataset_meta = MockDatasetMeta(episodes)
with pytest.raises(ValueError, match="Training would produce all-zero targets"):
SARMEncodingProcessorStep(config=config, dataset_meta=dataset_meta)
def test_rejects_dense_annotations_when_all_null(self, mock_clip_model):
"""Present-but-null annotation columns must also fail before training."""
from lerobot.rewards.sarm.processor_sarm import SARMEncodingProcessorStep
config = MockConfig(
annotation_mode="dense_only",
dense_subtask_names=["d1", "d2", "d3", "d4"],
dense_temporal_proportions=[0.25, 0.25, 0.25, 0.25],
)
episodes = [
{
"dataset_from_index": 0,
"dataset_to_index": 100,
"task": "t",
"dense_subtask_names": None,
"dense_subtask_start_frames": None,
"dense_subtask_end_frames": None,
},
{
"dataset_from_index": 100,
"dataset_to_index": 200,
"task": "t",
"dense_subtask_names": None,
"dense_subtask_start_frames": None,
"dense_subtask_end_frames": None,
},
]
dataset_meta = MockDatasetMeta(episodes)
with pytest.raises(ValueError, match="none of the 2 episodes have usable annotations"):
SARMEncodingProcessorStep(config=config, dataset_meta=dataset_meta)
def test_rejects_dense_annotations_with_missing_frame_column(self, mock_clip_model):
"""Names alone are not usable when a required frame-boundary column is absent."""
from lerobot.rewards.sarm.processor_sarm import SARMEncodingProcessorStep
config = MockConfig(
annotation_mode="dense_only",
dense_subtask_names=["d1", "d2"],
dense_temporal_proportions=[0.5, 0.5],
)
episodes = [
{
"dataset_from_index": 0,
"dataset_to_index": 100,
"task": "t",
"dense_subtask_names": ["d1", "d2"],
"dense_subtask_start_frames": [0, 50],
}
]
with pytest.raises(ValueError, match="Missing required columns: subtask_end_frames"):
SARMEncodingProcessorStep(config=config, dataset_meta=MockDatasetMeta(episodes))
def test_warns_when_dense_annotations_are_partial(self, mock_clip_model, caplog):
"""Partially annotated datasets remain supported but report exact coverage."""
import logging
from lerobot.rewards.sarm.processor_sarm import SARMEncodingProcessorStep
config = MockConfig(
annotation_mode="dense_only",
dense_subtask_names=["d1", "d2"],
dense_temporal_proportions=[0.5, 0.5],
)
episodes = [
{
"dataset_from_index": 0,
"dataset_to_index": 100,
"task": "t",
"dense_subtask_names": ["d1", "d2"],
"dense_subtask_start_frames": [0, 50],
"dense_subtask_end_frames": [49, 99],
},
{
"dataset_from_index": 100,
"dataset_to_index": 200,
"task": "t",
"dense_subtask_names": None,
"dense_subtask_start_frames": None,
"dense_subtask_end_frames": None,
},
]
with caplog.at_level(logging.WARNING, logger="lerobot.rewards.sarm.processor_sarm"):
SARMEncodingProcessorStep(config=config, dataset_meta=MockDatasetMeta(episodes))
assert "1/2 episodes have unusable annotations" in caplog.text
assert "only the 1 annotated episodes will train the head" in caplog.text
def test_no_warning_when_dense_annotations_present(self, mock_clip_model, caplog):
"""A fully annotated dataset must not emit an annotation-coverage warning."""
import logging
from lerobot.rewards.sarm.processor_sarm import SARMEncodingProcessorStep
config = MockConfig(
annotation_mode="dense_only",
dense_subtask_names=["d1", "d2", "d3", "d4"],
dense_temporal_proportions=[0.25, 0.25, 0.25, 0.25],
)
episodes = [
{
"dataset_from_index": 0,
"dataset_to_index": 100,
"task": "t",
"dense_subtask_names": ["d1", "d2", "d3", "d4"],
"dense_subtask_start_frames": [0, 25, 50, 75],
"dense_subtask_end_frames": [25, 50, 75, 100],
}
]
dataset_meta = MockDatasetMeta(episodes)
with caplog.at_level(logging.WARNING, logger="lerobot.rewards.sarm.processor_sarm"):
SARMEncodingProcessorStep(config=config, dataset_meta=dataset_meta)
assert not any("unusable annotations" in m for m in caplog.messages)