Compare commits

...

3 Commits

Author SHA1 Message Date
Steven Palma a3feb09b08 fix(dataset): bump versions + improvements 2026-07-30 15:55:21 +02:00
Sundar Raghavan 13a261a08a test(streaming): make bucket get_safe_version test actually exercise the except branch
The prior test patched _load_metadata as a plain no-op, so __init__'s try
block succeeded and never entered the except branch where the
repo_type != "bucket" guard and get_safe_version live - the assertion passed
vacuously (verified: it still passed with the guard removed).

Use side_effect=[FileNotFoundError, None] so the first _load_metadata raises
(forcing the except path) and the second succeeds after the meta pull, and
stub _pull_from_repo so the path runs without a network call. Now the test
fails if the bucket guard is removed. Thanks @mohitydv09 for the catch.

Signed-off-by: Sundar Raghavan <sdraghav@amazon.com>
2026-07-30 15:45:43 +02:00
Sundar Raghavan 9782bfd64b Support streaming from HF Storage Buckets in StreamingLeRobotDataset
Add an opt-in repo_type="bucket" parameter to StreamingLeRobotDataset
and LeRobotDatasetMetadata so a dataset can be streamed directly from an
HF Storage Bucket (hf://buckets/...) with no local download.

When repo_type="bucket":
- skip git-version resolution (buckets have no refs/tags),
- pull the meta/ directory via HfFileSystem.get,
- point url_root at hf://buckets/{repo_id},
- read parquet shards via load_dataset("parquet",
  data_files="hf://buckets/{repo_id}/data/*/*.parquet", ...).

The default repo_type="dataset" preserves all existing behavior.
LeRobotDataset (non-streaming) and create() are unchanged.

Closes #3969

Signed-off-by: Sundar Raghavan <sdraghav@amazon.com>
2026-07-30 15:45:42 +02:00
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
@@ -257,15 +258,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.
@@ -280,6 +282,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``
@@ -287,10 +291,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
@@ -317,6 +325,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
@@ -345,15 +354,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:
@@ -430,3 +432,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" },