fix(dataset): bump versions + improvements

This commit is contained in:
Steven Palma
2026-07-30 15:55:21 +02:00
parent 13a261a08a
commit a3feb09b08
5 changed files with 147 additions and 55 deletions
+2 -2
View File
@@ -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]",
+28 -18
View File
@@ -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
@@ -73,8 +75,8 @@ 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",
*, *,
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.
@@ -104,17 +106,31 @@ class LeRobotDatasetMetadata:
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.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
metadata_lock = contextlib.nullcontext()
if self.repo_type == "bucket":
self.root.parent.mkdir(parents=True, exist_ok=True)
metadata_lock = WeakFileLock(self.root.parent / f".{self.root.name}.lock")
with metadata_lock:
try: try:
if force_cache_sync or ( if force_cache_sync or (
self._requested_root is None and has_legacy_hub_download_metadata(self.root) self._requested_root is None and has_legacy_hub_download_metadata(self.root)
@@ -236,21 +252,15 @@ class LeRobotDatasetMetadata:
*, *,
token: str | bool | None = None, token: str | bool | None = None,
) -> None: ) -> None:
if getattr(self, "repo_type", "dataset") == "bucket": if self.repo_type == "bucket":
from huggingface_hub import HfFileSystem self.root.mkdir(parents=True, exist_ok=True)
sync_bucket(
fs = HfFileSystem() f"hf://buckets/{self.repo_id}/meta",
dest = ( str(self.root / "meta"),
self._requested_root delete=True,
if self._requested_root is not None quiet=True,
else (HF_LEROBOT_HUB_CACHE / ("buckets--" + self.repo_id.replace("/", "--"))) token=token,
) )
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 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:
@@ -282,7 +292,7 @@ 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": if self.repo_type == "bucket":
return f"hf://buckets/{self.repo_id}" return f"hf://buckets/{self.repo_id}"
return f"hf://datasets/{self.repo_id}" return f"hf://datasets/{self.repo_id}"
+16 -9
View File
@@ -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
@@ -242,7 +243,6 @@ 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,
@@ -258,17 +258,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``.
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.
@@ -283,6 +282,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``
@@ -290,11 +291,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.repo_type = repo_type self.repo_type = repo_type
self._requested_root = Path(root) if root 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
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
@@ -350,15 +354,18 @@ 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 else {"token": token}
if self.repo_type == "bucket": if self.repo_type == "bucket":
self.hf_dataset: datasets.IterableDataset = load_dataset( self.hf_dataset: datasets.IterableDataset = load_dataset(
"parquet", "parquet",
data_files=f"hf://buckets/{self.repo_id}/data/*/*.parquet", data_files=f"hf://buckets/{self.repo_id}/data/*/*.parquet",
split="train", split="train",
streaming=self.streaming, streaming=self.streaming,
**token_kwargs,
) )
else: else:
token_kwargs = {} if token is None or self.streaming_from_local else {"token": token} if self.streaming_from_local:
token_kwargs = {}
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),
split="train", split="train",
+84 -9
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 MagicMock, Mock, patch from unittest.mock import Mock, patch
import numpy as np import numpy as np
import pytest import pytest
@@ -441,8 +441,10 @@ class _StopConstructionError(Exception):
def _fake_meta(*args, **kwargs): def _fake_meta(*args, **kwargs):
"""Minimal LeRobotDatasetMetadata stand-in exposing only what __init__ reads.""" """Minimal LeRobotDatasetMetadata stand-in exposing only what __init__ reads."""
meta = type("_Meta", (), {})() meta = type("_Meta", (), {})()
meta.root = kwargs.get("root") or "/tmp/_streaming_meta" root = kwargs.get("root", args[1] if len(args) > 1 else None)
meta.revision = kwargs.get("revision") or "v0" 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._version = "v3.0"
meta.depth_keys = [] meta.depth_keys = []
meta.image_keys = [] meta.image_keys = []
@@ -461,10 +463,12 @@ def _fake_meta(*args, **kwargs):
def test_streaming_repo_type_routes_load_dataset(repo_type, expected_source, expected_data_files): 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.""" """repo_type='bucket' loads parquet from hf://buckets/...; 'dataset' keeps the Hub-repo path."""
captured = {} captured = {}
token = "hf_test_token"
def fake_load_dataset(source, **kwargs): def fake_load_dataset(source, **kwargs):
captured["source"] = source captured["source"] = source
captured["data_files"] = kwargs.get("data_files") captured["data_files"] = kwargs.get("data_files")
captured["token"] = kwargs.get("token")
raise _StopConstructionError raise _StopConstructionError
with ( with (
@@ -473,19 +477,16 @@ def test_streaming_repo_type_routes_load_dataset(repo_type, expected_source, exp
patch("lerobot.datasets.streaming_dataset.load_dataset", fake_load_dataset), patch("lerobot.datasets.streaming_dataset.load_dataset", fake_load_dataset),
pytest.raises(_StopConstructionError), pytest.raises(_StopConstructionError),
): ):
StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type=repo_type) StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type=repo_type, token=token)
assert captured["source"] == expected_source.format(repo_id=DUMMY_REPO_ID) 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["data_files"] == expected_data_files.format(repo_id=DUMMY_REPO_ID)
assert captured["token"] == token
def test_bucket_metadata_url_root(tmp_path): def test_bucket_metadata_url_root(tmp_path):
"""repo_type='bucket' produces url_root pointing at hf://buckets/...""" """repo_type='bucket' produces url_root pointing at hf://buckets/..."""
mock_fs = MagicMock() with patch.object(LeRobotDatasetMetadata, "_load_metadata"):
with (
patch("huggingface_hub.HfFileSystem", return_value=mock_fs),
patch.object(LeRobotDatasetMetadata, "_load_metadata"),
):
meta = LeRobotDatasetMetadata( meta = LeRobotDatasetMetadata(
DUMMY_REPO_ID, DUMMY_REPO_ID,
root=tmp_path, root=tmp_path,
@@ -522,3 +523,77 @@ def test_bucket_skips_get_safe_version(tmp_path):
# version resolution. # version resolution.
mock_pull.assert_called_once() mock_pull.assert_called_once()
mock_gsv.assert_not_called() 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")
Generated
+3 -3
View File
@@ -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" },