feat(dataset): Support streaming from HF Storage Buckets + Bump HF hub & datasets (#4236)

This commit is contained in:
Steven Palma
2026-07-31 01:18:51 +02:00
committed by GitHub
parent 1fe58f2d3a
commit 0d0737ab57
5 changed files with 255 additions and 36 deletions
+2 -2
View File
@@ -68,7 +68,7 @@ dependencies = [
# Config & Hub
"draccus>=0.11.6,<0.12.0",
"huggingface-hub>=1.0.0,<2.0.0",
"huggingface-hub>=1.6.0,<2.0.0",
"requests>=2.32.0,<3.0.0",
# Environments
@@ -95,7 +95,7 @@ dependencies = [
# ── Feature-scoped extras ──────────────────────────────────
dataset = [
"datasets>=4.7.0,<5.0.0",
"datasets>=4.8.0,<5.0.0",
"pandas>=2.0.0,<3.0.0", # NOTE: Transitive dependency of datasets
"pyarrow>=21.0.0,<30.0.0", # NOTE: Transitive dependency of datasets
"lerobot[av-dep]",
+48 -16
View File
@@ -18,13 +18,15 @@ import logging
from collections.abc import Callable, Iterable
from copy import deepcopy
from pathlib import Path
from typing import Literal
import numpy as np
import packaging.version
import pandas as pd
import pyarrow as pa
import pyarrow.parquet as pq
from huggingface_hub import snapshot_download
from huggingface_hub import snapshot_download, sync_bucket
from huggingface_hub.utils import WeakFileLock
from lerobot.configs import DEPTH_METER_UNIT, VideoEncoderConfig
from lerobot.utils.constants import DEFAULT_FEATURES, HF_LEROBOT_HOME, HF_LEROBOT_HUB_CACHE
@@ -74,6 +76,7 @@ class LeRobotDatasetMetadata:
force_cache_sync: bool = False,
metadata_buffer_size: int = 10,
*,
repo_type: Literal["dataset", "bucket"] = "dataset",
token: str | bool | None = None,
):
"""Load or download metadata for an existing LeRobot dataset.
@@ -96,36 +99,53 @@ class LeRobotDatasetMetadata:
even when local files exist.
metadata_buffer_size: Number of episode metadata records to buffer
in memory before flushing to parquet.
repo_type: Repository type: "dataset" (default) or "bucket" for an
HF Storage Bucket streamed over hf://buckets/.
token: Authentication token used for Hub requests. Pass a string
token, ``True`` to require the locally stored token, ``False``
to disable authentication, or ``None`` to use the Hugging Face
Hub default.
"""
if repo_type not in ("dataset", "bucket"):
raise ValueError(f"repo_type must be 'dataset' or 'bucket', got {repo_type!r}")
self.repo_id = repo_id
self.repo_type = repo_type
self.revision = revision if revision else CODEBASE_VERSION
self._requested_root = Path(root) if root is not None else None
self.root = self._requested_root if self._requested_root is not None else HF_LEROBOT_HOME / repo_id
if self._requested_root is not None:
self.root = self._requested_root
elif self.repo_type == "bucket":
self.root = HF_LEROBOT_HUB_CACHE / ("buckets--" + self.repo_id.replace("/", "--"))
else:
self.root = HF_LEROBOT_HOME / repo_id
self._pq_writer = None
self.latest_episode = None
self._metadata_buffer: list[dict] = []
self._metadata_buffer_size = metadata_buffer_size
self._finalized = False
try:
if force_cache_sync or (
self._requested_root is None and has_legacy_hub_download_metadata(self.root)
):
raise FileNotFoundError
self._load_metadata()
except (FileNotFoundError, NotADirectoryError):
if is_valid_version(self.revision):
if token is None:
self.revision = get_safe_version(self.repo_id, self.revision)
else:
self.revision = get_safe_version(self.repo_id, self.revision, token=token)
metadata_lock = contextlib.nullcontext()
if self.repo_type == "bucket":
self.root.parent.mkdir(parents=True, exist_ok=True)
metadata_lock = WeakFileLock(self.root.parent / f".{self.root.name}.lock")
self._pull_from_repo(allow_patterns="meta/", token=token)
self._load_metadata()
with metadata_lock:
try:
if force_cache_sync or (
self._requested_root is None and has_legacy_hub_download_metadata(self.root)
):
raise FileNotFoundError
self._load_metadata()
except (FileNotFoundError, NotADirectoryError):
if self.repo_type != "bucket" and is_valid_version(self.revision):
if token is None:
self.revision = get_safe_version(self.repo_id, self.revision)
else:
self.revision = get_safe_version(self.repo_id, self.revision, token=token)
self._pull_from_repo(allow_patterns="meta/", token=token)
self._load_metadata()
def _flush_metadata_buffer(self) -> None:
"""Write all buffered episode metadata to parquet file."""
@@ -232,6 +252,16 @@ class LeRobotDatasetMetadata:
*,
token: str | bool | None = None,
) -> None:
if self.repo_type == "bucket":
self.root.mkdir(parents=True, exist_ok=True)
sync_bucket(
f"hf://buckets/{self.repo_id}/meta",
str(self.root / "meta"),
delete=True,
quiet=True,
token=token,
)
return
token_kwargs = {} if token is None else {"token": token}
if self._requested_root is None:
self.root = Path(
@@ -262,6 +292,8 @@ class LeRobotDatasetMetadata:
@property
def url_root(self) -> str:
"""Hugging Face Hub URL root for this dataset."""
if self.repo_type == "bucket":
return f"hf://buckets/{self.repo_id}"
return f"hf://datasets/{self.repo_id}"
@property
+34 -14
View File
@@ -16,6 +16,7 @@
from collections import deque
from collections.abc import Callable, Generator, Iterable, Iterator
from pathlib import Path
from typing import Literal
import datasets
import numpy as np
@@ -261,15 +262,16 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
*,
repo_type: Literal["dataset", "bucket"] = "dataset",
token: str | bool | None = None,
):
"""Initialize a StreamingLeRobotDataset.
Args:
repo_id (str): This is the repo id that will be used to fetch the dataset.
root (Path | None, optional): Local directory to use for local datasets. When omitted, Hub
metadata is resolved through a revision-safe snapshot cache under
``$HF_LEROBOT_HOME/hub``.
root (Path | None, optional): Local directory to use for local datasets. In bucket mode,
this is an optional local metadata-cache directory; parquet and video data remain remote.
When omitted, Hub metadata is resolved through the cache under ``$HF_LEROBOT_HOME/hub``.
episodes (list[int] | None, optional): If specified, this will only load episodes specified by
their episode_index in this list.
image_transforms (Callable | None, optional): Transform to apply to image data.
@@ -284,6 +286,8 @@ 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".
repo_type: "dataset" (default) or "bucket" to stream from an HF Storage Bucket
over ``hf://buckets/``.
token: Authentication token used while streaming this dataset from
the Hub. Pass a string token, ``True`` to require the locally
stored token, ``False`` to disable authentication, or ``None``
@@ -291,10 +295,14 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
on the dataset instance after initialization.
"""
super().__init__()
if repo_type not in ("dataset", "bucket"):
raise ValueError(f"repo_type must be 'dataset' or 'bucket', got {repo_type!r}")
self.repo_id = repo_id
self._requested_root = Path(root) if root else None
self.repo_type = repo_type
self._requested_root = Path(root) if root is not None else None
self.root = self._requested_root if self._requested_root is not None else HF_LEROBOT_HOME / repo_id
self.streaming_from_local = root is not None
self.streaming_from_local = root is not None and self.repo_type == "dataset"
self.image_transforms = image_transforms
self.episodes = episodes
@@ -321,6 +329,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self._requested_root,
self.revision,
force_cache_sync=force_cache_sync,
repo_type=self.repo_type,
token=token,
)
self.root = self.meta.root
@@ -349,15 +358,26 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self.delta_timestamps = delta_timestamps
self.delta_indices = get_delta_indices(self.delta_timestamps, self.fps)
token_kwargs = {} if token is None or self.streaming_from_local else {"token": token}
self.hf_dataset: datasets.IterableDataset = load_dataset(
self.repo_id if not self.streaming_from_local else str(self.root),
split="train",
streaming=self.streaming,
data_files="data/*/*.parquet",
revision=self.revision,
**token_kwargs,
)
token_kwargs = {} if token is None else {"token": token}
if self.repo_type == "bucket":
self.hf_dataset: datasets.IterableDataset = load_dataset(
"parquet",
data_files=f"hf://buckets/{self.repo_id}/data/*/*.parquet",
split="train",
streaming=self.streaming,
**token_kwargs,
)
else:
if self.streaming_from_local:
token_kwargs = {}
self.hf_dataset: datasets.IterableDataset = load_dataset(
self.repo_id if not self.streaming_from_local else str(self.root),
split="train",
streaming=self.streaming,
data_files="data/*/*.parquet",
revision=self.revision,
**token_kwargs,
)
self.num_shards = min(self.hf_dataset.num_shards, max_num_shards)
+168 -1
View File
@@ -14,7 +14,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from types import SimpleNamespace
from unittest.mock import Mock
from unittest.mock import Mock, patch
import numpy as np
import pytest
@@ -23,6 +23,7 @@ import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
import lerobot.datasets.streaming_dataset as streaming_dataset_module
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
from lerobot.datasets.utils import safe_shard
from lerobot.utils.constants import ACTION
@@ -100,6 +101,7 @@ def test_streaming_dataset_forwards_hub_token_only_for_remote_data(tmp_path, mon
requested_root,
streaming_dataset_module.CODEBASE_VERSION,
force_cache_sync=False,
repo_type="dataset",
token=token,
)
if from_local:
@@ -478,3 +480,168 @@ def test_frames_with_delta_consistency_with_shards(
assert all(t[1] for t in key_checks), (
f"Checking {list(filter(lambda t: not t[1], key_checks))[0][0]} left and right were found different (i: {i}, frame_idx: {frame_idx})"
)
class _StopConstructionError(Exception):
"""Sentinel raised from a patched load_dataset to halt __init__ after the branch under test."""
def _fake_meta(*args, **kwargs):
"""Minimal LeRobotDatasetMetadata stand-in exposing only what __init__ reads."""
meta = type("_Meta", (), {})()
root = kwargs.get("root", args[1] if len(args) > 1 else None)
revision = kwargs.get("revision", args[2] if len(args) > 2 else None)
meta.root = root or "/tmp/_streaming_meta"
meta.revision = revision or "v0"
meta._version = "v3.0"
meta.depth_keys = []
meta.image_keys = []
meta.rescale_depth_stats = lambda *_a, **_k: None
meta.repo_type = kwargs.get("repo_type", "dataset")
return meta
@pytest.mark.parametrize(
"repo_type, expected_source, expected_data_files",
[
("bucket", "parquet", "hf://buckets/{repo_id}/data/*/*.parquet"),
("dataset", "{repo_id}", "data/*/*.parquet"),
],
)
def test_streaming_repo_type_routes_load_dataset(repo_type, expected_source, expected_data_files):
"""repo_type='bucket' loads parquet from hf://buckets/...; 'dataset' keeps the Hub-repo path."""
captured = {}
token = "hf_test_token"
def fake_load_dataset(source, **kwargs):
captured["source"] = source
captured["data_files"] = kwargs.get("data_files")
captured["token"] = kwargs.get("token")
raise _StopConstructionError
with (
patch("lerobot.datasets.streaming_dataset.LeRobotDatasetMetadata", _fake_meta),
patch("lerobot.datasets.streaming_dataset.check_version_compatibility", lambda *a, **k: None),
patch("lerobot.datasets.streaming_dataset.load_dataset", fake_load_dataset),
pytest.raises(_StopConstructionError),
):
StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type=repo_type, token=token)
assert captured["source"] == expected_source.format(repo_id=DUMMY_REPO_ID)
assert captured["data_files"] == expected_data_files.format(repo_id=DUMMY_REPO_ID)
assert captured["token"] == token
def test_bucket_metadata_url_root(tmp_path):
"""repo_type='bucket' produces url_root pointing at hf://buckets/..."""
with patch.object(LeRobotDatasetMetadata, "_load_metadata"):
meta = LeRobotDatasetMetadata(
DUMMY_REPO_ID,
root=tmp_path,
repo_type="bucket",
)
assert meta.url_root == f"hf://buckets/{DUMMY_REPO_ID}"
def test_bucket_skips_get_safe_version(tmp_path):
"""repo_type='bucket' must NOT call get_safe_version (buckets have no git refs).
The first ``_load_metadata()`` raises ``FileNotFoundError`` so ``__init__``
enters the except branch where the ``repo_type != "bucket"`` guard and
``get_safe_version`` live; the second call (after the bucket meta pull)
succeeds. ``_pull_from_repo`` is stubbed so the except path runs without a
network call. Without the raise, the try block would succeed and the guard
branch would never execute, so the assertion would pass vacuously.
"""
with (
patch("lerobot.datasets.dataset_metadata.get_safe_version") as mock_gsv,
patch.object(
LeRobotDatasetMetadata,
"_load_metadata",
side_effect=[FileNotFoundError, None],
),
patch.object(LeRobotDatasetMetadata, "_pull_from_repo") as mock_pull,
):
LeRobotDatasetMetadata(
DUMMY_REPO_ID,
root=tmp_path,
repo_type="bucket",
)
# The except branch ran (proven by the pull), but the bucket guard skipped
# version resolution.
mock_pull.assert_called_once()
mock_gsv.assert_not_called()
def test_bucket_metadata_sync_uses_stable_cache_and_token(tmp_path):
hub_cache = tmp_path / "hub"
token = "hf_test_token"
expected_root = hub_cache / f"buckets--{DUMMY_REPO_ID.replace('/', '--')}"
with (
patch("lerobot.datasets.dataset_metadata.HF_LEROBOT_HUB_CACHE", hub_cache),
patch("lerobot.datasets.dataset_metadata.sync_bucket") as mock_sync,
patch.object(
LeRobotDatasetMetadata,
"_load_metadata",
side_effect=[FileNotFoundError, None],
),
):
meta = LeRobotDatasetMetadata(DUMMY_REPO_ID, repo_type="bucket", token=token)
assert meta.root == expected_root
mock_sync.assert_called_once_with(
f"hf://buckets/{DUMMY_REPO_ID}/meta",
str(expected_root / "meta"),
delete=True,
quiet=True,
token=token,
)
with (
patch("lerobot.datasets.dataset_metadata.HF_LEROBOT_HUB_CACHE", hub_cache),
patch("lerobot.datasets.dataset_metadata.sync_bucket") as mock_sync,
patch.object(LeRobotDatasetMetadata, "_load_metadata"),
):
cached_meta = LeRobotDatasetMetadata(DUMMY_REPO_ID, repo_type="bucket", token=token)
assert cached_meta.root == expected_root
mock_sync.assert_not_called()
def test_repo_type_is_keyword_only_and_preserves_positional_episodes():
episodes = [1, 2]
with (
patch("lerobot.datasets.streaming_dataset.LeRobotDatasetMetadata", _fake_meta),
patch("lerobot.datasets.streaming_dataset.check_version_compatibility"),
patch(
"lerobot.datasets.streaming_dataset.load_dataset",
return_value=SimpleNamespace(num_shards=1),
),
):
dataset = StreamingLeRobotDataset(DUMMY_REPO_ID, None, episodes)
assert dataset.episodes == episodes
assert dataset.repo_type == "dataset"
def test_bucket_root_caches_metadata_without_switching_to_local_streaming(tmp_path):
with (
patch("lerobot.datasets.streaming_dataset.LeRobotDatasetMetadata", _fake_meta),
patch("lerobot.datasets.streaming_dataset.check_version_compatibility"),
patch(
"lerobot.datasets.streaming_dataset.load_dataset",
return_value=SimpleNamespace(num_shards=1),
) as mock_load_dataset,
):
dataset = StreamingLeRobotDataset(DUMMY_REPO_ID, root=tmp_path, repo_type="bucket")
assert dataset.root == tmp_path
assert not dataset.streaming_from_local
assert mock_load_dataset.call_args.args == ("parquet",)
assert mock_load_dataset.call_args.kwargs["data_files"].startswith("hf://buckets/")
def test_invalid_repo_type_fails_before_io():
with pytest.raises(ValueError, match="repo_type must be 'dataset' or 'bucket'"):
StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type="space")
Generated
+3 -3
View File
@@ -1,5 +1,5 @@
version = 1
revision = 2
revision = 3
requires-python = ">=3.12"
resolution-markers = [
"(python_full_version >= '3.15' and platform_machine == 'AMD64' and sys_platform == 'linux') or (python_full_version >= '3.15' and platform_machine == 'x86_64' and sys_platform == 'linux')",
@@ -3282,7 +3282,7 @@ requires-dist = [
{ name = "cmeel-tinyxml2", marker = "extra == 'placo-dep'", specifier = "<11" },
{ name = "cmeel-urdfdom", marker = "extra == 'placo-dep'", specifier = ">=4,<5" },
{ name = "contourpy", marker = "extra == 'matplotlib-dep'", specifier = ">=1.3.0,<2.0.0" },
{ name = "datasets", marker = "extra == 'dataset'", specifier = ">=4.7.0,<5.0.0" },
{ name = "datasets", marker = "extra == 'dataset'", specifier = ">=4.8.0,<5.0.0" },
{ name = "debugpy", marker = "extra == 'dev'", specifier = ">=1.8.1,<1.9.0" },
{ name = "decord", marker = "(platform_machine == 'AMD64' and extra == 'groot') or (platform_machine == 'x86_64' and extra == 'groot')", specifier = ">=0.6.0,<1.0.0" },
{ name = "deepdiff", marker = "extra == 'deepdiff-dep'", specifier = ">=7.0.1,<9.0.0" },
@@ -3305,7 +3305,7 @@ requires-dist = [
{ name = "hebi-py", marker = "extra == 'phone'", specifier = ">=2.8.0,<2.12.0" },
{ name = "hf-libero", marker = "sys_platform == 'linux' and extra == 'libero'", specifier = ">=0.1.4,<0.2.0" },
{ name = "hidapi", marker = "extra == 'gamepad'", specifier = ">=0.14.0,<0.15.0" },
{ name = "huggingface-hub", specifier = ">=1.0.0,<2.0.0" },
{ name = "huggingface-hub", specifier = ">=1.6.0,<2.0.0" },
{ name = "ipykernel", marker = "extra == 'notebook'", specifier = ">=6.0.0,<7.0.0" },
{ name = "jsonlines", marker = "extra == 'dataset'", specifier = ">=4.0.0,<5.0.0" },
{ name = "jupyter", marker = "extra == 'notebook'", specifier = ">=1.0.0,<2.0.0" },