mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 21:19:40 +00:00
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:
committed by
Steven Palma
parent
7e0fd0d653
commit
9782bfd64b
@@ -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
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user