Fix streaming Parquet imports and video fallback

This commit is contained in:
Pepijn
2026-07-24 16:41:54 +02:00
parent a3edab661b
commit 1e7e0b6de5
8 changed files with 320 additions and 16 deletions
+2
View File
@@ -114,6 +114,7 @@ def make_dataset(
tolerance_s=cfg.tolerance_s, tolerance_s=cfg.tolerance_s,
return_uint8=True, return_uint8=True,
depth_output_unit=cfg.dataset.depth_output_unit, depth_output_unit=cfg.dataset.depth_output_unit,
video_backend=cfg.dataset.video_backend,
data_root=cfg.dataset.streaming_data_root, data_root=cfg.dataset.streaming_data_root,
episode_pool_size=cfg.dataset.streaming_episode_pool_size, episode_pool_size=cfg.dataset.streaming_episode_pool_size,
prefetch_episodes=cfg.dataset.streaming_prefetch_episodes, prefetch_episodes=cfg.dataset.streaming_prefetch_episodes,
@@ -204,6 +205,7 @@ def make_train_eval_datasets(
tolerance_s=cfg.tolerance_s, tolerance_s=cfg.tolerance_s,
return_uint8=True, return_uint8=True,
depth_output_unit=cfg.dataset.depth_output_unit, depth_output_unit=cfg.dataset.depth_output_unit,
video_backend=cfg.dataset.video_backend,
data_root=cfg.dataset.streaming_data_root, data_root=cfg.dataset.streaming_data_root,
episode_pool_size=cfg.dataset.streaming_episode_pool_size, episode_pool_size=cfg.dataset.streaming_episode_pool_size,
prefetch_episodes=cfg.dataset.streaming_prefetch_episodes, prefetch_episodes=cfg.dataset.streaming_prefetch_episodes,
+13 -1
View File
@@ -25,13 +25,14 @@ import torch
from lerobot.configs import DEFAULT_DEPTH_UNIT, DEPTH_METER_UNIT, DepthEncoderConfig from lerobot.configs import DEFAULT_DEPTH_UNIT, DEPTH_METER_UNIT, DepthEncoderConfig
from lerobot.streaming.episode_cache import EpisodeByteCache from lerobot.streaming.episode_cache import EpisodeByteCache
from lerobot.streaming.episode_parquet import EpisodeParquetReader
from lerobot.streaming.episode_pool import ExactCoveragePool from lerobot.streaming.episode_pool import ExactCoveragePool
from lerobot.streaming.manifest import EpisodeVideoManifest from lerobot.streaming.manifest import EpisodeVideoManifest
from lerobot.utils.constants import HF_LEROBOT_HOME 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 .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata
from .depth_utils import MM_PER_METRE, dequantize_depth 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 .feature_utils import check_delta_timestamps, get_delta_indices, get_hf_features_from_features
from .io_utils import hf_transform_to_torch from .io_utils import hf_transform_to_torch
from .streaming_sidecar import ( from .streaming_sidecar import (
@@ -113,6 +114,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
shuffle: bool = True, shuffle: bool = True,
return_uint8: bool = False, return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT, depth_output_unit: str = DEFAULT_DEPTH_UNIT,
video_backend: str | None = None,
data_root: str | Path | None = None, data_root: str | Path | None = None,
episode_pool_size: int | None = None, episode_pool_size: int | None = None,
prefetch_episodes: int = 8, 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. 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"). depth_output_unit (str, optional): Physical unit depth maps are dequantized to ("m" or "mm").
Defaults to "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://``, data_root (str | Path | None, optional): Dataset payload root. Supports local paths, ``hf://``,
and fsspec URLs. and fsspec URLs.
episode_pool_size (int | None, optional): Number of complete episodes in the sampling pool. 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.max_num_shards = max_num_shards
self._return_uint8 = return_uint8 self._return_uint8 = return_uint8
self._depth_output_unit = depth_output_unit 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: if buffer_size <= 0:
raise ValueError("buffer_size must be positive") raise ValueError("buffer_size must be positive")
if max_num_shards <= 0: if max_num_shards <= 0:
@@ -463,6 +473,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
workers=workers, workers=workers,
range_backend=range_backend, range_backend=range_backend,
max_open_decoders=decoder_limit, max_open_decoders=decoder_limit,
video_backend=self._video_backend,
tolerance_s=self.tolerance_s,
) )
def _make_episode_item( def _make_episode_item(
+230 -5
View File
@@ -10,17 +10,22 @@
from __future__ import annotations from __future__ import annotations
import io import io
import logging
import threading import threading
import time import time
from collections import OrderedDict from collections import OrderedDict
from collections.abc import Callable
from concurrent.futures import Future, ThreadPoolExecutor from concurrent.futures import Future, ThreadPoolExecutor
from pathlib import Path from pathlib import Path
from types import SimpleNamespace
from typing import Any from typing import Any
from lerobot.streaming.manifest import EpisodeVideoManifest from lerobot.streaming.manifest import EpisodeVideoManifest
from lerobot.streaming.mp4 import Mp4SampleSlice, synthesize_mp4 from lerobot.streaming.mp4 import Mp4SampleSlice, synthesize_mp4
from lerobot.streaming.range_fetch import make_range_fetcher from lerobot.streaming.range_fetch import make_range_fetcher
logger = logging.getLogger(__name__)
class EpisodeByteCache: class EpisodeByteCache:
def __init__( def __init__(
@@ -37,11 +42,19 @@ class EpisodeByteCache:
native_http_subranges: int = 1, native_http_subranges: int = 1,
open_decoders: bool = True, open_decoders: bool = True,
max_open_decoders: int = 64, max_open_decoders: int = 64,
video_backend: str = "torchcodec",
tolerance_s: float = 1e-4,
): ):
if byte_budget <= 0: if byte_budget <= 0:
raise ValueError("byte_budget must be positive") raise ValueError("byte_budget must be positive")
if max_open_decoders <= 0: if max_open_decoders <= 0:
raise ValueError("max_open_decoders must be positive") 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.manifest = manifest
self.fetcher = make_range_fetcher( self.fetcher = make_range_fetcher(
data_root, data_root,
@@ -55,11 +68,16 @@ class EpisodeByteCache:
self.byte_budget = byte_budget self.byte_budget = byte_budget
self.open_decoders = open_decoders self.open_decoders = open_decoders
self.max_open_decoders = max_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._pool = ThreadPoolExecutor(max_workers=workers)
self._cache: OrderedDict[tuple[int, str], dict[str, Any]] = OrderedDict() self._cache: OrderedDict[tuple[int, str], dict[str, Any]] = OrderedDict()
self._decoders: OrderedDict[tuple[int, str], Any] = OrderedDict() self._decoders: OrderedDict[tuple[int, str], Any] = OrderedDict()
self._futures: dict[tuple[int, str], Future[dict[str, Any]]] = {} self._futures: dict[tuple[int, str], Future[dict[str, Any]]] = {}
self._retained_episodes: dict[int, int] = {} 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._bytes = 0
self._lock = threading.Lock() self._lock = threading.Lock()
self._timing_totals = { self._timing_totals = {
@@ -73,11 +91,15 @@ class EpisodeByteCache:
def close(self) -> None: def close(self) -> None:
self._pool.shutdown(wait=True, cancel_futures=True) self._pool.shutdown(wait=True, cancel_futures=True)
with self._lock: with self._lock:
decoders = list(self._decoders.values())
self._cache.clear() self._cache.clear()
self._decoders.clear() self._decoders.clear()
self._futures.clear() self._futures.clear()
self._retained_episodes.clear() self._retained_episodes.clear()
self._fallback_decoders.clear()
self._bytes = 0 self._bytes = 0
for decoder in decoders:
_close_decoder(decoder)
self.fetcher.close() self.fetcher.close()
def __enter__(self) -> EpisodeByteCache: def __enter__(self) -> EpisodeByteCache:
@@ -113,6 +135,11 @@ class EpisodeByteCache:
with self._lock: with self._lock:
return len(self._decoders) 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: def ensure_ready(self, episode_index: int) -> None:
for camera_key in self.manifest.video_keys: for camera_key in self.manifest.video_keys:
self.get_bytes(episode_index, camera_key) self.get_bytes(episode_index, camera_key)
@@ -147,31 +174,95 @@ class EpisodeByteCache:
self._decoders.move_to_end(key) self._decoders.move_to_end(key)
return decoder return decoder
decoder = open_video_decoder(io.BytesIO(entry["bytes"])) decoder = self._open_decoder(key, entry["bytes"])
with self._lock: with self._lock:
existing = self._decoders.get(key) existing = self._decoders.get(key)
if existing is not None: if existing is not None:
self._decoders.move_to_end(key) self._decoders.move_to_end(key)
_close_decoder(decoder)
return existing return existing
self._decoders[key] = decoder self._decoders[key] = decoder
while len(self._decoders) > self.max_open_decoders: 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 return decoder
def get_frames(self, episode_index: int, camera_key: str, timestamps: list[float]): 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) span = self.manifest.lookup(episode_index, camera_key)
local_ts = [ts - span.source_start_pts for ts in timestamps] local_ts = [ts - span.source_start_pts for ts in timestamps]
decoder = self.get_decoder(episode_index, camera_key) 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 metadata = decoder.metadata
fps = getattr(metadata, "average_fps", None) fps = getattr(metadata, "average_fps", None)
if fps is None: if fps is None:
duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9) duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9)
fps = metadata.num_frames / duration fps = metadata.num_frames / duration
return decoder.get_frames_at(indices=[round(ts * fps) for ts in local_ts]).data 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]: def timing_summary(self) -> dict[str, float]:
with self._lock: with self._lock:
summary = dict(self._timing_totals) summary = dict(self._timing_totals)
summary["decoder_fallbacks"] = float(self._decoder_fallback_count)
fetcher_summary = getattr(self.fetcher, "timing_summary", None) fetcher_summary = getattr(self.fetcher, "timing_summary", None)
if fetcher_summary is not None: if fetcher_summary is not None:
summary.update(fetcher_summary()) summary.update(fetcher_summary())
@@ -236,7 +327,10 @@ class EpisodeByteCache:
) )
entry = self._cache.pop(key) entry = self._cache.pop(key)
self._bytes -= len(entry["bytes"]) 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]: def _fetch_and_synthesize(self, episode_index: int, camera_key: str) -> dict[str, Any]:
lookup_start = time.perf_counter() lookup_start = time.perf_counter()
@@ -271,9 +365,140 @@ class EpisodeByteCache:
return entry 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: if frame_mappings is not None:
raise ValueError("Synthesized episode videos use a local timeline; pass frame_mappings=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 from torchcodec.decoders import VideoDecoder
return VideoDecoder(file_like_or_bytesio, seek_mode="approximate") return VideoDecoder(file_like_or_bytesio, seek_mode="approximate")
@@ -6,7 +6,7 @@
# #
# http://www.apache.org/licenses/LICENSE-2.0 # 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 from __future__ import annotations
@@ -20,7 +20,7 @@ pytest.importorskip("pyarrow", reason="pyarrow is required (install lerobot[data
import pyarrow as pa import pyarrow as pa
import pyarrow.parquet as pq 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: def _table(episodes: list[int]) -> pa.Table:
+57 -1
View File
@@ -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])) 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={}) manifest = EpisodeVideoManifest(video_keys=["camera"], files=[], spans={})
cache = EpisodeByteCache( cache = EpisodeByteCache(
manifest, manifest,
@@ -112,6 +119,7 @@ def _fake_cache(monkeypatch, tmp_path, *, byte_budget=8, max_open_decoders=1):
workers=1, workers=1,
open_decoders=False, open_decoders=False,
max_open_decoders=max_open_decoders, max_open_decoders=max_open_decoders,
video_backend=video_backend,
) )
monkeypatch.setattr( monkeypatch.setattr(
cache, cache,
@@ -161,6 +169,54 @@ def test_decoder_count_has_independent_limit(monkeypatch, tmp_path):
assert cache.open_decoder_count == 1 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): def test_releasing_episode_allows_immediate_eviction(monkeypatch, tmp_path):
with _fake_cache(monkeypatch, tmp_path, byte_budget=5) as cache: with _fake_cache(monkeypatch, tmp_path, byte_budget=5) as cache:
cache.retain_episode(0) cache.retain_episode(0)
+2
View File
@@ -33,6 +33,7 @@ def test_factory_wires_production_streaming_settings(monkeypatch):
dataset_config = DatasetConfig( dataset_config = DatasetConfig(
repo_id="owner/dataset", repo_id="owner/dataset",
streaming=True, streaming=True,
video_backend="pyav",
streaming_data_root="memory://payload", streaming_data_root="memory://payload",
streaming_episode_pool_size=7, streaming_episode_pool_size=7,
streaming_prefetch_episodes=3, 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"]["prefetch_episodes"] == 3
assert captured["kwargs"]["byte_budget_gb"] == 2.5 assert captured["kwargs"]["byte_budget_gb"] == 2.5
assert captured["kwargs"]["max_num_shards"] == 1 assert captured["kwargs"]["max_num_shards"] == 1
assert captured["kwargs"]["video_backend"] == "pyav"
assert captured["kwargs"]["return_uint8"] is True assert captured["kwargs"]["return_uint8"] is True
assert captured["kwargs"]["repeat"] is True assert captured["kwargs"]["repeat"] is True
+8 -1
View File
@@ -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"])]) _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" root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory( map_dataset = lerobot_dataset_factory(
root=root, root=root,
repo_id=DUMMY_REPO_ID, repo_id=DUMMY_REPO_ID,
total_episodes=2, total_episodes=2,
total_frames=20, total_frames=20,
video_backend=video_backend,
) )
streaming = StreamingLeRobotDataset( streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID, DUMMY_REPO_ID,
root=root, root=root,
shuffle=False, shuffle=False,
buffer_size=2, buffer_size=2,
video_backend=video_backend,
) )
for sample in streaming: for sample in streaming: