feat(dataset): accept token argument for private HF Hub datasets (#4136)

This commit is contained in:
Steven Palma
2026-07-24 18:51:35 +02:00
committed by GitHub
parent ab2b5b04dd
commit 0d383d09f2
9 changed files with 191 additions and 13 deletions
+16 -2
View File
@@ -73,6 +73,8 @@ class LeRobotDatasetMetadata:
revision: str | None = None,
force_cache_sync: bool = False,
metadata_buffer_size: int = 10,
*,
token: str | bool | None = None,
):
"""Load or download metadata for an existing LeRobot dataset.
@@ -94,6 +96,10 @@ class LeRobotDatasetMetadata:
even when local files exist.
metadata_buffer_size: Number of episode metadata records to buffer
in memory before flushing to parquet.
token: Authentication token used for Hub requests. Pass a string
token, ``True`` to require the locally stored token, ``False``
to disable authentication, or ``None`` to use the Hugging Face
Hub default.
"""
self.repo_id = repo_id
self.revision = revision if revision else CODEBASE_VERSION
@@ -113,9 +119,12 @@ class LeRobotDatasetMetadata:
self._load_metadata()
except (FileNotFoundError, NotADirectoryError):
if is_valid_version(self.revision):
self.revision = get_safe_version(self.repo_id, 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/")
self._pull_from_repo(allow_patterns="meta/", token=token)
self._load_metadata()
def _flush_metadata_buffer(self) -> None:
@@ -220,7 +229,10 @@ class LeRobotDatasetMetadata:
self,
allow_patterns: list[str] | str | None = None,
ignore_patterns: list[str] | str | None = None,
*,
token: str | bool | None = None,
) -> None:
token_kwargs = {} if token is None else {"token": token}
if self._requested_root is None:
self.root = Path(
snapshot_download(
@@ -230,6 +242,7 @@ class LeRobotDatasetMetadata:
cache_dir=HF_LEROBOT_HUB_CACHE,
allow_patterns=allow_patterns,
ignore_patterns=ignore_patterns,
**token_kwargs,
)
)
return
@@ -242,6 +255,7 @@ class LeRobotDatasetMetadata:
local_dir=self._requested_root,
allow_patterns=allow_patterns,
ignore_patterns=ignore_patterns,
**token_kwargs,
)
self.root = self._requested_root
+30 -5
View File
@@ -65,6 +65,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
encoder_threads: int | None = None,
streaming_encoding: bool = False,
encoder_queue_maxsize: int = 30,
*,
token: str | bool | None = None,
):
"""
2 modes are available for instantiating this class, depending on 2 different use cases:
@@ -197,6 +199,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
instead of writing PNG images first. This makes save_episode() near-instant. Defaults to False.
encoder_queue_maxsize (int, optional): Maximum number of frames to buffer per camera when using
streaming encoding. Defaults to 30 (~1s at 30fps).
token: Authentication token used while downloading this dataset
from the Hub. Pass a string token, ``True`` to require the
locally stored token, ``False`` to disable authentication, or
``None`` to use the Hugging Face Hub default. The token is not
retained on the dataset instance after initialization.
Note:
Write-mode parameters (``streaming_encoding``, ``batch_encoding_size``) passed to
@@ -220,7 +227,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
# Load metadata (sets self.root once from the resolved metadata root)
self.meta = LeRobotDatasetMetadata(
self.repo_id, self._requested_root, self.revision, force_cache_sync=force_cache_sync
self.repo_id,
self._requested_root,
self.revision,
force_cache_sync=force_cache_sync,
token=token,
)
self.root = self.meta.root
self.revision = self.meta.revision
@@ -260,8 +271,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
# Load actual data
if force_cache_sync or not self.reader.try_load():
if is_valid_version(self.revision):
self.revision = get_safe_version(self.repo_id, self.revision)
self._download(download_videos)
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._download(download_videos, token=token)
self.reader.load_and_activate()
# Detect write-mode params for backward compatibility
@@ -626,10 +640,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
hub_api.delete_tag(self.repo_id, tag=CODEBASE_VERSION, repo_type="dataset")
hub_api.create_tag(self.repo_id, tag=CODEBASE_VERSION, revision=branch, repo_type="dataset")
def _download(self, download_videos: bool = True) -> None:
def _download(self, download_videos: bool = True, *, token: str | bool | None = None) -> None:
"""Downloads the dataset from the given 'repo_id' at the provided version."""
ignore_patterns = None if download_videos else "videos/"
files = None
token_kwargs = {} if token is None else {"token": token}
if self.episodes is not None:
# Reader is guaranteed to exist here (created in __init__ before _download)
files = self.reader.get_episodes_file_paths()
@@ -643,6 +658,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
cache_dir=HF_LEROBOT_HUB_CACHE,
allow_patterns=files,
ignore_patterns=ignore_patterns,
**token_kwargs,
)
)
else:
@@ -654,6 +670,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
local_dir=self._requested_root,
allow_patterns=files,
ignore_patterns=ignore_patterns,
**token_kwargs,
)
self.meta.root = self._requested_root
@@ -793,6 +810,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
image_writer_threads: int = 0,
streaming_encoding: bool = False,
encoder_queue_maxsize: int = 30,
*,
token: str | bool | None = None,
) -> "LeRobotDataset":
"""Resume recording on an existing dataset.
@@ -826,6 +845,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
streaming_encoding: If ``True``, encode video in real-time during
capture.
encoder_queue_maxsize: Max buffered frames per camera for streaming.
token: Authentication token used if metadata must be downloaded
from the Hub. The token is not retained on the dataset instance.
Returns:
A :class:`LeRobotDataset` in write mode, ready to append episodes.
@@ -854,7 +875,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
# Load metadata (revision-safe when root is not provided)
obj.meta = LeRobotDatasetMetadata(
obj.repo_id, obj._requested_root, obj.revision, force_cache_sync=force_cache_sync
obj.repo_id,
obj._requested_root,
obj.revision,
force_cache_sync=force_cache_sync,
token=token,
)
obj._encoder_threads = encoder_threads
+3
View File
@@ -48,6 +48,8 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
tolerances_s: dict | None = None,
download_videos: bool = True,
video_backend: str | None = None,
*,
token: str | bool | None = None,
):
super().__init__()
self.repo_ids = repo_ids
@@ -65,6 +67,7 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
tolerance_s=self.tolerances_s[repo_id],
download_videos=download_videos,
video_backend=video_backend,
token=token,
)
for repo_id in repo_ids
]
+14 -1
View File
@@ -256,6 +256,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
shuffle: bool = True,
return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
*,
token: str | bool | None = None,
):
"""Initialize a StreamingLeRobotDataset.
@@ -278,6 +280,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
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").
Defaults to "mm".
token: Authentication token used while streaming this dataset from
the Hub. Pass a string token, ``True`` to require the locally
stored token, ``False`` to disable authentication, or ``None``
to use the Hugging Face Hub default. The token is not retained
on the dataset instance after initialization.
"""
super().__init__()
self.repo_id = repo_id
@@ -306,7 +313,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
# Load metadata
self.meta = LeRobotDatasetMetadata(
self.repo_id, self._requested_root, self.revision, force_cache_sync=force_cache_sync
self.repo_id,
self._requested_root,
self.revision,
force_cache_sync=force_cache_sync,
token=token,
)
self.root = self.meta.root
self.revision = self.meta.revision
@@ -334,12 +345,14 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self.delta_timestamps = delta_timestamps
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}
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)
+13 -4
View File
@@ -325,16 +325,19 @@ def check_version_compatibility(
logging.warning(FUTURE_MESSAGE.format(repo_id=repo_id, version=v_check))
def get_repo_versions(repo_id: str) -> list[packaging.version.Version]:
def get_repo_versions(repo_id: str, *, token: str | bool | None = None) -> list[packaging.version.Version]:
"""Return available valid versions (branches and tags) on a given Hub repo.
Args:
repo_id (str): The repository ID on the Hugging Face Hub.
token: Authentication token used for Hub requests. Pass a string token,
``True`` to require the locally stored token, ``False`` to disable
authentication, or ``None`` to use the Hugging Face Hub default.
Returns:
list[packaging.version.Version]: A list of valid versions found.
"""
api = HfApi()
api = HfApi() if token is None else HfApi(token=token)
repo_refs = api.list_repo_refs(repo_id, repo_type="dataset")
repo_refs = [b.name for b in repo_refs.branches + repo_refs.tags]
repo_versions = []
@@ -345,7 +348,12 @@ def get_repo_versions(repo_id: str) -> list[packaging.version.Version]:
return repo_versions
def get_safe_version(repo_id: str, version: str | packaging.version.Version) -> str:
def get_safe_version(
repo_id: str,
version: str | packaging.version.Version,
*,
token: str | bool | None = None,
) -> str:
"""Return the specified version if available on repo, or the latest compatible one.
If the exact version is not found, it looks for the latest version with the
@@ -354,6 +362,7 @@ def get_safe_version(repo_id: str, version: str | packaging.version.Version) ->
Args:
repo_id (str): The repository ID on the Hugging Face Hub.
version (str | packaging.version.Version): The target version.
token: Authentication token forwarded to the Hub version lookup.
Returns:
str: The safe version string (e.g., "v1.2.3") to use as a revision.
@@ -366,7 +375,7 @@ def get_safe_version(repo_id: str, version: str | packaging.version.Version) ->
target_version = (
packaging.version.parse(version) if not isinstance(version, packaging.version.Version) else version
)
hub_versions = get_repo_versions(repo_id)
hub_versions = get_repo_versions(repo_id) if token is None else get_repo_versions(repo_id, token=token)
if not hub_versions:
raise RevisionNotFoundError(