mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-31 05:29:40 +00:00
feat(dataset): Support streaming from HF Storage Buckets + Bump HF hub & datasets (#4236)
This commit is contained in:
+2
-2
@@ -68,7 +68,7 @@ dependencies = [
|
|||||||
|
|
||||||
# Config & Hub
|
# Config & Hub
|
||||||
"draccus>=0.11.6,<0.12.0",
|
"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",
|
"requests>=2.32.0,<3.0.0",
|
||||||
|
|
||||||
# Environments
|
# Environments
|
||||||
@@ -95,7 +95,7 @@ dependencies = [
|
|||||||
|
|
||||||
# ── Feature-scoped extras ──────────────────────────────────
|
# ── Feature-scoped extras ──────────────────────────────────
|
||||||
dataset = [
|
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
|
"pandas>=2.0.0,<3.0.0", # NOTE: Transitive dependency of datasets
|
||||||
"pyarrow>=21.0.0,<30.0.0", # NOTE: Transitive dependency of datasets
|
"pyarrow>=21.0.0,<30.0.0", # NOTE: Transitive dependency of datasets
|
||||||
"lerobot[av-dep]",
|
"lerobot[av-dep]",
|
||||||
|
|||||||
@@ -18,13 +18,15 @@ import logging
|
|||||||
from collections.abc import Callable, Iterable
|
from collections.abc import Callable, Iterable
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import packaging.version
|
import packaging.version
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import pyarrow as pa
|
import pyarrow as pa
|
||||||
import pyarrow.parquet as pq
|
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.configs import DEPTH_METER_UNIT, VideoEncoderConfig
|
||||||
from lerobot.utils.constants import DEFAULT_FEATURES, HF_LEROBOT_HOME, HF_LEROBOT_HUB_CACHE
|
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,
|
force_cache_sync: bool = False,
|
||||||
metadata_buffer_size: int = 10,
|
metadata_buffer_size: int = 10,
|
||||||
*,
|
*,
|
||||||
|
repo_type: Literal["dataset", "bucket"] = "dataset",
|
||||||
token: str | bool | None = None,
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""Load or download metadata for an existing LeRobot dataset.
|
"""Load or download metadata for an existing LeRobot dataset.
|
||||||
@@ -96,36 +99,53 @@ 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.
|
||||||
"""
|
"""
|
||||||
|
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_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
|
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._pq_writer = None
|
||||||
self.latest_episode = None
|
self.latest_episode = None
|
||||||
self._metadata_buffer: list[dict] = []
|
self._metadata_buffer: list[dict] = []
|
||||||
self._metadata_buffer_size = metadata_buffer_size
|
self._metadata_buffer_size = metadata_buffer_size
|
||||||
self._finalized = False
|
self._finalized = False
|
||||||
|
|
||||||
try:
|
metadata_lock = contextlib.nullcontext()
|
||||||
if force_cache_sync or (
|
if self.repo_type == "bucket":
|
||||||
self._requested_root is None and has_legacy_hub_download_metadata(self.root)
|
self.root.parent.mkdir(parents=True, exist_ok=True)
|
||||||
):
|
metadata_lock = WeakFileLock(self.root.parent / f".{self.root.name}.lock")
|
||||||
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)
|
|
||||||
|
|
||||||
self._pull_from_repo(allow_patterns="meta/", token=token)
|
with metadata_lock:
|
||||||
self._load_metadata()
|
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:
|
def _flush_metadata_buffer(self) -> None:
|
||||||
"""Write all buffered episode metadata to parquet file."""
|
"""Write all buffered episode metadata to parquet file."""
|
||||||
@@ -232,6 +252,16 @@ class LeRobotDatasetMetadata:
|
|||||||
*,
|
*,
|
||||||
token: str | bool | None = None,
|
token: str | bool | None = 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}
|
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 +292,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 self.repo_type == "bucket":
|
||||||
|
return f"hf://buckets/{self.repo_id}"
|
||||||
return f"hf://datasets/{self.repo_id}"
|
return f"hf://datasets/{self.repo_id}"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -16,6 +16,7 @@
|
|||||||
from collections import deque
|
from collections import deque
|
||||||
from collections.abc import Callable, Generator, Iterable, Iterator
|
from collections.abc import Callable, Generator, Iterable, Iterator
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
import datasets
|
import datasets
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -261,15 +262,16 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
return_uint8: bool = False,
|
return_uint8: bool = False,
|
||||||
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
|
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
|
||||||
*,
|
*,
|
||||||
|
repo_type: Literal["dataset", "bucket"] = "dataset",
|
||||||
token: str | bool | None = None,
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""Initialize a StreamingLeRobotDataset.
|
"""Initialize a StreamingLeRobotDataset.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
repo_id (str): This is the repo id that will be used to fetch the dataset.
|
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
|
root (Path | None, optional): Local directory to use for local datasets. In bucket mode,
|
||||||
metadata is resolved through a revision-safe snapshot cache under
|
this is an optional local metadata-cache directory; parquet and video data remain remote.
|
||||||
``$HF_LEROBOT_HOME/hub``.
|
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
|
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.
|
||||||
@@ -284,6 +286,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
shuffle (bool, optional): Whether to shuffle the dataset across exhaustions. Defaults to True.
|
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").
|
depth_output_unit (str, optional): Physical unit depth maps are dequantized to ("m" or "mm").
|
||||||
Defaults to "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
|
token: Authentication token used while streaming this dataset from
|
||||||
the Hub. Pass a string token, ``True`` to require the locally
|
the Hub. Pass a string token, ``True`` to require the locally
|
||||||
stored token, ``False`` to disable authentication, or ``None``
|
stored token, ``False`` to disable authentication, or ``None``
|
||||||
@@ -291,10 +295,14 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
on the dataset instance after initialization.
|
on the dataset instance after initialization.
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
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_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.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.image_transforms = image_transforms
|
||||||
self.episodes = episodes
|
self.episodes = episodes
|
||||||
@@ -321,6 +329,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
|
||||||
@@ -349,15 +358,26 @@ 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)
|
||||||
|
|
||||||
token_kwargs = {} if token is None or self.streaming_from_local else {"token": token}
|
token_kwargs = {} if token is None else {"token": token}
|
||||||
self.hf_dataset: datasets.IterableDataset = load_dataset(
|
if self.repo_type == "bucket":
|
||||||
self.repo_id if not self.streaming_from_local else str(self.root),
|
self.hf_dataset: datasets.IterableDataset = load_dataset(
|
||||||
split="train",
|
"parquet",
|
||||||
streaming=self.streaming,
|
data_files=f"hf://buckets/{self.repo_id}/data/*/*.parquet",
|
||||||
data_files="data/*/*.parquet",
|
split="train",
|
||||||
revision=self.revision,
|
streaming=self.streaming,
|
||||||
**token_kwargs,
|
**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)
|
self.num_shards = min(self.hf_dataset.num_shards, max_num_shards)
|
||||||
|
|
||||||
|
|||||||
@@ -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 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:
|
||||||
@@ -478,3 +480,168 @@ 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", (), {})()
|
||||||
|
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")
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
version = 1
|
version = 1
|
||||||
revision = 2
|
revision = 3
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
resolution-markers = [
|
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')",
|
"(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-tinyxml2", marker = "extra == 'placo-dep'", specifier = "<11" },
|
||||||
{ name = "cmeel-urdfdom", marker = "extra == 'placo-dep'", specifier = ">=4,<5" },
|
{ name = "cmeel-urdfdom", marker = "extra == 'placo-dep'", specifier = ">=4,<5" },
|
||||||
{ name = "contourpy", marker = "extra == 'matplotlib-dep'", specifier = ">=1.3.0,<2.0.0" },
|
{ 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 = "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 = "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" },
|
{ 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 = "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 = "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 = "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 = "ipykernel", marker = "extra == 'notebook'", specifier = ">=6.0.0,<7.0.0" },
|
||||||
{ name = "jsonlines", marker = "extra == 'dataset'", specifier = ">=4.0.0,<5.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" },
|
{ name = "jupyter", marker = "extra == 'notebook'", specifier = ">=1.0.0,<2.0.0" },
|
||||||
|
|||||||
Reference in New Issue
Block a user