mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
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:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user