mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
Fix streaming Parquet imports and video fallback
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
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]):
|
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)
|
||||||
metadata = decoder.metadata
|
with self._lock:
|
||||||
fps = getattr(metadata, "average_fps", None)
|
uses_pyav_timestamps = self.video_backend == "pyav" and key not in self._fallback_decoders
|
||||||
if fps is None:
|
try:
|
||||||
duration = max(getattr(metadata, "end_stream_seconds", 0.0), 1e-9)
|
if uses_pyav_timestamps:
|
||||||
fps = metadata.num_frames / duration
|
return decoder.get_frames_played_at(local_ts, tolerance_s=self.tolerance_s).data
|
||||||
return decoder.get_frames_at(indices=[round(ts * fps) for ts in local_ts]).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]:
|
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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user