diff --git a/src/lerobot/datasets/factory.py b/src/lerobot/datasets/factory.py index b5a3c8db6..fb8b79382 100644 --- a/src/lerobot/datasets/factory.py +++ b/src/lerobot/datasets/factory.py @@ -114,6 +114,7 @@ def make_dataset( tolerance_s=cfg.tolerance_s, return_uint8=True, depth_output_unit=cfg.dataset.depth_output_unit, + video_backend=cfg.dataset.video_backend, data_root=cfg.dataset.streaming_data_root, episode_pool_size=cfg.dataset.streaming_episode_pool_size, prefetch_episodes=cfg.dataset.streaming_prefetch_episodes, @@ -204,6 +205,7 @@ def make_train_eval_datasets( tolerance_s=cfg.tolerance_s, return_uint8=True, depth_output_unit=cfg.dataset.depth_output_unit, + video_backend=cfg.dataset.video_backend, data_root=cfg.dataset.streaming_data_root, episode_pool_size=cfg.dataset.streaming_episode_pool_size, prefetch_episodes=cfg.dataset.streaming_prefetch_episodes, diff --git a/src/lerobot/datasets/streaming_dataset.py b/src/lerobot/datasets/streaming_dataset.py index cd47eecf5..0a21c94e0 100644 --- a/src/lerobot/datasets/streaming_dataset.py +++ b/src/lerobot/datasets/streaming_dataset.py @@ -25,13 +25,14 @@ import torch from lerobot.configs import DEFAULT_DEPTH_UNIT, DEPTH_METER_UNIT, DepthEncoderConfig from lerobot.streaming.episode_cache import EpisodeByteCache +from lerobot.streaming.episode_parquet import EpisodeParquetReader from lerobot.streaming.episode_pool import ExactCoveragePool from lerobot.streaming.manifest import EpisodeVideoManifest from lerobot.utils.constants import HF_LEROBOT_HOME +from lerobot.utils.import_utils import get_safe_default_video_backend from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata from .depth_utils import MM_PER_METRE, dequantize_depth -from .episode_parquet import EpisodeParquetReader from .feature_utils import check_delta_timestamps, get_delta_indices, get_hf_features_from_features from .io_utils import hf_transform_to_torch from .streaming_sidecar import ( @@ -113,6 +114,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): shuffle: bool = True, return_uint8: bool = False, depth_output_unit: str = DEFAULT_DEPTH_UNIT, + video_backend: str | None = None, data_root: str | Path | None = None, episode_pool_size: int | None = None, prefetch_episodes: int = 8, @@ -140,6 +142,9 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): shuffle (bool, optional): Whether to shuffle the dataset across exhaustions. Defaults to True. depth_output_unit (str, optional): Physical unit depth maps are dequantized to ("m" or "mm"). Defaults to "mm". + video_backend (str | None, optional): Decoder backend for synthesized episode videos. + Defaults to the same platform-safe backend as map-style loading. If TorchCodec + rejects a synthesized MP4, the byte cache falls back to its bounded PyAV decoder. data_root (str | Path | None, optional): Dataset payload root. Supports local paths, ``hf://``, and fsspec URLs. episode_pool_size (int | None, optional): Number of complete episodes in the sampling pool. @@ -168,6 +173,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): self.max_num_shards = max_num_shards self._return_uint8 = return_uint8 self._depth_output_unit = depth_output_unit + self._video_backend = video_backend if video_backend is not None else get_safe_default_video_backend() + if self._video_backend == "video_reader": + self._video_backend = "pyav" + if self._video_backend not in {"torchcodec", "pyav"}: + raise ValueError(f"Unsupported video backend: {self._video_backend}") if buffer_size <= 0: raise ValueError("buffer_size must be positive") if max_num_shards <= 0: @@ -463,6 +473,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): workers=workers, range_backend=range_backend, max_open_decoders=decoder_limit, + video_backend=self._video_backend, + tolerance_s=self.tolerance_s, ) def _make_episode_item( diff --git a/src/lerobot/streaming/episode_cache.py b/src/lerobot/streaming/episode_cache.py index 8c54edbab..dccd3337e 100644 --- a/src/lerobot/streaming/episode_cache.py +++ b/src/lerobot/streaming/episode_cache.py @@ -10,17 +10,22 @@ from __future__ import annotations import io +import logging import threading import time from collections import OrderedDict +from collections.abc import Callable from concurrent.futures import Future, ThreadPoolExecutor from pathlib import Path +from types import SimpleNamespace from typing import Any from lerobot.streaming.manifest import EpisodeVideoManifest from lerobot.streaming.mp4 import Mp4SampleSlice, synthesize_mp4 from lerobot.streaming.range_fetch import make_range_fetcher +logger = logging.getLogger(__name__) + class EpisodeByteCache: def __init__( @@ -37,11 +42,19 @@ class EpisodeByteCache: native_http_subranges: int = 1, open_decoders: bool = True, max_open_decoders: int = 64, + video_backend: str = "torchcodec", + tolerance_s: float = 1e-4, ): if byte_budget <= 0: raise ValueError("byte_budget must be positive") if max_open_decoders <= 0: raise ValueError("max_open_decoders must be positive") + if video_backend == "video_reader": + video_backend = "pyav" + if video_backend not in {"torchcodec", "pyav"}: + raise ValueError(f"Unsupported video backend: {video_backend}") + if tolerance_s <= 0: + raise ValueError("tolerance_s must be positive") self.manifest = manifest self.fetcher = make_range_fetcher( data_root, @@ -55,11 +68,16 @@ class EpisodeByteCache: self.byte_budget = byte_budget self.open_decoders = open_decoders self.max_open_decoders = max_open_decoders + self.video_backend = video_backend + self.tolerance_s = tolerance_s self._pool = ThreadPoolExecutor(max_workers=workers) self._cache: OrderedDict[tuple[int, str], dict[str, Any]] = OrderedDict() self._decoders: OrderedDict[tuple[int, str], Any] = OrderedDict() self._futures: dict[tuple[int, str], Future[dict[str, Any]]] = {} self._retained_episodes: dict[int, int] = {} + self._decoder_fallback_count = 0 + self._fallback_decoders: set[tuple[int, str]] = set() + self._fallback_warning_emitted = False self._bytes = 0 self._lock = threading.Lock() self._timing_totals = { @@ -73,11 +91,15 @@ class EpisodeByteCache: def close(self) -> None: self._pool.shutdown(wait=True, cancel_futures=True) with self._lock: + decoders = list(self._decoders.values()) self._cache.clear() self._decoders.clear() self._futures.clear() self._retained_episodes.clear() + self._fallback_decoders.clear() self._bytes = 0 + for decoder in decoders: + _close_decoder(decoder) self.fetcher.close() def __enter__(self) -> EpisodeByteCache: @@ -113,6 +135,11 @@ class EpisodeByteCache: with self._lock: return len(self._decoders) + @property + def decoder_fallback_count(self) -> int: + with self._lock: + return self._decoder_fallback_count + def ensure_ready(self, episode_index: int) -> None: for camera_key in self.manifest.video_keys: self.get_bytes(episode_index, camera_key) @@ -147,31 +174,95 @@ class EpisodeByteCache: self._decoders.move_to_end(key) return decoder - decoder = open_video_decoder(io.BytesIO(entry["bytes"])) + decoder = self._open_decoder(key, entry["bytes"]) with self._lock: existing = self._decoders.get(key) if existing is not None: self._decoders.move_to_end(key) + _close_decoder(decoder) return existing self._decoders[key] = decoder while len(self._decoders) > self.max_open_decoders: - self._decoders.popitem(last=False) + evicted_key, evicted_decoder = self._decoders.popitem(last=False) + self._fallback_decoders.discard(evicted_key) + _close_decoder(evicted_decoder) return decoder + def _open_decoder(self, key: tuple[int, str], data: bytes) -> Any: + try: + if self.video_backend == "torchcodec": + return open_video_decoder(io.BytesIO(data)) + return open_video_decoder(io.BytesIO(data), backend=self.video_backend) + except Exception as primary_error: + if self.video_backend != "torchcodec": + raise + try: + decoder = open_video_decoder(io.BytesIO(data), backend="pyav") + except Exception as fallback_error: + raise RuntimeError( + "Both TorchCodec and PyAV rejected synthesized episode video " + f"{key}: TorchCodec error: {primary_error}" + ) from fallback_error + with self._lock: + self._decoder_fallback_count += 1 + self._fallback_decoders.add(key) + should_warn = not self._fallback_warning_emitted + self._fallback_warning_emitted = True + if should_warn: + logger.warning( + "TorchCodec rejected a synthesized episode MP4; using the bounded PyAV " + "decoder fallback for affected videos. First error: %s", + primary_error, + ) + else: + logger.debug("Using PyAV decoder fallback for synthesized episode video %s", key) + return decoder + def get_frames(self, episode_index: int, camera_key: str, timestamps: list[float]): + key = (episode_index, camera_key) span = self.manifest.lookup(episode_index, camera_key) local_ts = [ts - span.source_start_pts for ts in timestamps] - decoder = self.get_decoder(episode_index, camera_key) - metadata = decoder.metadata - fps = getattr(metadata, "average_fps", None) - if fps is None: - duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9) - fps = metadata.num_frames / duration - return decoder.get_frames_at(indices=[round(ts * fps) for ts in local_ts]).data + decoder, release = self._decoder_for_frames(episode_index, camera_key) + with self._lock: + uses_pyav_timestamps = self.video_backend == "pyav" and key not in self._fallback_decoders + try: + if uses_pyav_timestamps: + return decoder.get_frames_played_at(local_ts, tolerance_s=self.tolerance_s).data + metadata = decoder.metadata + fps = getattr(metadata, "average_fps", None) + if fps is None: + duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9) + fps = metadata.num_frames / duration + return decoder.get_frames_at(indices=[round(ts * fps) for ts in local_ts]).data + finally: + if release is not None: + release() + + def _decoder_for_frames( + self, episode_index: int, camera_key: str + ) -> tuple[Any, Callable[[], None] | None]: + key = (episode_index, camera_key) + while True: + decoder = self.get_decoder(episode_index, camera_key) + acquire = getattr(decoder, "acquire", None) + if acquire is None: + return decoder, None + try: + acquire() + except RuntimeError: + # The decoder was evicted between lookup and lease acquisition. Remove a stale + # cached reference if it raced with close, then retry with a fresh decoder. + with self._lock: + if self._decoders.get(key) is decoder: + self._decoders.pop(key) + self._fallback_decoders.discard(key) + continue + return decoder, decoder.release def timing_summary(self) -> dict[str, float]: with self._lock: summary = dict(self._timing_totals) + summary["decoder_fallbacks"] = float(self._decoder_fallback_count) fetcher_summary = getattr(self.fetcher, "timing_summary", None) if fetcher_summary is not None: summary.update(fetcher_summary()) @@ -236,7 +327,10 @@ class EpisodeByteCache: ) entry = self._cache.pop(key) self._bytes -= len(entry["bytes"]) - self._decoders.pop(key, None) + decoder = self._decoders.pop(key, None) + self._fallback_decoders.discard(key) + if decoder is not None: + _close_decoder(decoder) def _fetch_and_synthesize(self, episode_index: int, camera_key: str) -> dict[str, Any]: lookup_start = time.perf_counter() @@ -271,9 +365,140 @@ class EpisodeByteCache: return entry -def open_video_decoder(file_like_or_bytesio, frame_mappings=None): +class _PyAVVideoDecoder: + """Small seekable PyAV adapter matching the TorchCodec calls used by the byte cache.""" + + def __init__(self, file_like_or_bytesio: Any): + import av + + self._source = file_like_or_bytesio + self._container = av.open(file_like_or_bytesio) + self._stream = self._container.streams.video[0] + average_rate = self._stream.average_rate + if average_rate is None: + raise ValueError("PyAV video stream does not expose an average frame rate") + self._fps = float(average_rate) + duration = ( + float(self._stream.duration * self._stream.time_base) + if self._stream.duration is not None + else 0.0 + ) + self.metadata = SimpleNamespace( + average_fps=self._fps, + num_frames=int(self._stream.frames or round(duration * self._fps)), + begin_stream_seconds=0.0, + end_stream_seconds=duration, + ) + self._decode_lock = threading.Lock() + self._state_lock = threading.Lock() + self._users = 0 + self._close_requested = False + self._closed = False + + def acquire(self) -> None: + with self._state_lock: + if self._close_requested or self._closed: + raise RuntimeError("PyAV decoder is closing") + self._users += 1 + + def release(self) -> None: + with self._state_lock: + self._users -= 1 + if self._users < 0: + raise RuntimeError("Unbalanced PyAV decoder release") + if self._users == 0 and self._close_requested: + self._close_resources() + + def get_frames_at(self, *, indices: list[int]) -> SimpleNamespace: + if not indices: + import torch + + return SimpleNamespace(data=torch.empty((0, 3, 0, 0), dtype=torch.uint8)) + timestamps = [index / self._fps for index in indices] + return self._get_frames_played_at(timestamps, tolerance_s=0.5 / self._fps + 1e-6) + + def get_frames_played_at( + self, + timestamps: list[float], + *, + tolerance_s: float, + ) -> SimpleNamespace: + return self._get_frames_played_at(timestamps, tolerance_s=tolerance_s) + + def _get_frames_played_at( + self, + timestamps: list[float], + *, + tolerance_s: float, + ) -> SimpleNamespace: + import torch + + first_ts = min(timestamps) + last_ts = max(timestamps) + loaded_frames: list[torch.Tensor] = [] + loaded_ts: list[float] = [] + with self._decode_lock: + self._container.seek( + round(first_ts / self._stream.time_base) - 1, + backward=True, + any_frame=False, + stream=self._stream, + ) + for frame in self._container.decode(self._stream): + if frame.pts is None: + continue + current_ts = float(frame.pts * self._stream.time_base) + array = frame.to_ndarray(format="rgb24") + loaded_frames.append(torch.from_numpy(array).permute(2, 0, 1).contiguous()) + loaded_ts.append(current_ts) + if current_ts >= last_ts: + break + + if not loaded_frames: + raise ValueError(f"PyAV decoded no frames for timestamps {timestamps}") + query_ts = torch.tensor(timestamps) + loaded_ts_tensor = torch.tensor(loaded_ts) + distances = torch.cdist(query_ts[:, None], loaded_ts_tensor[:, None], p=1) + minimum, closest = distances.min(1) + if not (minimum <= tolerance_s).all(): + raise ValueError( + f"PyAV frame timestamps exceed tolerance {tolerance_s}: " + f"queries={query_ts}, decoded={loaded_ts_tensor}" + ) + return SimpleNamespace(data=torch.stack([loaded_frames[index] for index in closest])) + + def close(self) -> None: + with self._state_lock: + self._close_requested = True + if self._users == 0: + self._close_resources() + + def _close_resources(self) -> None: + if self._closed: + return + self._container.close() + close = getattr(self._source, "close", None) + if close is not None: + close() + self._closed = True + + +def _close_decoder(decoder: Any) -> None: + close = getattr(decoder, "close", None) + if close is not None: + try: + close() + except Exception: + logger.debug("Failed to close video decoder", exc_info=True) + + +def open_video_decoder(file_like_or_bytesio, frame_mappings=None, *, backend: str = "torchcodec"): if frame_mappings is not None: raise ValueError("Synthesized episode videos use a local timeline; pass frame_mappings=None.") + if backend == "pyav": + return _PyAVVideoDecoder(file_like_or_bytesio) + if backend != "torchcodec": + raise ValueError(f"Unsupported video backend: {backend}") from torchcodec.decoders import VideoDecoder return VideoDecoder(file_like_or_bytesio, seek_mode="approximate") diff --git a/src/lerobot/datasets/episode_parquet.py b/src/lerobot/streaming/episode_parquet.py similarity index 98% rename from src/lerobot/datasets/episode_parquet.py rename to src/lerobot/streaming/episode_parquet.py index fc24798ef..09616c57e 100644 --- a/src/lerobot/datasets/episode_parquet.py +++ b/src/lerobot/streaming/episode_parquet.py @@ -6,7 +6,7 @@ # # http://www.apache.org/licenses/LICENSE-2.0 -"""Episode-scoped Parquet reads for training-time dataset streaming.""" +"""Pure episode-scoped Parquet reads for training-time dataset streaming.""" from __future__ import annotations diff --git a/tests/datasets/test_episode_parquet_reader.py b/tests/datasets/test_episode_parquet_reader.py index 80ab1ffef..ee6a4ae9e 100644 --- a/tests/datasets/test_episode_parquet_reader.py +++ b/tests/datasets/test_episode_parquet_reader.py @@ -20,7 +20,7 @@ pytest.importorskip("pyarrow", reason="pyarrow is required (install lerobot[data import pyarrow as pa import pyarrow.parquet as pq -from lerobot.datasets.episode_parquet import EpisodeParquetReader +from lerobot.streaming.episode_parquet import EpisodeParquetReader def _table(episodes: list[int]) -> pa.Table: diff --git a/tests/datasets/test_episode_video_streaming.py b/tests/datasets/test_episode_video_streaming.py index d2a7b1fec..38e5ef056 100644 --- a/tests/datasets/test_episode_video_streaming.py +++ b/tests/datasets/test_episode_video_streaming.py @@ -103,7 +103,14 @@ def test_parser_accepts_co64_chunk_offsets(): np.testing.assert_array_equal(mp4.sample_offsets, np.array([10_000, 10_050, 10_025])) -def _fake_cache(monkeypatch, tmp_path, *, byte_budget=8, max_open_decoders=1): +def _fake_cache( + monkeypatch, + tmp_path, + *, + byte_budget=8, + max_open_decoders=1, + video_backend="torchcodec", +): manifest = EpisodeVideoManifest(video_keys=["camera"], files=[], spans={}) cache = EpisodeByteCache( manifest, @@ -112,6 +119,7 @@ def _fake_cache(monkeypatch, tmp_path, *, byte_budget=8, max_open_decoders=1): workers=1, open_decoders=False, max_open_decoders=max_open_decoders, + video_backend=video_backend, ) monkeypatch.setattr( cache, @@ -161,6 +169,54 @@ def test_decoder_count_has_independent_limit(monkeypatch, tmp_path): assert cache.open_decoder_count == 1 +def test_decoder_eviction_and_cache_shutdown_close_backend_resources(monkeypatch, tmp_path): + opened = [] + + class FakeDecoder: + def __init__(self): + self.closed = False + + def close(self): + self.closed = True + + def open_decoder(_data): + decoder = FakeDecoder() + opened.append(decoder) + return decoder + + monkeypatch.setattr("lerobot.streaming.episode_cache.open_video_decoder", open_decoder) + with _fake_cache(monkeypatch, tmp_path, byte_budget=20, max_open_decoders=1) as cache: + cache.get_decoder(0, "camera") + cache.get_decoder(1, "camera") + + assert opened[0].closed + assert not opened[1].closed + + assert opened[1].closed + + +def test_decoder_falls_back_to_pyav_when_torchcodec_rejects_mini_mp4(monkeypatch, tmp_path): + opened_backends = [] + + class FakeDecoder: + pass + + def open_decoder(_data, frame_mappings=None, *, backend="torchcodec"): + assert frame_mappings is None + opened_backends.append(backend) + if backend == "torchcodec": + raise ValueError("No valid stream found") + return FakeDecoder() + + monkeypatch.setattr("lerobot.streaming.episode_cache.open_video_decoder", open_decoder) + with _fake_cache(monkeypatch, tmp_path, video_backend="torchcodec") as cache: + decoder = cache.get_decoder(0, "camera") + + assert isinstance(decoder, FakeDecoder) + assert opened_backends == ["torchcodec", "pyav"] + assert cache.decoder_fallback_count == 1 + + def test_releasing_episode_allows_immediate_eviction(monkeypatch, tmp_path): with _fake_cache(monkeypatch, tmp_path, byte_budget=5) as cache: cache.retain_episode(0) diff --git a/tests/datasets/test_streaming_factory.py b/tests/datasets/test_streaming_factory.py index 3126473d2..bb4d4c677 100644 --- a/tests/datasets/test_streaming_factory.py +++ b/tests/datasets/test_streaming_factory.py @@ -33,6 +33,7 @@ def test_factory_wires_production_streaming_settings(monkeypatch): dataset_config = DatasetConfig( repo_id="owner/dataset", streaming=True, + video_backend="pyav", streaming_data_root="memory://payload", streaming_episode_pool_size=7, streaming_prefetch_episodes=3, @@ -54,5 +55,6 @@ def test_factory_wires_production_streaming_settings(monkeypatch): assert captured["kwargs"]["prefetch_episodes"] == 3 assert captured["kwargs"]["byte_budget_gb"] == 2.5 assert captured["kwargs"]["max_num_shards"] == 1 + assert captured["kwargs"]["video_backend"] == "pyav" assert captured["kwargs"]["return_uint8"] is True assert captured["kwargs"]["repeat"] is True diff --git a/tests/datasets/test_streaming_production.py b/tests/datasets/test_streaming_production.py index 197e16916..2ee4a4338 100644 --- a/tests/datasets/test_streaming_production.py +++ b/tests/datasets/test_streaming_production.py @@ -61,19 +61,26 @@ def test_streaming_matches_map_style_with_exact_coverage(tmp_path: Path, lerobot _assert_item_equal(sample, map_dataset[int(sample["index"])]) -def test_streaming_rgb_video_matches_map_style(tmp_path: Path, lerobot_dataset_factory) -> None: +@pytest.mark.parametrize("video_backend", ["torchcodec", "pyav"]) +def test_streaming_rgb_video_matches_map_style( + tmp_path: Path, + lerobot_dataset_factory, + video_backend: str, +) -> None: root = tmp_path / "dataset" map_dataset = lerobot_dataset_factory( root=root, repo_id=DUMMY_REPO_ID, total_episodes=2, total_frames=20, + video_backend=video_backend, ) streaming = StreamingLeRobotDataset( DUMMY_REPO_ID, root=root, shuffle=False, buffer_size=2, + video_backend=video_backend, ) for sample in streaming: