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 b339637a4..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 @@ -73,8 +75,8 @@ class LeRobotDatasetMetadata: revision: str | None = None, force_cache_sync: bool = False, metadata_buffer_size: int = 10, - repo_type: str = "dataset", *, + repo_type: Literal["dataset", "bucket"] = "dataset", token: str | bool | None = None, ): """Load or download metadata for an existing LeRobot dataset. @@ -104,32 +106,46 @@ class LeRobotDatasetMetadata: 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 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) + 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.""" @@ -236,21 +252,15 @@ class LeRobotDatasetMetadata: *, token: str | bool | None = None, ) -> None: - if getattr(self, "repo_type", "dataset") == "bucket": - from huggingface_hub import HfFileSystem - - fs = HfFileSystem() - dest = ( - self._requested_root - if self._requested_root is not None - else (HF_LEROBOT_HUB_CACHE / ("buckets--" + self.repo_id.replace("/", "--"))) + 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, ) - dest = Path(dest) - dest.mkdir(parents=True, exist_ok=True) - # fs.get copies the SOURCE dir INTO the target, so target=dest (not - # dest/meta) lands the tree as dest/meta/... rather than dest/meta/meta. - fs.get(f"hf://buckets/{self.repo_id}/meta", str(dest), recursive=True) - self.root = dest return token_kwargs = {} if token is None else {"token": token} if self._requested_root is None: @@ -282,7 +292,7 @@ class LeRobotDatasetMetadata: @property def url_root(self) -> str: """Hugging Face Hub URL root for this dataset.""" - if getattr(self, "repo_type", "dataset") == "bucket": + if self.repo_type == "bucket": return f"hf://buckets/{self.repo_id}" return f"hf://datasets/{self.repo_id}" diff --git a/src/lerobot/datasets/streaming_dataset.py b/src/lerobot/datasets/streaming_dataset.py index aadf9177f..b91042f8e 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 @@ -242,7 +243,6 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): self, repo_id: str, root: str | Path | None = None, - repo_type: str = "dataset", episodes: list[int] | None = None, image_transforms: Callable | None = None, delta_timestamps: dict[list[float]] | None = None, @@ -258,17 +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``. - repo_type (str, optional): "dataset" (default) or "bucket" to stream - from an HF Storage Bucket over hf://buckets/. + 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. @@ -283,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`` @@ -290,11 +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.repo_type = repo_type - self._requested_root = Path(root) if root else None + 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 @@ -350,15 +354,18 @@ 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 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: - token_kwargs = {} if token is None or self.streaming_from_local else {"token": token} + 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", diff --git a/tests/datasets/test_streaming.py b/tests/datasets/test_streaming.py index ee4df4981..a23461606 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 MagicMock, Mock, patch +from unittest.mock import Mock, patch import numpy as np import pytest @@ -441,8 +441,10 @@ class _StopConstructionError(Exception): def _fake_meta(*args, **kwargs): """Minimal LeRobotDatasetMetadata stand-in exposing only what __init__ reads.""" meta = type("_Meta", (), {})() - meta.root = kwargs.get("root") or "/tmp/_streaming_meta" - meta.revision = kwargs.get("revision") or "v0" + 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 = [] @@ -461,10 +463,12 @@ def _fake_meta(*args, **kwargs): 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 ( @@ -473,19 +477,16 @@ def test_streaming_repo_type_routes_load_dataset(repo_type, expected_source, exp patch("lerobot.datasets.streaming_dataset.load_dataset", fake_load_dataset), pytest.raises(_StopConstructionError), ): - StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type=repo_type) + 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/...""" - mock_fs = MagicMock() - with ( - patch("huggingface_hub.HfFileSystem", return_value=mock_fs), - patch.object(LeRobotDatasetMetadata, "_load_metadata"), - ): + with patch.object(LeRobotDatasetMetadata, "_load_metadata"): meta = LeRobotDatasetMetadata( DUMMY_REPO_ID, root=tmp_path, @@ -522,3 +523,77 @@ def test_bucket_skips_get_safe_version(tmp_path): # 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" },