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>
This commit is contained in:
Sundar Raghavan
2026-07-17 23:34:38 -07:00
committed by Steven Palma
parent 7e0fd0d653
commit 9782bfd64b
3 changed files with 124 additions and 11 deletions
+23 -1
View File
@@ -73,6 +73,7 @@ class LeRobotDatasetMetadata:
revision: str | None = None, revision: str | None = None,
force_cache_sync: bool = False, force_cache_sync: bool = False,
metadata_buffer_size: int = 10, metadata_buffer_size: int = 10,
repo_type: str = "dataset",
*, *,
token: str | bool | None = None, token: str | bool | None = None,
): ):
@@ -96,12 +97,15 @@ class LeRobotDatasetMetadata:
even when local files exist. even when local files exist.
metadata_buffer_size: Number of episode metadata records to buffer metadata_buffer_size: Number of episode metadata records to buffer
in memory before flushing to parquet. 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: Authentication token used for Hub requests. Pass a string
token, ``True`` to require the locally stored token, ``False`` token, ``True`` to require the locally stored token, ``False``
to disable authentication, or ``None`` to use the Hugging Face to disable authentication, or ``None`` to use the Hugging Face
Hub default. Hub default.
""" """
self.repo_id = repo_id self.repo_id = repo_id
self.repo_type = repo_type
self.revision = revision if revision else CODEBASE_VERSION self.revision = revision if revision else CODEBASE_VERSION
self._requested_root = Path(root) if root is not None 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.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 raise FileNotFoundError
self._load_metadata() self._load_metadata()
except (FileNotFoundError, NotADirectoryError): except (FileNotFoundError, NotADirectoryError):
if is_valid_version(self.revision): if self.repo_type != "bucket" and is_valid_version(self.revision):
if token is None: if token is None:
self.revision = get_safe_version(self.repo_id, self.revision) self.revision = get_safe_version(self.repo_id, self.revision)
else: else:
@@ -232,6 +236,22 @@ class LeRobotDatasetMetadata:
*, *,
token: str | bool | None = None, token: str | bool | None = 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} token_kwargs = {} if token is None else {"token": token}
if self._requested_root is None: if self._requested_root is None:
self.root = Path( self.root = Path(
@@ -262,6 +282,8 @@ class LeRobotDatasetMetadata:
@property @property
def url_root(self) -> str: def url_root(self) -> str:
"""Hugging Face Hub URL root for this dataset.""" """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}" return f"hf://datasets/{self.repo_id}"
@property @property
+13
View File
@@ -242,6 +242,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self, self,
repo_id: str, repo_id: str,
root: str | Path | None = None, root: str | Path | None = None,
repo_type: str = "dataset",
episodes: list[int] | None = None, episodes: list[int] | None = None,
image_transforms: Callable | None = None, image_transforms: Callable | None = None,
delta_timestamps: dict[list[float]] | 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 root (Path | None, optional): Local directory to use for local datasets. When omitted, Hub
metadata is resolved through a revision-safe snapshot cache under metadata is resolved through a revision-safe snapshot cache under
``$HF_LEROBOT_HOME/hub``. ``$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 episodes (list[int] | None, optional): If specified, this will only load episodes specified by
their episode_index in this list. their episode_index in this list.
image_transforms (Callable | None, optional): Transform to apply to image data. image_transforms (Callable | None, optional): Transform to apply to image data.
@@ -288,6 +291,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
""" """
super().__init__() super().__init__()
self.repo_id = repo_id 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 else None
self.root = self._requested_root if self._requested_root is not None else HF_LEROBOT_HOME / repo_id 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
@@ -317,6 +321,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self._requested_root, self._requested_root,
self.revision, self.revision,
force_cache_sync=force_cache_sync, force_cache_sync=force_cache_sync,
repo_type=self.repo_type,
token=token, token=token,
) )
self.root = self.meta.root self.root = self.meta.root
@@ -345,6 +350,14 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self.delta_timestamps = delta_timestamps self.delta_timestamps = delta_timestamps
self.delta_indices = get_delta_indices(self.delta_timestamps, self.fps) self.delta_indices = get_delta_indices(self.delta_timestamps, self.fps)
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} token_kwargs = {} if token is None or self.streaming_from_local else {"token": token}
self.hf_dataset: datasets.IterableDataset = load_dataset( self.hf_dataset: datasets.IterableDataset = load_dataset(
self.repo_id if not self.streaming_from_local else str(self.root), self.repo_id if not self.streaming_from_local else str(self.root),
+79 -1
View File
@@ -14,7 +14,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import Mock from unittest.mock import MagicMock, Mock, patch
import numpy as np import numpy as np
import pytest import pytest
@@ -23,6 +23,7 @@ import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])") pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
import lerobot.datasets.streaming_dataset as streaming_dataset_module 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.streaming_dataset import StreamingLeRobotDataset
from lerobot.datasets.utils import safe_shard from lerobot.datasets.utils import safe_shard
from lerobot.utils.constants import ACTION 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, requested_root,
streaming_dataset_module.CODEBASE_VERSION, streaming_dataset_module.CODEBASE_VERSION,
force_cache_sync=False, force_cache_sync=False,
repo_type="dataset",
token=token, token=token,
) )
if from_local: if from_local:
@@ -430,3 +432,79 @@ def test_frames_with_delta_consistency_with_shards(
assert all(t[1] for t in key_checks), ( 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})" 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()