mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-28 12:15:59 +00:00
Merge origin/main into streaming byte-cache branch
This commit is contained in:
@@ -51,6 +51,7 @@ pre-commit run --all-files # Lint + format (ruff, typo
|
|||||||
## Notes
|
## Notes
|
||||||
|
|
||||||
- **Mypy is gradual**: strict only for `lerobot.envs`, `lerobot.configs`, `lerobot.optim`, `lerobot.model`, `lerobot.cameras`, `lerobot.motors`, `lerobot.transport`. Add type annotations when modifying these modules.
|
- **Mypy is gradual**: strict only for `lerobot.envs`, `lerobot.configs`, `lerobot.optim`, `lerobot.model`, `lerobot.cameras`, `lerobot.motors`, `lerobot.transport`. Add type annotations when modifying these modules.
|
||||||
- **Optional dependencies**: many policies, envs, and robots are behind extras (e.g., `lerobot[aloha]`). New imports for optional packages must be guarded or lazy. See `pyproject.toml [project.optional-dependencies]`.
|
- **Imports**: prefer top-level imports; relative (`from .sibling import X`) across sibling files within a module, absolute (`from lerobot.module import X`) across modules.
|
||||||
|
- **Optional dependencies**: many policies, envs, and robots are behind extras (e.g., `lerobot[aloha]`, see `pyproject.toml`). Guard optional imports with `TYPE_CHECKING or _foo_available` at module top + a `require_package(...)` check at use time. Reuse the `_foo_available` flags in `utils/import_utils.py`; don't call `is_package_available`.
|
||||||
- **Video decoding**: datasets can store observations as video files. `LeRobotDataset` handles frame extraction, but tests need ffmpeg installed.
|
- **Video decoding**: datasets can store observations as video files. `LeRobotDataset` handles frame extraction, but tests need ffmpeg installed.
|
||||||
- **Prioritize use of `uv run`** to execute Python commands (not raw `python` or `pip`).
|
- **Prioritize use of `uv run`** to execute Python commands (not raw `python` or `pip`).
|
||||||
|
|||||||
@@ -165,6 +165,8 @@ Batches are flat dictionaries keyed by the constants in [`lerobot.utils.constant
|
|||||||
|
|
||||||
LeRobot uses `PolicyProcessorPipeline`s to normalize inputs and de-normalize outputs around your policy. For a concrete reference, see [`processor_act.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/act/processor_act.py) or [`processor_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/processor_diffusion.py).
|
LeRobot uses `PolicyProcessorPipeline`s to normalize inputs and de-normalize outputs around your policy. For a concrete reference, see [`processor_act.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/act/processor_act.py) or [`processor_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/processor_diffusion.py).
|
||||||
|
|
||||||
|
Pay close attention here: processors are the most common reproducibility pain point. A mismatch in normalization mode (`IDENTITY` vs `MEAN_STD` vs `MIN_MAX` vs `QUANTILES`/`QUANTILE10`) or in which features get normalized will train and eval without erroring, yet silently wreck results. Make sure the modes match how the checkpoint was trained, that the required stats exist (e.g. `QUANTILES` needs `q01`/`q99`), and that the pre- and post-processors stay consistent.
|
||||||
|
|
||||||
```python
|
```python
|
||||||
# processor_my_policy.py
|
# processor_my_policy.py
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -304,7 +306,9 @@ Mirror an existing policy that's structurally similar to yours; the diff is smal
|
|||||||
|
|
||||||
### Heavy / optional dependencies
|
### Heavy / optional dependencies
|
||||||
|
|
||||||
Most policies need a heavy backbone (transformers, diffusers, a specific VLM SDK). The convention is **two-step gating**: a `TYPE_CHECKING`-guarded import at module top, and a `require_package` runtime check in the constructor. [`modeling_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/modeling_diffusion.py) is the canonical reference:
|
Most policies need a heavy backbone (transformers, diffusers, a specific VLM SDK). Wherever one exists, prefer loading it e.g from `transformers` or `diffusers` rather than re-implementing the architecture in-tree.
|
||||||
|
|
||||||
|
The convention is **two-step gating**: a `TYPE_CHECKING`-guarded import at module top, and a `require_package` runtime check in the constructor. [`modeling_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/modeling_diffusion.py) is the canonical reference:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
@@ -374,6 +378,7 @@ The general expectations are in [`CONTRIBUTING.md`](https://github.com/huggingfa
|
|||||||
- [ ] Optional deps live behind a `[project.optional-dependencies]` extra and the `TYPE_CHECKING + require_package` guard.
|
- [ ] Optional deps live behind a `[project.optional-dependencies]` extra and the `TYPE_CHECKING + require_package` guard.
|
||||||
- [ ] `tests/policies/` updated; backward-compat artifact committed & policy-specific tests.
|
- [ ] `tests/policies/` updated; backward-compat artifact committed & policy-specific tests.
|
||||||
- [ ] `src/lerobot/policies/<name>/README.md` symlinked into `docs/source/policy_<name>_README.md`; user-facing `docs/source/<name>.mdx` written and added to `_toctree.yml`.
|
- [ ] `src/lerobot/policies/<name>/README.md` symlinked into `docs/source/policy_<name>_README.md`; user-facing `docs/source/<name>.mdx` written and added to `_toctree.yml`.
|
||||||
|
- [ ] `lerobot-train --policy.type my_policy ...` runs end-to-end for at least a few steps + save a checkpoint that can be loaded and run by `lerobot-eval` or `lerobot-rollout`.
|
||||||
- [ ] `templates/lerobot_modelcard_template.md` has a description entry and a `policy_docs` link for your policy.
|
- [ ] `templates/lerobot_modelcard_template.md` has a description entry and a `policy_docs` link for your policy.
|
||||||
- [ ] The models table in the root `README.md` lists your policy in the right category, linking to your doc page.
|
- [ ] The models table in the root `README.md` lists your policy in the right category, linking to your doc page.
|
||||||
- [ ] At least one reproducible benchmark eval in the policy MDX with a published checkpoint (sim benchmark, or real-robot dataset + checkpoint).
|
- [ ] At least one reproducible benchmark eval in the policy MDX with a published checkpoint (sim benchmark, or real-robot dataset + checkpoint).
|
||||||
|
|||||||
@@ -252,6 +252,10 @@ lerobot-dataset-viz \
|
|||||||
--episode-index 0
|
--episode-index 0
|
||||||
```
|
```
|
||||||
|
|
||||||
|
For a private or gated dataset, authenticate first with `hf auth login`, or set the
|
||||||
|
`HF_TOKEN` environment variable. The Hub client then discovers the credential
|
||||||
|
automatically; no token argument is needed.
|
||||||
|
|
||||||
**From a local folder:**
|
**From a local folder:**
|
||||||
Add the `--root` option and set `--mode local`. For example, to search in `./my_local_data_dir/lerobot/pusht`:
|
Add the `--root` option and set `--mode local`. For example, to search in `./my_local_data_dir/lerobot/pusht`:
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -158,7 +158,7 @@ accelerate-dep = ["accelerate>=1.14.0,<2.0.0"]
|
|||||||
can-dep = ["python-can>=4.2.0,<5.0.0"]
|
can-dep = ["python-can>=4.2.0,<5.0.0"]
|
||||||
peft-dep = ["peft>=0.18.0,<1.0.0"]
|
peft-dep = ["peft>=0.18.0,<1.0.0"]
|
||||||
scipy-dep = ["scipy>=1.14.0,<2.0.0"]
|
scipy-dep = ["scipy>=1.14.0,<2.0.0"]
|
||||||
diffusers-dep = ["diffusers>=0.27.2,<0.36.0"]
|
diffusers-dep = ["diffusers>=0.38.0,<0.40.0"]
|
||||||
qwen-vl-utils-dep = ["qwen-vl-utils>=0.0.11,<0.1.0"]
|
qwen-vl-utils-dep = ["qwen-vl-utils>=0.0.11,<0.1.0"]
|
||||||
matplotlib-dep = ["matplotlib>=3.10.3,<4.0.0", "contourpy>=1.3.0,<2.0.0"] # NOTE: Explicitly listing contourpy helps the resolver converge faster.
|
matplotlib-dep = ["matplotlib>=3.10.3,<4.0.0", "contourpy>=1.3.0,<2.0.0"] # NOTE: Explicitly listing contourpy helps the resolver converge faster.
|
||||||
pyserial-dep = ["pyserial>=3.5,<4.0"]
|
pyserial-dep = ["pyserial>=3.5,<4.0"]
|
||||||
|
|||||||
@@ -173,7 +173,8 @@ class Reachy2Camera(Camera):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Invalid color mode '{self.color_mode}'. Expected {ColorMode.RGB} or {ColorMode.BGR}."
|
f"Invalid color mode '{self.color_mode}'. Expected {ColorMode.RGB} or {ColorMode.BGR}."
|
||||||
)
|
)
|
||||||
if self.color_mode == ColorMode.RGB:
|
is_depth_frame = self.config.name == "depth" and self.config.image_type == "depth"
|
||||||
|
if not is_depth_frame and self.color_mode == ColorMode.RGB:
|
||||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||||
|
|
||||||
self.latest_frame = frame
|
self.latest_frame = frame
|
||||||
|
|||||||
@@ -453,7 +453,7 @@ class RealSenseCamera(Camera):
|
|||||||
)
|
)
|
||||||
|
|
||||||
processed_image = image
|
processed_image = image
|
||||||
if self.color_mode == ColorMode.BGR:
|
if not depth_frame and self.color_mode == ColorMode.BGR:
|
||||||
processed_image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
|
processed_image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
|
||||||
|
|
||||||
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE, cv2.ROTATE_180]:
|
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE, cv2.ROTATE_180]:
|
||||||
|
|||||||
@@ -73,6 +73,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,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""Load or download metadata for an existing LeRobot dataset.
|
"""Load or download metadata for an existing LeRobot dataset.
|
||||||
|
|
||||||
@@ -94,6 +96,10 @@ 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.
|
||||||
|
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.repo_id = repo_id
|
||||||
self.revision = revision if revision else CODEBASE_VERSION
|
self.revision = revision if revision else CODEBASE_VERSION
|
||||||
@@ -113,9 +119,12 @@ class LeRobotDatasetMetadata:
|
|||||||
self._load_metadata()
|
self._load_metadata()
|
||||||
except (FileNotFoundError, NotADirectoryError):
|
except (FileNotFoundError, NotADirectoryError):
|
||||||
if is_valid_version(self.revision):
|
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()
|
self._load_metadata()
|
||||||
|
|
||||||
def _flush_metadata_buffer(self) -> None:
|
def _flush_metadata_buffer(self) -> None:
|
||||||
@@ -220,7 +229,10 @@ class LeRobotDatasetMetadata:
|
|||||||
self,
|
self,
|
||||||
allow_patterns: list[str] | str | None = None,
|
allow_patterns: list[str] | str | None = None,
|
||||||
ignore_patterns: list[str] | str | None = None,
|
ignore_patterns: list[str] | str | None = None,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
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(
|
||||||
snapshot_download(
|
snapshot_download(
|
||||||
@@ -230,6 +242,7 @@ class LeRobotDatasetMetadata:
|
|||||||
cache_dir=HF_LEROBOT_HUB_CACHE,
|
cache_dir=HF_LEROBOT_HUB_CACHE,
|
||||||
allow_patterns=allow_patterns,
|
allow_patterns=allow_patterns,
|
||||||
ignore_patterns=ignore_patterns,
|
ignore_patterns=ignore_patterns,
|
||||||
|
**token_kwargs,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -242,6 +255,7 @@ class LeRobotDatasetMetadata:
|
|||||||
local_dir=self._requested_root,
|
local_dir=self._requested_root,
|
||||||
allow_patterns=allow_patterns,
|
allow_patterns=allow_patterns,
|
||||||
ignore_patterns=ignore_patterns,
|
ignore_patterns=ignore_patterns,
|
||||||
|
**token_kwargs,
|
||||||
)
|
)
|
||||||
self.root = self._requested_root
|
self.root = self._requested_root
|
||||||
|
|
||||||
|
|||||||
@@ -65,6 +65,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
encoder_threads: int | None = None,
|
encoder_threads: int | None = None,
|
||||||
streaming_encoding: bool = False,
|
streaming_encoding: bool = False,
|
||||||
encoder_queue_maxsize: int = 30,
|
encoder_queue_maxsize: int = 30,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
2 modes are available for instantiating this class, depending on 2 different use cases:
|
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.
|
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
|
encoder_queue_maxsize (int, optional): Maximum number of frames to buffer per camera when using
|
||||||
streaming encoding. Defaults to 30 (~1s at 30fps).
|
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:
|
Note:
|
||||||
Write-mode parameters (``streaming_encoding``, ``batch_encoding_size``) passed to
|
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)
|
# Load metadata (sets self.root once from the resolved metadata root)
|
||||||
self.meta = LeRobotDatasetMetadata(
|
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.root = self.meta.root
|
||||||
self.revision = self.meta.revision
|
self.revision = self.meta.revision
|
||||||
@@ -260,8 +271,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
# Load actual data
|
# Load actual data
|
||||||
if force_cache_sync or not self.reader.try_load():
|
if force_cache_sync or not self.reader.try_load():
|
||||||
if is_valid_version(self.revision):
|
if is_valid_version(self.revision):
|
||||||
self.revision = get_safe_version(self.repo_id, self.revision)
|
if token is None:
|
||||||
self._download(download_videos)
|
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()
|
self.reader.load_and_activate()
|
||||||
|
|
||||||
# Detect write-mode params for backward compatibility
|
# 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.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")
|
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."""
|
"""Downloads the dataset from the given 'repo_id' at the provided version."""
|
||||||
ignore_patterns = None if download_videos else "videos/"
|
ignore_patterns = None if download_videos else "videos/"
|
||||||
files = None
|
files = None
|
||||||
|
token_kwargs = {} if token is None else {"token": token}
|
||||||
if self.episodes is not None:
|
if self.episodes is not None:
|
||||||
# Reader is guaranteed to exist here (created in __init__ before _download)
|
# Reader is guaranteed to exist here (created in __init__ before _download)
|
||||||
files = self.reader.get_episodes_file_paths()
|
files = self.reader.get_episodes_file_paths()
|
||||||
@@ -643,6 +658,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
cache_dir=HF_LEROBOT_HUB_CACHE,
|
cache_dir=HF_LEROBOT_HUB_CACHE,
|
||||||
allow_patterns=files,
|
allow_patterns=files,
|
||||||
ignore_patterns=ignore_patterns,
|
ignore_patterns=ignore_patterns,
|
||||||
|
**token_kwargs,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -654,6 +670,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
local_dir=self._requested_root,
|
local_dir=self._requested_root,
|
||||||
allow_patterns=files,
|
allow_patterns=files,
|
||||||
ignore_patterns=ignore_patterns,
|
ignore_patterns=ignore_patterns,
|
||||||
|
**token_kwargs,
|
||||||
)
|
)
|
||||||
self.meta.root = self._requested_root
|
self.meta.root = self._requested_root
|
||||||
|
|
||||||
@@ -793,6 +810,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
image_writer_threads: int = 0,
|
image_writer_threads: int = 0,
|
||||||
streaming_encoding: bool = False,
|
streaming_encoding: bool = False,
|
||||||
encoder_queue_maxsize: int = 30,
|
encoder_queue_maxsize: int = 30,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
) -> "LeRobotDataset":
|
) -> "LeRobotDataset":
|
||||||
"""Resume recording on an existing dataset.
|
"""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
|
streaming_encoding: If ``True``, encode video in real-time during
|
||||||
capture.
|
capture.
|
||||||
encoder_queue_maxsize: Max buffered frames per camera for streaming.
|
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:
|
Returns:
|
||||||
A :class:`LeRobotDataset` in write mode, ready to append episodes.
|
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)
|
# Load metadata (revision-safe when root is not provided)
|
||||||
obj.meta = LeRobotDatasetMetadata(
|
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
|
obj._encoder_threads = encoder_threads
|
||||||
|
|||||||
@@ -48,6 +48,8 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
|||||||
tolerances_s: dict | None = None,
|
tolerances_s: dict | None = None,
|
||||||
download_videos: bool = True,
|
download_videos: bool = True,
|
||||||
video_backend: str | None = None,
|
video_backend: str | None = None,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.repo_ids = repo_ids
|
self.repo_ids = repo_ids
|
||||||
@@ -65,6 +67,7 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
|||||||
tolerance_s=self.tolerances_s[repo_id],
|
tolerance_s=self.tolerances_s[repo_id],
|
||||||
download_videos=download_videos,
|
download_videos=download_videos,
|
||||||
video_backend=video_backend,
|
video_backend=video_backend,
|
||||||
|
token=token,
|
||||||
)
|
)
|
||||||
for repo_id in repo_ids
|
for repo_id in repo_ids
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -120,6 +120,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
prefetch_episodes: int = 8,
|
prefetch_episodes: int = 8,
|
||||||
byte_budget_gb: float = 8.0,
|
byte_budget_gb: float = 8.0,
|
||||||
repeat: bool = False,
|
repeat: bool = False,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""Initialize a StreamingLeRobotDataset.
|
"""Initialize a StreamingLeRobotDataset.
|
||||||
|
|
||||||
@@ -153,6 +155,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
repeat (bool, optional): Repeat rank-local exact-coverage epochs without yielding a
|
repeat (bool, optional): Repeat rank-local exact-coverage epochs without yielding a
|
||||||
short final training batch. The training factory enables this; direct iteration is
|
short final training batch. The training factory enables this; direct iteration is
|
||||||
finite by default.
|
finite by default.
|
||||||
|
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__()
|
super().__init__()
|
||||||
self.repo_id = repo_id
|
self.repo_id = repo_id
|
||||||
@@ -203,7 +210,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
|
|
||||||
# Load metadata
|
# Load metadata
|
||||||
self.meta = LeRobotDatasetMetadata(
|
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.root = self.meta.root
|
||||||
self.revision = self.meta.revision
|
self.revision = self.meta.revision
|
||||||
|
|||||||
@@ -325,16 +325,19 @@ def check_version_compatibility(
|
|||||||
logging.warning(FUTURE_MESSAGE.format(repo_id=repo_id, version=v_check))
|
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.
|
"""Return available valid versions (branches and tags) on a given Hub repo.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
repo_id (str): The repository ID on the Hugging Face Hub.
|
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:
|
Returns:
|
||||||
list[packaging.version.Version]: A list of valid versions found.
|
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 = api.list_repo_refs(repo_id, repo_type="dataset")
|
||||||
repo_refs = [b.name for b in repo_refs.branches + repo_refs.tags]
|
repo_refs = [b.name for b in repo_refs.branches + repo_refs.tags]
|
||||||
repo_versions = []
|
repo_versions = []
|
||||||
@@ -345,7 +348,12 @@ def get_repo_versions(repo_id: str) -> list[packaging.version.Version]:
|
|||||||
return repo_versions
|
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.
|
"""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
|
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:
|
Args:
|
||||||
repo_id (str): The repository ID on the Hugging Face Hub.
|
repo_id (str): The repository ID on the Hugging Face Hub.
|
||||||
version (str | packaging.version.Version): The target version.
|
version (str | packaging.version.Version): The target version.
|
||||||
|
token: Authentication token forwarded to the Hub version lookup.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
str: The safe version string (e.g., "v1.2.3") to use as a revision.
|
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 = (
|
target_version = (
|
||||||
packaging.version.parse(version) if not isinstance(version, packaging.version.Version) else 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:
|
if not hub_versions:
|
||||||
raise RevisionNotFoundError(
|
raise RevisionNotFoundError(
|
||||||
|
|||||||
@@ -58,6 +58,9 @@ class BiSOFollower(BimanualMixin, Robot):
|
|||||||
port=config.left_arm_config.port,
|
port=config.left_arm_config.port,
|
||||||
disable_torque_on_disconnect=config.left_arm_config.disable_torque_on_disconnect,
|
disable_torque_on_disconnect=config.left_arm_config.disable_torque_on_disconnect,
|
||||||
max_relative_target=config.left_arm_config.max_relative_target,
|
max_relative_target=config.left_arm_config.max_relative_target,
|
||||||
|
position_p_coefficient=config.left_arm_config.position_p_coefficient,
|
||||||
|
position_i_coefficient=config.left_arm_config.position_i_coefficient,
|
||||||
|
position_d_coefficient=config.left_arm_config.position_d_coefficient,
|
||||||
use_degrees=config.left_arm_config.use_degrees,
|
use_degrees=config.left_arm_config.use_degrees,
|
||||||
cameras=left_arm_cameras,
|
cameras=left_arm_cameras,
|
||||||
)
|
)
|
||||||
@@ -68,6 +71,9 @@ class BiSOFollower(BimanualMixin, Robot):
|
|||||||
port=config.right_arm_config.port,
|
port=config.right_arm_config.port,
|
||||||
disable_torque_on_disconnect=config.right_arm_config.disable_torque_on_disconnect,
|
disable_torque_on_disconnect=config.right_arm_config.disable_torque_on_disconnect,
|
||||||
max_relative_target=config.right_arm_config.max_relative_target,
|
max_relative_target=config.right_arm_config.max_relative_target,
|
||||||
|
position_p_coefficient=config.right_arm_config.position_p_coefficient,
|
||||||
|
position_i_coefficient=config.right_arm_config.position_i_coefficient,
|
||||||
|
position_d_coefficient=config.right_arm_config.position_d_coefficient,
|
||||||
use_degrees=config.right_arm_config.use_degrees,
|
use_degrees=config.right_arm_config.use_degrees,
|
||||||
cameras=config.right_arm_config.cameras,
|
cameras=config.right_arm_config.cameras,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -41,6 +41,11 @@ class SOFollowerConfig:
|
|||||||
# Set to `True` for backward compatibility with previous policies/dataset
|
# Set to `True` for backward compatibility with previous policies/dataset
|
||||||
use_degrees: bool = True
|
use_degrees: bool = True
|
||||||
|
|
||||||
|
# Position-mode PID gains written to Feetech STS3215 motors at connect time.
|
||||||
|
position_p_coefficient: int = 16
|
||||||
|
position_i_coefficient: int = 0
|
||||||
|
position_d_coefficient: int = 32
|
||||||
|
|
||||||
|
|
||||||
@RobotConfig.register_subclass("so101_follower")
|
@RobotConfig.register_subclass("so101_follower")
|
||||||
@RobotConfig.register_subclass("so100_follower")
|
@RobotConfig.register_subclass("so100_follower")
|
||||||
|
|||||||
@@ -161,11 +161,9 @@ class SOFollower(Robot):
|
|||||||
self.bus.configure_motors()
|
self.bus.configure_motors()
|
||||||
for motor in self.bus.motors:
|
for motor in self.bus.motors:
|
||||||
self.bus.write("Operating_Mode", motor, OperatingMode.POSITION.value)
|
self.bus.write("Operating_Mode", motor, OperatingMode.POSITION.value)
|
||||||
# Set P_Coefficient to lower value to avoid shakiness (Default is 32)
|
self.bus.write("P_Coefficient", motor, self.config.position_p_coefficient)
|
||||||
self.bus.write("P_Coefficient", motor, 16)
|
self.bus.write("I_Coefficient", motor, self.config.position_i_coefficient)
|
||||||
# Set I_Coefficient and D_Coefficient to default value 0 and 32
|
self.bus.write("D_Coefficient", motor, self.config.position_d_coefficient)
|
||||||
self.bus.write("I_Coefficient", motor, 0)
|
|
||||||
self.bus.write("D_Coefficient", motor, 32)
|
|
||||||
|
|
||||||
if motor == "gripper":
|
if motor == "gripper":
|
||||||
self.bus.write("Max_Torque_Limit", motor, 500) # 50% of max torque to avoid burnout
|
self.bus.write("Max_Torque_Limit", motor, 500) # 50% of max torque to avoid burnout
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ import cv2
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from lerobot.cameras.configs import Cv2Rotation
|
from lerobot.cameras.configs import ColorMode, Cv2Rotation
|
||||||
from lerobot.cameras.opencv import OpenCVCamera, OpenCVCameraConfig
|
from lerobot.cameras.opencv import OpenCVCamera, OpenCVCameraConfig
|
||||||
from lerobot.utils.errors import DeviceAlreadyConnectedError, DeviceNotConnectedError
|
from lerobot.utils.errors import DeviceAlreadyConnectedError, DeviceNotConnectedError
|
||||||
|
|
||||||
@@ -132,6 +132,28 @@ def test_read(index_or_path):
|
|||||||
assert isinstance(img, np.ndarray)
|
assert isinstance(img, np.ndarray)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("index_or_path", TEST_IMAGE_PATHS, ids=TEST_IMAGE_SIZES)
|
||||||
|
def test_color_mode_conversion(index_or_path):
|
||||||
|
"""RGB and BGR reads of the same frame must differ only by a channel-axis reversal."""
|
||||||
|
rgb_config = OpenCVCameraConfig(index_or_path=index_or_path, color_mode=ColorMode.RGB, warmup_s=0)
|
||||||
|
bgr_config = OpenCVCameraConfig(index_or_path=index_or_path, color_mode=ColorMode.BGR, warmup_s=0)
|
||||||
|
with OpenCVCamera(rgb_config) as rgb_cam:
|
||||||
|
rgb = rgb_cam.read()
|
||||||
|
with OpenCVCamera(bgr_config) as bgr_cam:
|
||||||
|
bgr = bgr_cam.read()
|
||||||
|
|
||||||
|
assert rgb.shape == bgr.shape
|
||||||
|
np.testing.assert_array_equal(rgb, bgr[..., ::-1])
|
||||||
|
|
||||||
|
|
||||||
|
def test_postprocess_invalid_color_mode():
|
||||||
|
config = OpenCVCameraConfig(index_or_path=DEFAULT_PNG_FILE_PATH)
|
||||||
|
camera = OpenCVCamera(config)
|
||||||
|
camera.color_mode = "invalid"
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
camera._postprocess_image(np.zeros((120, 160, 3), dtype=np.uint8))
|
||||||
|
|
||||||
|
|
||||||
def test_read_before_connect():
|
def test_read_before_connect():
|
||||||
config = OpenCVCameraConfig(index_or_path=DEFAULT_PNG_FILE_PATH)
|
config = OpenCVCameraConfig(index_or_path=DEFAULT_PNG_FILE_PATH)
|
||||||
|
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ import pytest
|
|||||||
|
|
||||||
pytest.importorskip("reachy2_sdk")
|
pytest.importorskip("reachy2_sdk")
|
||||||
|
|
||||||
|
from lerobot.cameras.configs import ColorMode
|
||||||
from lerobot.cameras.reachy2_camera import Reachy2Camera, Reachy2CameraConfig
|
from lerobot.cameras.reachy2_camera import Reachy2Camera, Reachy2CameraConfig
|
||||||
from lerobot.utils.errors import DeviceNotConnectedError
|
from lerobot.utils.errors import DeviceNotConnectedError
|
||||||
|
|
||||||
@@ -33,28 +34,19 @@ PARAMS = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def _make_cam_manager_mock():
|
def _make_cam_manager_mock(color_frame, depth_frame=None):
|
||||||
c = MagicMock(name="CameraManagerMock")
|
c = MagicMock(name="CameraManagerMock")
|
||||||
|
|
||||||
teleop = MagicMock(name="TeleopCam")
|
teleop = MagicMock(name="TeleopCam")
|
||||||
teleop.width = 640
|
teleop.width = 640
|
||||||
teleop.height = 480
|
teleop.height = 480
|
||||||
teleop.get_frame = MagicMock(
|
teleop.get_frame = MagicMock(side_effect=lambda *_, **__: (color_frame, time.time()))
|
||||||
side_effect=lambda *_, **__: (
|
|
||||||
np.zeros((480, 640, 3), dtype=np.uint8),
|
|
||||||
time.time(),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
depth = MagicMock(name="DepthCam")
|
depth = MagicMock(name="DepthCam")
|
||||||
depth.width = 640
|
depth.width = 640
|
||||||
depth.height = 480
|
depth.height = 480
|
||||||
depth.get_frame = MagicMock(
|
depth.get_frame = MagicMock(side_effect=lambda *_, **__: (color_frame, time.time()))
|
||||||
side_effect=lambda *_, **__: (
|
depth.get_depth_frame = MagicMock(side_effect=lambda *_, **__: (depth_frame, time.time()))
|
||||||
np.zeros((480, 640, 3), dtype=np.uint8),
|
|
||||||
time.time(),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
c.is_connected.return_value = True
|
c.is_connected.return_value = True
|
||||||
c.teleop = teleop
|
c.teleop = teleop
|
||||||
@@ -84,12 +76,14 @@ def _make_cam_manager_mock():
|
|||||||
# ids=["teleop-left", "teleop-right", "torso-rgb", "torso-depth"],
|
# ids=["teleop-left", "teleop-right", "torso-rgb", "torso-depth"],
|
||||||
ids=["teleop-left", "teleop-right", "torso-rgb"],
|
ids=["teleop-left", "teleop-right", "torso-rgb"],
|
||||||
)
|
)
|
||||||
def camera(request):
|
def camera(request, img_array_factory):
|
||||||
name, image_type = request.param
|
name, image_type = request.param
|
||||||
|
color_frame = img_array_factory(height=480, width=640)
|
||||||
|
depth_frame = img_array_factory(height=480, width=640, channels=1, dtype=np.uint16)[..., 0]
|
||||||
with (
|
with (
|
||||||
patch(
|
patch(
|
||||||
"lerobot.cameras.reachy2_camera.reachy2_camera.CameraManager",
|
"lerobot.cameras.reachy2_camera.reachy2_camera.CameraManager",
|
||||||
side_effect=lambda *a, **k: _make_cam_manager_mock(),
|
side_effect=lambda *a, **k: _make_cam_manager_mock(color_frame, depth_frame),
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
config = Reachy2CameraConfig(name=name, image_type=image_type)
|
config = Reachy2CameraConfig(name=name, image_type=image_type)
|
||||||
@@ -188,6 +182,41 @@ def test_read_latest_too_old(camera):
|
|||||||
_ = camera.read_latest(max_age_ms=0) # immediately too old
|
_ = camera.read_latest(max_age_ms=0) # immediately too old
|
||||||
|
|
||||||
|
|
||||||
|
def test_color_mode_conversion(img_array_factory):
|
||||||
|
"""teleop frames are native BGR: RGB reverses the channel axis, BGR is passed through."""
|
||||||
|
frame = img_array_factory(height=8, width=8)
|
||||||
|
|
||||||
|
outputs = {}
|
||||||
|
for color_mode in (ColorMode.RGB, ColorMode.BGR):
|
||||||
|
with patch(
|
||||||
|
"lerobot.cameras.reachy2_camera.reachy2_camera.CameraManager",
|
||||||
|
side_effect=lambda *a, **k: _make_cam_manager_mock(frame),
|
||||||
|
):
|
||||||
|
cam = Reachy2Camera(Reachy2CameraConfig(name="teleop", image_type="left", color_mode=color_mode))
|
||||||
|
cam.connect()
|
||||||
|
outputs[color_mode] = cam.read()
|
||||||
|
cam.disconnect()
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(outputs[ColorMode.BGR], frame)
|
||||||
|
np.testing.assert_array_equal(outputs[ColorMode.RGB], frame[..., ::-1])
|
||||||
|
|
||||||
|
|
||||||
|
def test_depth_frame_not_color_converted(img_array_factory):
|
||||||
|
"""A depth/depth frame must be returned as-is, without BGR<->RGB conversion."""
|
||||||
|
color_frame = img_array_factory(height=8, width=8)
|
||||||
|
depth = img_array_factory(height=8, width=8, channels=1, dtype=np.uint16)[..., 0]
|
||||||
|
with patch(
|
||||||
|
"lerobot.cameras.reachy2_camera.reachy2_camera.CameraManager",
|
||||||
|
side_effect=lambda *a, **k: _make_cam_manager_mock(color_frame, depth_frame=depth),
|
||||||
|
):
|
||||||
|
cam = Reachy2Camera(Reachy2CameraConfig(name="depth", image_type="depth"))
|
||||||
|
cam.connect()
|
||||||
|
out = cam.read()
|
||||||
|
cam.disconnect()
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(out, depth)
|
||||||
|
|
||||||
|
|
||||||
def test_wrong_camera_name():
|
def test_wrong_camera_name():
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
_ = Reachy2CameraConfig(name="wrong-name", image_type="left")
|
_ = Reachy2CameraConfig(name="wrong-name", image_type="left")
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from unittest.mock import patch
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from lerobot.cameras.configs import Cv2Rotation
|
from lerobot.cameras.configs import ColorMode, Cv2Rotation
|
||||||
from lerobot.utils.errors import DeviceAlreadyConnectedError, DeviceNotConnectedError
|
from lerobot.utils.errors import DeviceAlreadyConnectedError, DeviceNotConnectedError
|
||||||
|
|
||||||
pytest.importorskip("pyrealsense2")
|
pytest.importorskip("pyrealsense2")
|
||||||
@@ -109,6 +109,32 @@ def test_read_depth():
|
|||||||
assert isinstance(img, np.ndarray)
|
assert isinstance(img, np.ndarray)
|
||||||
|
|
||||||
|
|
||||||
|
# These exercise _postprocess_image directly rather than read(): the bag playback returns
|
||||||
|
# non-deterministic frames we can't compare against, and the depth read() path is skipped
|
||||||
|
# (see test_read_depth) with the current pyrealsense2 version.
|
||||||
|
def test_color_mode_conversion(img_array_factory):
|
||||||
|
"""RGB (native for RealSense) is passed through; BGR reverses the channel axis."""
|
||||||
|
color = img_array_factory(height=3, width=4)
|
||||||
|
|
||||||
|
outputs = {}
|
||||||
|
for color_mode in (ColorMode.RGB, ColorMode.BGR):
|
||||||
|
camera = RealSenseCamera(RealSenseCameraConfig(serial_number_or_name="042", color_mode=color_mode))
|
||||||
|
camera.capture_height, camera.capture_width = color.shape[:2]
|
||||||
|
outputs[color_mode] = camera._postprocess_image(color)
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(outputs[ColorMode.RGB], color)
|
||||||
|
np.testing.assert_array_equal(outputs[ColorMode.BGR], color[..., ::-1])
|
||||||
|
|
||||||
|
|
||||||
|
def test_depth_frame_not_color_converted(img_array_factory):
|
||||||
|
"""Depth frames must bypass color conversion, even when a BGR color_mode is set."""
|
||||||
|
camera = RealSenseCamera(RealSenseCameraConfig(serial_number_or_name="042", color_mode=ColorMode.BGR))
|
||||||
|
depth = img_array_factory(height=3, width=4, channels=1, dtype=np.uint16)[..., 0]
|
||||||
|
camera.capture_height, camera.capture_width = depth.shape
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(camera._postprocess_image(depth, depth_frame=True), depth)
|
||||||
|
|
||||||
|
|
||||||
def test_read_before_connect():
|
def test_read_before_connect():
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042")
|
config = RealSenseCameraConfig(serial_number_or_name="042")
|
||||||
camera = RealSenseCamera(config)
|
camera = RealSenseCamera(config)
|
||||||
|
|||||||
@@ -14,16 +14,21 @@
|
|||||||
# 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 unittest.mock import Mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
from packaging.version import Version
|
||||||
|
|
||||||
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
||||||
|
|
||||||
from datasets import Dataset # noqa: E402
|
from datasets import Dataset # noqa: E402
|
||||||
from huggingface_hub import DatasetCard
|
from huggingface_hub import DatasetCard
|
||||||
|
|
||||||
|
import lerobot.datasets.utils as dataset_utils
|
||||||
from lerobot.datasets.io_utils import hf_transform_to_torch
|
from lerobot.datasets.io_utils import hf_transform_to_torch
|
||||||
from lerobot.datasets.utils import create_lerobot_dataset_card
|
from lerobot.datasets.utils import create_lerobot_dataset_card, get_repo_versions, get_safe_version
|
||||||
from lerobot.utils.constants import ACTION, OBS_IMAGES
|
from lerobot.utils.constants import ACTION, OBS_IMAGES
|
||||||
from lerobot.utils.feature_utils import combine_feature_dicts
|
from lerobot.utils.feature_utils import combine_feature_dicts
|
||||||
|
|
||||||
@@ -57,6 +62,30 @@ def test_default_parameters():
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
def test_get_repo_versions_forwards_token(monkeypatch, token):
|
||||||
|
api = Mock()
|
||||||
|
api.list_repo_refs.return_value = SimpleNamespace(
|
||||||
|
branches=[SimpleNamespace(name="v3.0")],
|
||||||
|
tags=[],
|
||||||
|
)
|
||||||
|
hf_api = Mock(return_value=api)
|
||||||
|
monkeypatch.setattr(dataset_utils, "HfApi", hf_api)
|
||||||
|
|
||||||
|
assert get_repo_versions("private/repo", token=token) == [Version("3.0")]
|
||||||
|
hf_api.assert_called_once_with(token=token)
|
||||||
|
api.list_repo_refs.assert_called_once_with("private/repo", repo_type="dataset")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
def test_get_safe_version_forwards_token(monkeypatch, token):
|
||||||
|
get_versions = Mock(return_value=[Version("3.0")])
|
||||||
|
monkeypatch.setattr(dataset_utils, "get_repo_versions", get_versions)
|
||||||
|
|
||||||
|
assert get_safe_version("private/repo", "v3.0", token=token) == "v3.0"
|
||||||
|
get_versions.assert_called_once_with("private/repo", token=token)
|
||||||
|
|
||||||
|
|
||||||
def test_with_tags():
|
def test_with_tags():
|
||||||
tags = ["tag1", "tag2"]
|
tags = ["tag1", "tag2"]
|
||||||
card = create_lerobot_dataset_card(tags=tags)
|
card = create_lerobot_dataset_card(tags=tags)
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ property delegation, and the full create-record-finalize-read lifecycle.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -191,6 +192,48 @@ def test_metadata_without_root_uses_hub_cache_snapshot_download(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
def test_metadata_download_forwards_token(tmp_path, monkeypatch, token):
|
||||||
|
snapshot_root = tmp_path / "snapshot"
|
||||||
|
snapshot_download = Mock(return_value=str(snapshot_root))
|
||||||
|
get_safe_version = Mock(return_value="v3.0")
|
||||||
|
load_metadata = Mock(side_effect=[FileNotFoundError, None])
|
||||||
|
monkeypatch.setattr(dataset_metadata_module, "snapshot_download", snapshot_download)
|
||||||
|
monkeypatch.setattr(dataset_metadata_module, "get_safe_version", get_safe_version)
|
||||||
|
monkeypatch.setattr(LeRobotDatasetMetadata, "_load_metadata", load_metadata)
|
||||||
|
|
||||||
|
meta = LeRobotDatasetMetadata(
|
||||||
|
repo_id=DUMMY_REPO_ID,
|
||||||
|
revision="v3.0",
|
||||||
|
token=token,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert meta.root == snapshot_root
|
||||||
|
assert not hasattr(meta, "_token")
|
||||||
|
get_safe_version.assert_called_once_with(DUMMY_REPO_ID, "v3.0", token=token)
|
||||||
|
assert snapshot_download.call_args.kwargs["token"] is token
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
def test_data_download_forwards_token(tmp_path, monkeypatch, token):
|
||||||
|
snapshot_root = tmp_path / "snapshot"
|
||||||
|
snapshot_download = Mock(return_value=str(snapshot_root))
|
||||||
|
monkeypatch.setattr(lerobot_dataset_module, "snapshot_download", snapshot_download)
|
||||||
|
|
||||||
|
dataset = LeRobotDataset.__new__(LeRobotDataset)
|
||||||
|
dataset.repo_id = DUMMY_REPO_ID
|
||||||
|
dataset.revision = "main"
|
||||||
|
dataset.episodes = None
|
||||||
|
dataset._requested_root = None
|
||||||
|
dataset.meta = SimpleNamespace(root=None)
|
||||||
|
dataset.reader = SimpleNamespace(root=None)
|
||||||
|
|
||||||
|
dataset._download(token=token)
|
||||||
|
|
||||||
|
assert dataset.root == snapshot_root
|
||||||
|
assert snapshot_download.call_args.kwargs["token"] is token
|
||||||
|
|
||||||
|
|
||||||
def test_without_root_reads_different_revisions_from_distinct_snapshot_roots(
|
def test_without_root_reads_different_revisions_from_distinct_snapshot_roots(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
info_factory,
|
info_factory,
|
||||||
|
|||||||
@@ -13,16 +13,53 @@
|
|||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
# 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 unittest.mock import Mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
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
|
||||||
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
|
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
|
||||||
from lerobot.utils.constants import ACTION
|
from lerobot.utils.constants import ACTION
|
||||||
from tests.fixtures.constants import DUMMY_REPO_ID
|
from tests.fixtures.constants import DUMMY_REPO_ID
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
@pytest.mark.parametrize("from_local", [False, True])
|
||||||
|
def test_streaming_dataset_forwards_token_to_metadata_without_retaining_it(
|
||||||
|
tmp_path, monkeypatch, token, from_local
|
||||||
|
):
|
||||||
|
requested_root = tmp_path / "local" if from_local else None
|
||||||
|
metadata = SimpleNamespace(
|
||||||
|
repo_id=DUMMY_REPO_ID,
|
||||||
|
root=requested_root or tmp_path / "snapshot",
|
||||||
|
revision=streaming_dataset_module.CODEBASE_VERSION,
|
||||||
|
_version=streaming_dataset_module.CODEBASE_VERSION,
|
||||||
|
features={},
|
||||||
|
total_episodes=0,
|
||||||
|
video_keys=[],
|
||||||
|
depth_keys=[],
|
||||||
|
image_keys=[],
|
||||||
|
rescale_depth_stats=Mock(),
|
||||||
|
)
|
||||||
|
metadata_cls = Mock(return_value=metadata)
|
||||||
|
monkeypatch.setattr(streaming_dataset_module, "LeRobotDatasetMetadata", metadata_cls)
|
||||||
|
|
||||||
|
dataset = StreamingLeRobotDataset(DUMMY_REPO_ID, root=requested_root, token=token)
|
||||||
|
|
||||||
|
metadata_cls.assert_called_once_with(
|
||||||
|
DUMMY_REPO_ID,
|
||||||
|
requested_root,
|
||||||
|
streaming_dataset_module.CODEBASE_VERSION,
|
||||||
|
force_cache_sync=False,
|
||||||
|
token=token,
|
||||||
|
)
|
||||||
|
assert not hasattr(dataset, "_token")
|
||||||
|
|
||||||
|
|
||||||
def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
|
def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
|
||||||
"""Test if are correctly accessed"""
|
"""Test if are correctly accessed"""
|
||||||
ds_num_frames = 400
|
ds_num_frames = 400
|
||||||
|
|||||||
@@ -109,3 +109,22 @@ def test_send_action(follower):
|
|||||||
|
|
||||||
goal_pos = {m: (i + 1) * 10 for i, m in enumerate(follower.bus.motors)}
|
goal_pos = {m: (i + 1) * 10 for i, m in enumerate(follower.bus.motors)}
|
||||||
follower.bus.sync_write.assert_called_once_with("Goal_Position", goal_pos)
|
follower.bus.sync_write.assert_called_once_with("Goal_Position", goal_pos)
|
||||||
|
|
||||||
|
|
||||||
|
def test_configure_writes_position_pid_coefficients():
|
||||||
|
bus_mock = _make_bus_mock()
|
||||||
|
bus_mock.motors = ["shoulder_pan"]
|
||||||
|
robot = MagicMock()
|
||||||
|
robot.bus = bus_mock
|
||||||
|
robot.config = SO100FollowerConfig(
|
||||||
|
port="/dev/null",
|
||||||
|
position_p_coefficient=32,
|
||||||
|
position_i_coefficient=1,
|
||||||
|
position_d_coefficient=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
SO100Follower.configure(robot)
|
||||||
|
|
||||||
|
bus_mock.write.assert_any_call("P_Coefficient", "shoulder_pan", 32)
|
||||||
|
bus_mock.write.assert_any_call("I_Coefficient", "shoulder_pan", 1)
|
||||||
|
bus_mock.write.assert_any_call("D_Coefficient", "shoulder_pan", 16)
|
||||||
|
|||||||
Reference in New Issue
Block a user