diff --git a/src/lerobot/datasets/dataset_metadata.py b/src/lerobot/datasets/dataset_metadata.py index 1f0e0add9..b339637a4 100644 --- a/src/lerobot/datasets/dataset_metadata.py +++ b/src/lerobot/datasets/dataset_metadata.py @@ -73,6 +73,7 @@ class LeRobotDatasetMetadata: revision: str | None = None, force_cache_sync: bool = False, metadata_buffer_size: int = 10, + repo_type: str = "dataset", *, token: str | bool | None = None, ): @@ -96,12 +97,15 @@ 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. """ 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 @@ -118,7 +122,7 @@ class LeRobotDatasetMetadata: raise FileNotFoundError self._load_metadata() except (FileNotFoundError, NotADirectoryError): - if is_valid_version(self.revision): + 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: @@ -232,6 +236,22 @@ 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("/", "--"))) + ) + 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: self.root = Path( @@ -262,6 +282,8 @@ class LeRobotDatasetMetadata: @property def url_root(self) -> str: """Hugging Face Hub URL root for this dataset.""" + if getattr(self, "repo_type", "dataset") == "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 806f2c24c..aadf9177f 100644 --- a/src/lerobot/datasets/streaming_dataset.py +++ b/src/lerobot/datasets/streaming_dataset.py @@ -242,6 +242,7 @@ 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, @@ -266,6 +267,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): 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/. 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. @@ -288,6 +291,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset): """ super().__init__() self.repo_id = repo_id + self.repo_type = repo_type self._requested_root = Path(root) if root 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 @@ -317,6 +321,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 +350,23 @@ 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, - ) + 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, + ) + else: + 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, + ) 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 cae4be5b6..ac9f06841 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 MagicMock, 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,79 @@ 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", (), {})() + meta.root = kwargs.get("root") or "/tmp/_streaming_meta" + meta.revision = kwargs.get("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 = {} + + def fake_load_dataset(source, **kwargs): + captured["source"] = source + captured["data_files"] = kwargs.get("data_files") + 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) + + assert captured["source"] == expected_source.format(repo_id=DUMMY_REPO_ID) + assert captured["data_files"] == expected_data_files.format(repo_id=DUMMY_REPO_ID) + + +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"), + ): + 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).""" + mock_fs = MagicMock() + with ( + patch("lerobot.datasets.dataset_metadata.get_safe_version") as mock_gsv, + patch("huggingface_hub.HfFileSystem", return_value=mock_fs), + patch.object(LeRobotDatasetMetadata, "_load_metadata"), + ): + LeRobotDatasetMetadata( + DUMMY_REPO_ID, + root=tmp_path, + repo_type="bucket", + ) + mock_gsv.assert_not_called()