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:
@@ -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 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:
|
||||
@@ -478,3 +480,168 @@ 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", (), {})()
|
||||
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")
|
||||
|
||||
Reference in New Issue
Block a user