Preserve private Hub auth in streaming readers

This commit is contained in:
Pepijn
2026-07-27 09:41:57 +02:00
parent b5a13e43ce
commit c0e9b0bbff
8 changed files with 74 additions and 10 deletions
+9 -1
View File
@@ -180,6 +180,7 @@ 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._streaming_io_token = None if self.streaming_from_local else token
self._video_backend = video_backend if video_backend is not None else get_safe_default_video_backend() self._video_backend = video_backend if video_backend is not None else get_safe_default_video_backend()
if self._video_backend == "video_reader": if self._video_backend == "video_reader":
self._video_backend = "pyav" self._video_backend = "pyav"
@@ -266,6 +267,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self._data_root, self._data_root,
workers=max_num_shards, workers=max_num_shards,
range_backend=sidecar_backend, range_backend=sidecar_backend,
token=self._streaming_io_token,
) )
self._hf_features = get_hf_features_from_features(self.meta.features) self._hf_features = get_hf_features_from_features(self.meta.features)
self._projected_columns = tuple(self._hf_features) self._projected_columns = tuple(self._hf_features)
@@ -349,7 +351,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
planner.newly_admitted.clear() planner.newly_admitted.clear()
planner.evicted.clear() planner.evicted.clear()
parquet_reader = EpisodeParquetReader(self._data_root, columns=self._projected_columns) parquet_reader = EpisodeParquetReader(
self._data_root,
columns=self._projected_columns,
token=self._streaming_io_token,
)
executor = ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix="lerobot-parquet") executor = ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix="lerobot-parquet")
episode_futures: dict[int, Future[datasets.Dataset]] = {} episode_futures: dict[int, Future[datasets.Dataset]] = {}
scheduled_episodes: set[int] = set() scheduled_episodes: set[int] = set()
@@ -475,6 +481,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
range_backend=range_backend, range_backend=range_backend,
workers=workers, workers=workers,
sidecar_path=self._sidecar_path, sidecar_path=self._sidecar_path,
token=self._streaming_io_token,
) )
decoder_limit = max(1, min(64, self.episode_pool_size * max(1, len(self.meta.video_keys)))) decoder_limit = max(1, min(64, self.episode_pool_size * max(1, len(self.meta.video_keys))))
return EpisodeByteCache( return EpisodeByteCache(
@@ -486,6 +493,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
max_open_decoders=decoder_limit, max_open_decoders=decoder_limit,
video_backend=self._video_backend, video_backend=self._video_backend,
tolerance_s=self.tolerance_s, tolerance_s=self.tolerance_s,
token=self._streaming_io_token,
) )
def _make_episode_item( def _make_episode_item(
+9 -1
View File
@@ -71,6 +71,7 @@ def build_mp4_sidecar(
workers: int = 8, workers: int = 8,
range_backend: str = "native-http", range_backend: str = "native-http",
max_probe_bytes: int = 64 * 1024 * 1024, max_probe_bytes: int = 64 * 1024 * 1024,
token: str | bool | None = None,
) -> None: ) -> None:
EpisodeVideoManifest.write_file_sidecar( EpisodeVideoManifest.write_file_sidecar(
destination, destination,
@@ -80,6 +81,7 @@ def build_mp4_sidecar(
range_backend=range_backend, range_backend=range_backend,
workers=workers, workers=workers,
max_probe_bytes=max_probe_bytes, max_probe_bytes=max_probe_bytes,
token=token,
) )
@@ -93,10 +95,13 @@ def download_published_sidecar(
spec: SidecarSpec, spec: SidecarSpec,
*, *,
cache_root: str | Path = DEFAULT_SIDECAR_CACHE, cache_root: str | Path = DEFAULT_SIDECAR_CACHE,
token: str | bool | None = None,
) -> bool: ) -> bool:
if Path(spec.data_root).expanduser().is_dir(): if Path(spec.data_root).expanduser().is_dir():
return False return False
filesystem, source = fsspec.core.url_to_fs(published_sidecar_url(spec, cache_root)) source_url = published_sidecar_url(spec, cache_root)
storage_options = {"token": token} if token is not None and source_url.startswith("hf://") else {}
filesystem, source = fsspec.core.url_to_fs(source_url, **storage_options)
if not filesystem.exists(source): if not filesystem.exists(source):
return False return False
with filesystem.open(source, "rb") as remote, destination.open("wb") as local: with filesystem.open(source, "rb") as remote, destination.open("wb") as local:
@@ -112,6 +117,7 @@ def ensure_dataset_mp4_sidecar(
workers: int = 8, workers: int = 8,
range_backend: str = "native-http", range_backend: str = "native-http",
lock_timeout_s: float = 30 * 60, lock_timeout_s: float = 30 * 60,
token: str | bool | None = None,
) -> Path | None: ) -> Path | None:
if not meta.video_keys: if not meta.video_keys:
return None return None
@@ -131,11 +137,13 @@ def ensure_dataset_mp4_sidecar(
target_spec, target_spec,
workers=workers, workers=workers,
range_backend=range_backend, range_backend=range_backend,
token=token,
), ),
download=lambda path, target_spec: download_published_sidecar( download=lambda path, target_spec: download_published_sidecar(
path, path,
target_spec, target_spec,
cache_root=cache_root, cache_root=cache_root,
token=token,
), ),
lock_timeout_s=lock_timeout_s, lock_timeout_s=lock_timeout_s,
) )
+2
View File
@@ -44,6 +44,7 @@ class EpisodeByteCache:
max_open_decoders: int = 64, max_open_decoders: int = 64,
video_backend: str = "torchcodec", video_backend: str = "torchcodec",
tolerance_s: float = 1e-4, tolerance_s: float = 1e-4,
token: str | bool | None = None,
): ):
if byte_budget <= 0: if byte_budget <= 0:
raise ValueError("byte_budget must be positive") raise ValueError("byte_budget must be positive")
@@ -64,6 +65,7 @@ class EpisodeByteCache:
native_http_timeout=native_http_timeout, native_http_timeout=native_http_timeout,
native_http_retries=native_http_retries, native_http_retries=native_http_retries,
native_http_subranges=native_http_subranges, native_http_subranges=native_http_subranges,
token=token,
) )
self.byte_budget = byte_budget self.byte_budget = byte_budget
self.open_decoders = open_decoders self.open_decoders = open_decoders
+10 -2
View File
@@ -23,14 +23,22 @@ import pyarrow.parquet as pq
class EpisodeParquetReader: class EpisodeParquetReader:
"""Read complete episodes with column projection from local or fsspec roots.""" """Read complete episodes with column projection from local or fsspec roots."""
def __init__(self, data_root: str | Path, *, columns: Sequence[str]): def __init__(
self,
data_root: str | Path,
*,
columns: Sequence[str],
token: str | bool | None = None,
):
if not columns: if not columns:
raise ValueError("EpisodeParquetReader requires at least one projected column") raise ValueError("EpisodeParquetReader requires at least one projected column")
self.columns = tuple(dict.fromkeys(columns)) self.columns = tuple(dict.fromkeys(columns))
self._read_columns = ( self._read_columns = (
self.columns if "episode_index" in self.columns else (*self.columns, "episode_index") self.columns if "episode_index" in self.columns else (*self.columns, "episode_index")
) )
self._filesystem, self._root_path = fsspec.core.url_to_fs(str(data_root)) data_root_str = str(data_root)
storage_options = {"token": token} if token is not None and data_root_str.startswith("hf://") else {}
self._filesystem, self._root_path = fsspec.core.url_to_fs(data_root_str, **storage_options)
def read_episode( def read_episode(
self, self,
+11 -1
View File
@@ -82,6 +82,7 @@ class EpisodeVideoManifest:
keyframe_pad_s: float = 0.1, keyframe_pad_s: float = 0.1,
keyframe_pad_fraction: float = 0.05, keyframe_pad_fraction: float = 0.05,
sidecar_path: str | Path | None = None, sidecar_path: str | Path | None = None,
token: str | bool | None = None,
) -> EpisodeVideoManifest: ) -> EpisodeVideoManifest:
meta.ensure_readable() meta.ensure_readable()
video_keys = list(meta.video_keys) video_keys = list(meta.video_keys)
@@ -99,6 +100,7 @@ class EpisodeVideoManifest:
workers=workers, workers=workers,
header_probe_bytes=header_probe_bytes, header_probe_bytes=header_probe_bytes,
max_probe_bytes=max_probe_bytes, max_probe_bytes=max_probe_bytes,
token=token,
) )
else: else:
records = cls.load_file_sidecar(sidecar_path) records = cls.load_file_sidecar(sidecar_path)
@@ -159,8 +161,14 @@ class EpisodeVideoManifest:
workers: int, workers: int,
header_probe_bytes: int, header_probe_bytes: int,
max_probe_bytes: int, max_probe_bytes: int,
token: str | bool | None,
) -> list[VideoFileRecord]: ) -> list[VideoFileRecord]:
fetcher = make_range_fetcher(data_root, range_backend=range_backend, workers=workers) fetcher = make_range_fetcher(
data_root,
range_backend=range_backend,
workers=workers,
token=token,
)
def build_file(path: str) -> VideoFileRecord: def build_file(path: str) -> VideoFileRecord:
file_size = fetcher.info_size(path) file_size = fetcher.info_size(path)
@@ -198,6 +206,7 @@ class EpisodeVideoManifest:
workers: int = 8, workers: int = 8,
header_probe_bytes: int = 4 * 1024 * 1024, header_probe_bytes: int = 4 * 1024 * 1024,
max_probe_bytes: int = 64 * 1024 * 1024, max_probe_bytes: int = 64 * 1024 * 1024,
token: str | bool | None = None,
) -> None: ) -> None:
records = cls._build_file_records( records = cls._build_file_records(
sorted(set(rel_paths)), sorted(set(rel_paths)),
@@ -206,6 +215,7 @@ class EpisodeVideoManifest:
workers=workers, workers=workers,
header_probe_bytes=header_probe_bytes, header_probe_bytes=header_probe_bytes,
max_probe_bytes=max_probe_bytes, max_probe_bytes=max_probe_bytes,
token=token,
) )
cls.save_file_sidecar(sidecar_path, records, spec=spec) cls.save_file_sidecar(sidecar_path, records, spec=spec)
+16 -5
View File
@@ -106,9 +106,17 @@ def _log_http_failure(
class ThreadLocalRangeFetcher: class ThreadLocalRangeFetcher:
"""Range reader that gives each worker thread independent file handles.""" """Range reader that gives each worker thread independent file handles."""
def __init__(self, data_root: str | Path, *, block_size: int = 2**20, cache_type: str = "none"): def __init__(
self,
data_root: str | Path,
*,
block_size: int = 2**20,
cache_type: str = "none",
token: str | bool | None = None,
):
self.data_root = str(data_root).rstrip("/") self.data_root = str(data_root).rstrip("/")
self.fs, self._root_path = fsspec.core.url_to_fs(self.data_root) storage_options = {"token": token} if token is not None and self.data_root.startswith("hf://") else {}
self.fs, self._root_path = fsspec.core.url_to_fs(self.data_root, **storage_options)
self._is_local = self.fs.protocol in ("file", "local") or ( self._is_local = self.fs.protocol in ("file", "local") or (
isinstance(self.fs.protocol, tuple) and "file" in self.fs.protocol isinstance(self.fs.protocol, tuple) and "file" in self.fs.protocol
) )
@@ -350,6 +358,7 @@ class NativeHTTPRangeFetcher:
max_retries: int = 4, max_retries: int = 4,
subrange_parts: int = 1, subrange_parts: int = 1,
subrange_min_bytes: int = 8 * 1024 * 1024, subrange_min_bytes: int = 8 * 1024 * 1024,
token: str | bool | None = None,
): ):
self.data_root = str(data_root).rstrip("/") self.data_root = str(data_root).rstrip("/")
if not self.data_root.startswith("hf://"): if not self.data_root.startswith("hf://"):
@@ -366,7 +375,7 @@ class NativeHTTPRangeFetcher:
if self.subrange_parts > 1 if self.subrange_parts > 1
else None else None
) )
self.api = HfApi() self.api = HfApi(token=token)
self.fs: HfFileSystem | None = None self.fs: HfFileSystem | None = None
self._bucket_id: str | None = None self._bucket_id: str | None = None
self._bucket_prefix = "" self._bucket_prefix = ""
@@ -378,7 +387,7 @@ class NativeHTTPRangeFetcher:
self._bucket_id = f"{parts[0]}/{parts[1]}" self._bucket_id = f"{parts[0]}/{parts[1]}"
self._bucket_prefix = parts[2].strip("/") if len(parts) == 3 else "" self._bucket_prefix = parts[2].strip("/") if len(parts) == 3 else ""
else: else:
self.fs = HfFileSystem() self.fs = HfFileSystem(token=token)
self.client = httpx.Client( self.client = httpx.Client(
timeout=timeout, timeout=timeout,
limits=httpx.Limits(max_connections=max_connections, max_keepalive_connections=max_connections), limits=httpx.Limits(max_connections=max_connections, max_keepalive_connections=max_connections),
@@ -717,9 +726,10 @@ def make_range_fetcher(
native_http_timeout: float = 60.0, native_http_timeout: float = 60.0,
native_http_retries: int = 4, native_http_retries: int = 4,
native_http_subranges: int = 1, native_http_subranges: int = 1,
token: str | bool | None = None,
): ):
if range_backend == "fsspec": if range_backend == "fsspec":
return ThreadLocalRangeFetcher(data_root) return ThreadLocalRangeFetcher(data_root, token=token)
if range_backend == "native-http": if range_backend == "native-http":
max_connections = native_http_connections or max(8, workers) max_connections = native_http_connections or max(8, workers)
return NativeHTTPRangeFetcher( return NativeHTTPRangeFetcher(
@@ -728,5 +738,6 @@ def make_range_fetcher(
timeout=native_http_timeout, timeout=native_http_timeout,
max_retries=native_http_retries, max_retries=native_http_retries,
subrange_parts=native_http_subranges, subrange_parts=native_http_subranges,
token=token,
) )
raise ValueError(f"Unknown range backend: {range_backend}") raise ValueError(f"Unknown range backend: {range_backend}")
@@ -11,6 +11,7 @@
from __future__ import annotations from __future__ import annotations
from pathlib import Path from pathlib import Path
from unittest.mock import Mock
import fsspec import fsspec
import pytest import pytest
@@ -108,3 +109,16 @@ def test_reader_supports_fsspec_remote_root() -> None:
table = reader.read_episode("data/chunk-000/file-000.parquet", episode_index=0, expected_rows=3) table = reader.read_episode("data/chunk-000/file-000.parquet", episode_index=0, expected_rows=3)
assert table.column("value").to_pylist() == [0, 1, 2] assert table.column("value").to_pylist() == [0, 1, 2]
def test_reader_forwards_explicit_token_to_hf_filesystem(monkeypatch) -> None:
url_to_fs = Mock(return_value=(fsspec.filesystem("memory"), "datasets/private@revision"))
monkeypatch.setattr(fsspec.core, "url_to_fs", url_to_fs)
EpisodeParquetReader(
"hf://datasets/private@revision",
columns=("episode_index",),
token="hf_test_token",
)
url_to_fs.assert_called_once_with("hf://datasets/private@revision", token="hf_test_token")
+3
View File
@@ -46,7 +46,9 @@ def test_streaming_dataset_forwards_token_to_metadata_without_retaining_it(
rescale_depth_stats=Mock(), rescale_depth_stats=Mock(),
) )
metadata_cls = Mock(return_value=metadata) metadata_cls = Mock(return_value=metadata)
ensure_sidecar = Mock(return_value=None)
monkeypatch.setattr(streaming_dataset_module, "LeRobotDatasetMetadata", metadata_cls) monkeypatch.setattr(streaming_dataset_module, "LeRobotDatasetMetadata", metadata_cls)
monkeypatch.setattr(streaming_dataset_module, "ensure_dataset_mp4_sidecar", ensure_sidecar)
dataset = StreamingLeRobotDataset(DUMMY_REPO_ID, root=requested_root, token=token) dataset = StreamingLeRobotDataset(DUMMY_REPO_ID, root=requested_root, token=token)
@@ -57,6 +59,7 @@ def test_streaming_dataset_forwards_token_to_metadata_without_retaining_it(
force_cache_sync=False, force_cache_sync=False,
token=token, token=token,
) )
assert ensure_sidecar.call_args.kwargs["token"] is (None if from_local else token)
assert not hasattr(dataset, "_token") assert not hasattr(dataset, "_token")