diff --git a/pyproject.toml b/pyproject.toml index 2c538d024..393fdbe9a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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]", diff --git a/src/lerobot/datasets/dataset_metadata.py b/src/lerobot/datasets/dataset_metadata.py index 1f0e0add9..87ddeeebc 100644 --- a/src/lerobot/datasets/dataset_metadata.py +++ b/src/lerobot/datasets/dataset_metadata.py @@ -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 diff --git a/src/lerobot/datasets/streaming_dataset.py b/src/lerobot/datasets/streaming_dataset.py index b63e34dc8..d6cc24197 100644 --- a/src/lerobot/datasets/streaming_dataset.py +++ b/src/lerobot/datasets/streaming_dataset.py @@ -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) diff --git a/tests/datasets/test_streaming.py b/tests/datasets/test_streaming.py index 08544326f..099f373ca 100644 --- a/tests/datasets/test_streaming.py +++ b/tests/datasets/test_streaming.py @@ -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") diff --git a/uv.lock b/uv.lock index f60ba17a9..c7a79d89d 100644 --- a/uv.lock +++ b/uv.lock @@ -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" },