refactor(datasets): resolve episode indices consistently

Make LeRobotDataset own episode selection policy: resolve the allowlist
via resolve_episode_indices (replacing the old warn-only out-of-range
handling, which now raises on an empty selection) and apply episode_filter,
then hand the finalized set to DatasetReader. The reader no longer
re-resolves, removing the duplicate resolution, and computes num_frames
from episode metadata so it stays consistent with the streaming path.
This commit is contained in:
CarolinePascal
2026-08-04 16:54:55 +02:00
parent ef71bc8ccb
commit 6eaea07046
2 changed files with 15 additions and 13 deletions
+7 -7
View File
@@ -39,7 +39,6 @@ from .io_utils import (
hf_transform_to_torch,
load_nested_dataset,
)
from .utils import resolve_episode_indices
from .video_utils import decode_video_frames
@@ -69,8 +68,9 @@ class DatasetReader:
Args:
meta: Dataset metadata instance.
root: Local dataset root directory.
episodes: Optional list of episode indices to select. ``None``
means all episodes.
episodes: Optional list of episode indices to select, assumed
already validated by the caller. ``None`` means
all episodes.
tolerance_s: Timestamp synchronization tolerance in seconds.
video_backend: Video decoding backend identifier.
delta_timestamps: Optional dict mapping feature keys to lists of
@@ -84,7 +84,7 @@ class DatasetReader:
"""
self._meta = meta
self.root = root
self.episodes = resolve_episode_indices(episodes, meta.total_episodes)
self.episodes = episodes
self._tolerance_s = tolerance_s
self._video_backend = video_backend
if image_transforms is not None and not callable(image_transforms):
@@ -152,9 +152,9 @@ class DatasetReader:
@property
def num_frames(self) -> int:
"""Number of frames in selected episodes."""
if self.episodes is not None and self.hf_dataset is not None:
return len(self.hf_dataset)
return self._meta.total_frames
if self.episodes is None:
return self._meta.total_frames
return sum(self._meta.episodes[ep]["length"] for ep in self.episodes)
@property
def num_episodes(self) -> int:
+8 -6
View File
@@ -34,6 +34,7 @@ from .utils import (
create_lerobot_dataset_card,
get_safe_version,
is_valid_version,
resolve_episode_indices,
)
from .video_utils import (
StreamingVideoEncoder,
@@ -237,13 +238,14 @@ class LeRobotDataset(torch.utils.data.Dataset):
self.revision = self.meta.revision
self.meta.rescale_depth_stats(self._depth_output_unit)
if episodes is not None and any(
episode >= self.meta.total_episodes or episode < 0 for episode in episodes
):
logger.warning(
f"Some episodes in the provided episodes list are out of range for this dataset ({self.meta.total_episodes})."
# Selection policy (allowlist resolution + predicate filter) is owned here;
# the reader just consumes the finalized episode index set.
episodes = resolve_episode_indices(episodes, self.meta.total_episodes)
if episodes is not None and not episodes:
raise ValueError(
"No valid episodes: the requested episode selection is empty after resolving "
f"against the dataset range [0, {self.meta.total_episodes})."
)
if episode_filter is not None:
resolved = self.meta.filter_episodes(episode_filter, candidates=episodes)
if not resolved: