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,
|
||||
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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user