mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 04:36:04 +00:00
Compare commits
32 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c266fe4f1a | |||
| 913664c320 | |||
| c1b6ea85d6 | |||
| ffe25afb8f | |||
| 3f093d8927 | |||
| 95211b98f1 | |||
| 95256d766d | |||
| fd53716688 | |||
| a96540a2c4 | |||
| acd42b4d85 | |||
| bbeacfe57d | |||
| 801346e18c | |||
| ab87fd9764 | |||
| 6c57dfd2ee | |||
| d63e6e67a5 | |||
| 0d383d09f2 | |||
| ab2b5b04dd | |||
| ac5c7b8600 | |||
| a6befef0ba | |||
| 53843007ea | |||
| d3bed0feee | |||
| a0eb860d1e | |||
| cfd9ff969c | |||
| f59eae4e27 | |||
| a993af9c51 | |||
| 392246feaf | |||
| 19dcbc19f1 | |||
| 679faeaafc | |||
| 228cb5ddb9 | |||
| ad176c6d41 | |||
| d6c605e8c5 | |||
| 9c82c39c7b |
@@ -0,0 +1,11 @@
|
|||||||
|
version: 2
|
||||||
|
updates:
|
||||||
|
- package-ecosystem: "github-actions"
|
||||||
|
directory: "/"
|
||||||
|
schedule:
|
||||||
|
interval: "weekly"
|
||||||
|
cooldown:
|
||||||
|
default-days: 7
|
||||||
|
groups:
|
||||||
|
actions:
|
||||||
|
patterns: ["*"]
|
||||||
@@ -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`).
|
||||||
|
|||||||
@@ -83,7 +83,7 @@ episode_index=0
|
|||||||
print(f"{dataset[episode_index]['action'].shape=}\n")
|
print(f"{dataset[episode_index]['action'].shape=}\n")
|
||||||
```
|
```
|
||||||
|
|
||||||
Learn more about it in the [LeRobotDataset Documentation](https://huggingface.co/docs/lerobot/lerobot-dataset-v3)
|
Learn more about it in the [LeRobotDataset Documentation](https://huggingface.co/docs/lerobot/lerobot-dataset-v3).
|
||||||
|
|
||||||
## SoTA Models
|
## SoTA Models
|
||||||
|
|
||||||
@@ -109,7 +109,7 @@ lerobot-train \
|
|||||||
| **World Models** | [VLA-JEPA](./docs/source/vla_jepa.mdx), [LingBot-VA](./docs/source/lingbot_va.mdx), [FastWAM](./docs/source/fastwam.mdx) |
|
| **World Models** | [VLA-JEPA](./docs/source/vla_jepa.mdx), [LingBot-VA](./docs/source/lingbot_va.mdx), [FastWAM](./docs/source/fastwam.mdx) |
|
||||||
| **Reward Models** | [SARM](./docs/source/sarm.mdx), [TOPReward](./docs/source/topreward.mdx), [Robometer](./docs/source/robometer.mdx) |
|
| **Reward Models** | [SARM](./docs/source/sarm.mdx), [TOPReward](./docs/source/topreward.mdx), [Robometer](./docs/source/robometer.mdx) |
|
||||||
|
|
||||||
Similarly to the hardware, you can easily implement your own policy & leverage LeRobot's data collection, training, and visualization tools, and share your model to the HF Hub
|
Similarly to the hardware, you can easily implement your own policy & leverage LeRobot's data collection, training, and visualization tools, and share your model to the HF Hub.
|
||||||
|
|
||||||
For detailed policy setup guides, see the [Policy Documentation](https://huggingface.co/docs/lerobot/bring_your_own_policies). For GPU/RAM requirements and expected training time per policy, see the [Compute Hardware Guide](https://huggingface.co/docs/lerobot/hardware_guide).
|
For detailed policy setup guides, see the [Policy Documentation](https://huggingface.co/docs/lerobot/bring_your_own_policies). For GPU/RAM requirements and expected training time per policy, see the [Compute Hardware Guide](https://huggingface.co/docs/lerobot/hardware_guide).
|
||||||
|
|
||||||
@@ -126,7 +126,7 @@ lerobot-eval \
|
|||||||
--eval.n_episodes=10
|
--eval.n_episodes=10
|
||||||
```
|
```
|
||||||
|
|
||||||
Learn how to implement your own simulation environment or benchmark and distribute it from the HF Hub by following the [EnvHub Documentation](https://huggingface.co/docs/lerobot/envhub)
|
Learn how to implement your own simulation environment or benchmark and distribute it from the HF Hub by following the [EnvHub Documentation](https://huggingface.co/docs/lerobot/envhub).
|
||||||
|
|
||||||
## Resources
|
## Resources
|
||||||
|
|
||||||
|
|||||||
@@ -89,8 +89,8 @@ subtask.
|
|||||||
|
|
||||||
The resulting spans are then stitched into a gap-free, full-episode
|
The resulting spans are then stitched into a gap-free, full-episode
|
||||||
cover, so **every frame has exactly one active subtask**. See
|
cover, so **every frame has exactly one active subtask**. See
|
||||||
[`run_hf_job.py`](https://github.com/huggingface/lerobot/blob/main/examples/annotations/run_hf_job.py)
|
[Running on Hugging Face Jobs](#running-on-hugging-face-jobs) for the
|
||||||
for the production settings (single camera, timestamped contact sheets,
|
production settings (single camera, timestamped contact sheets,
|
||||||
auto-windowed subtask generation).
|
auto-windowed subtask generation).
|
||||||
|
|
||||||
### Tools
|
### Tools
|
||||||
@@ -110,28 +110,67 @@ not-yet-implemented.
|
|||||||
|
|
||||||
## Running on Hugging Face Jobs
|
## Running on Hugging Face Jobs
|
||||||
|
|
||||||
Annotation runs on [Hugging Face Jobs](https://huggingface.co/docs/hub/en/jobs).
|
Annotating a real dataset needs a GPU big enough to serve the VLM, so
|
||||||
The repo ships a launcher script you copy and tweak for your dataset:
|
`lerobot-annotate` can dispatch itself to
|
||||||
|
[Hugging Face Jobs](https://huggingface.co/docs/hub/en/jobs) — same as
|
||||||
|
`lerobot-train`. Add `--job.target=<flavor>` to the exact command you'd
|
||||||
|
run locally and it runs on that hardware instead:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
HF_TOKEN=hf_... uv run python examples/annotations/run_hf_job.py
|
hf auth login # once
|
||||||
|
|
||||||
|
uv run lerobot-annotate \
|
||||||
|
--repo_id=user/my_dataset \
|
||||||
|
--new_repo_id=user/my_dataset_annotated \
|
||||||
|
--push_to_hub=true \
|
||||||
|
--vlm.model_id=Qwen/Qwen3.6-27B \
|
||||||
|
--vlm.num_gpus=1 \
|
||||||
|
--vlm.serve_command="vllm serve Qwen/Qwen3.6-27B --tensor-parallel-size 1 \
|
||||||
|
--max-model-len 32768 --gpu-memory-utilization 0.8 \
|
||||||
|
--uvicorn-log-level warning --port {port}" \
|
||||||
|
--vlm.serve_ready_timeout_s=1800 \
|
||||||
|
--vlm.chat_template_kwargs='{"enable_thinking": false}' \
|
||||||
|
--job.target=h200
|
||||||
```
|
```
|
||||||
|
|
||||||
[`run_hf_job.py`](https://github.com/huggingface/lerobot/blob/main/examples/annotations/run_hf_job.py)
|
That submits a single-GPU `h200` job that:
|
||||||
starts a single-GPU `h200` job (bump it to `h200x4` for big datasets)
|
|
||||||
that:
|
|
||||||
|
|
||||||
1. installs `lerobot` (from `main`) plus the annotation extras,
|
1. starts from the `vllm/vllm-openai` image and installs `lerobot` on top,
|
||||||
2. boots one vLLM server per GPU (using the `vllm/vllm-openai` image) and
|
2. boots one vLLM server per GPU and drives it over the OpenAI-compatible API,
|
||||||
drives it over the OpenAI-compatible API,
|
3. runs the `plan` / `interjections` / `vqa` modules across the dataset,
|
||||||
3. runs the `plan` / `interjections` / `vqa` modules across the dataset
|
|
||||||
with `lerobot-annotate`,
|
|
||||||
4. with `--push_to_hub=true`, uploads the result to `--new_repo_id` (or
|
4. with `--push_to_hub=true`, uploads the result to `--new_repo_id` (or
|
||||||
back to `--repo_id` in place if you leave that unset).
|
back to `--repo_id` in place if you leave that unset).
|
||||||
|
|
||||||
To use a different dataset, model, or hub repo, edit the `CMD` block in
|
The command streams the job's logs; `Ctrl-C` detaches without cancelling
|
||||||
the script. Every flag there maps directly to a `lerobot-annotate` flag
|
it. List the available flavors and their pricing with `hf jobs hardware`.
|
||||||
(run `lerobot-annotate --help` for the full list).
|
|
||||||
|
<Tip warning={true}>
|
||||||
|
|
||||||
|
Qwen3.6 ships with thinking enabled, which eats the token budget the
|
||||||
|
annotator needs for its JSON answer — `--vlm.chat_template_kwargs='{"enable_thinking": false}'`
|
||||||
|
turns it off. Without `--push_to_hub=true` the annotated dataset is
|
||||||
|
discarded when the pod exits.
|
||||||
|
|
||||||
|
</Tip>
|
||||||
|
|
||||||
|
### Job options
|
||||||
|
|
||||||
|
| Flag | Default | What it does |
|
||||||
|
| ------------------- | ------------------------- | ------------------------------------------------------------------------------- |
|
||||||
|
| `--job.target` | `local` | HF Jobs flavor to run on (e.g. `h200`, `h200x4`). Omitted/`local` runs here. |
|
||||||
|
| `--job.image` | `vllm/vllm-openai:latest` | Runtime image for the pod. |
|
||||||
|
| `--job.timeout` | `2h` | Wall-clock cap. Raise it for large datasets. |
|
||||||
|
| `--job.detach` | `false` | Submit and exit instead of streaming logs. |
|
||||||
|
| `--job.lerobot_ref` | `main` | Git ref of lerobot installed on the pod — point it at a branch to test changes. |
|
||||||
|
| `--job.tags` | `[]` | Extra tags on the job and on any dataset it pushes (`lerobot` is always added). |
|
||||||
|
|
||||||
|
For a bigger dataset, scale to `h200x4` and raise
|
||||||
|
`--vlm.parallel_servers` / `--vlm.num_gpus` to match, and give the job
|
||||||
|
more headroom with e.g. `--job.timeout=8h`.
|
||||||
|
|
||||||
|
Remote runs need `--repo_id` (the pod pulls the dataset from the Hub;
|
||||||
|
`--root` names a directory only your machine has). A dataset that exists
|
||||||
|
only in your local cache is pushed to a **private** repo first.
|
||||||
|
|
||||||
## Key options
|
## Key options
|
||||||
|
|
||||||
|
|||||||
@@ -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).
|
||||||
|
|||||||
@@ -141,6 +141,17 @@ sample["target_message_indices"]
|
|||||||
|
|
||||||
The renderer does not apply a tokenizer chat template. Policy processors decide how to serialize the messages for their backbone, which keeps the same dataset usable across SmolVLA, Pi0.5, and any future VLM that expects OpenAI-style chat messages.
|
The renderer does not apply a tokenizer chat template. Policy processors decide how to serialize the messages for their backbone, which keeps the same dataset usable across SmolVLA, Pi0.5, and any future VLM that expects OpenAI-style chat messages.
|
||||||
|
|
||||||
|
## Blends
|
||||||
|
|
||||||
|
Blend recipes select one weighted sub-recipe deterministically from the sample index.
|
||||||
|
`recipes/subtask_mem.yaml` trains the compact core blend — high-level subtask prediction, low-level execution, and memory. `recipes/subtask_mem_vqa_speech.yaml` is the fuller variant that also adds VQA and spoken interjection responses.
|
||||||
|
|
||||||
|
A message recipe with a supervised assistant turn on the `low_level` stream trains
|
||||||
|
the π0.5 paper's joint sequence instead of a blend: the target span gets text CE
|
||||||
|
while also conditioning the action losses in the same forward.
|
||||||
|
`recipes/subtask_joint.yaml` is the provided example; pair it with
|
||||||
|
`--policy.joint_subtask_conditioning=true` at inference.
|
||||||
|
|
||||||
## Graceful absence
|
## Graceful absence
|
||||||
|
|
||||||
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
||||||
|
|||||||
@@ -1,3 +1,11 @@
|
|||||||
|
# OMX
|
||||||
|
|
||||||
|
<img
|
||||||
|
src="https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/lerobot/omx_mainimage.png"
|
||||||
|
alt="OMX"
|
||||||
|
width=600
|
||||||
|
/>
|
||||||
|
|
||||||
## Order and Assemble the parts
|
## Order and Assemble the parts
|
||||||
|
|
||||||
First, assemble the OMX hardware following the official assembly guide.
|
First, assemble the OMX hardware following the official assembly guide.
|
||||||
|
|||||||
@@ -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,80 +0,0 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
"""Launch ``lerobot-annotate`` on a Hugging Face job (vllm + Qwen3.6-27B VLM).
|
|
||||||
|
|
||||||
Spawns one single-GPU ``h200`` job that:
|
|
||||||
|
|
||||||
1. installs ``lerobot`` from ``main`` plus the annotation extras,
|
|
||||||
2. boots one vllm server with Qwen3.6-27B (dense VLM),
|
|
||||||
3. runs the plan / interjections / vqa modules across the dataset
|
|
||||||
in free-form mode (each episode generates its own subtasks +
|
|
||||||
memory),
|
|
||||||
4. uploads the annotated dataset to ``--new_repo_id`` (when set)
|
|
||||||
or back to ``--repo_id``.
|
|
||||||
|
|
||||||
Usage:
|
|
||||||
|
|
||||||
HF_TOKEN=hf_... uv run python examples/annotations/run_hf_job.py
|
|
||||||
|
|
||||||
Adjust ``CMD`` (dataset, model, hub repo) and ``flavor`` below for your
|
|
||||||
run. For larger datasets, scale to ``h200x4`` and raise
|
|
||||||
``--vlm.parallel_servers`` / ``--vlm.num_gpus`` to match.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import os
|
|
||||||
|
|
||||||
from huggingface_hub import get_token, run_job
|
|
||||||
|
|
||||||
token = os.environ.get("HF_TOKEN") or get_token()
|
|
||||||
if not token:
|
|
||||||
raise RuntimeError("No HF token. Run `huggingface-cli login` or `export HF_TOKEN=hf_...`")
|
|
||||||
|
|
||||||
CMD = (
|
|
||||||
"apt-get update -qq && apt-get install -y -qq git ffmpeg && "
|
|
||||||
"pip install --no-deps "
|
|
||||||
"'lerobot @ git+https://github.com/huggingface/lerobot.git@main' && "
|
|
||||||
# Pins mirror pyproject.toml — unpinned installs pull av 18 / datasets 5 /
|
|
||||||
# draccus 0.11, which break lerobot at import time.
|
|
||||||
"pip install --upgrade-strategy only-if-needed "
|
|
||||||
"'datasets>=4.7.0,<5.0.0' 'pyarrow>=21.0.0,<30.0.0' 'av>=15.0.0,<16.0.0' 'draccus==0.10.0' "
|
|
||||||
"'pandas>=2.0.0,<3.0.0' jsonlines gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
|
||||||
"openai && "
|
|
||||||
"export VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS=0 && "
|
|
||||||
"export VLLM_VIDEO_BACKEND=pyav && "
|
|
||||||
"lerobot-annotate "
|
|
||||||
"--repo_id=pepijn223/robocasa_pretrain_human300_v4 "
|
|
||||||
"--new_repo_id=pepijn223/robocasa_pretrain_human300_v4_annotated "
|
|
||||||
"--push_to_hub=true "
|
|
||||||
"--vlm.backend=openai "
|
|
||||||
"--vlm.model_id=Qwen/Qwen3.6-27B "
|
|
||||||
"--vlm.num_gpus=1 "
|
|
||||||
'--vlm.serve_command="vllm serve Qwen/Qwen3.6-27B '
|
|
||||||
"--tensor-parallel-size 1 --max-model-len 32768 "
|
|
||||||
'--gpu-memory-utilization 0.8 --uvicorn-log-level warning --port {port}" '
|
|
||||||
"--vlm.serve_ready_timeout_s=1800 "
|
|
||||||
# Qwen3.6 ships with thinking on; annotation wants plain JSON answers.
|
|
||||||
"--vlm.chat_template_kwargs='{\"enable_thinking\": false}'"
|
|
||||||
)
|
|
||||||
|
|
||||||
job = run_job(
|
|
||||||
image="vllm/vllm-openai:latest",
|
|
||||||
command=["bash", "-c", CMD],
|
|
||||||
flavor="h200",
|
|
||||||
secrets={"HF_TOKEN": token},
|
|
||||||
timeout="2h",
|
|
||||||
)
|
|
||||||
print(f"Job URL: {job.url}")
|
|
||||||
print(f"Job ID: {job.id}")
|
|
||||||
+1
-1
@@ -155,7 +155,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"]
|
||||||
|
|||||||
@@ -20,6 +20,29 @@ from dataclasses import dataclass, field
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
from lerobot.configs.default import JobConfig
|
||||||
|
|
||||||
|
# The annotation pipeline boots its own vLLM server, so the pod starts from the
|
||||||
|
# official vLLM runtime rather than the prebuilt `lerobot-gpu` training image;
|
||||||
|
# `lerobot` is pip-installed on top (see `lerobot.jobs.annotate`).
|
||||||
|
DEFAULT_ANNOTATE_JOB_IMAGE = "vllm/vllm-openai:latest"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class AnnotationJobConfig(JobConfig):
|
||||||
|
"""`JobConfig` with the annotation runtime's defaults.
|
||||||
|
|
||||||
|
Adds `lerobot_ref` because the vLLM image ships no lerobot: the pod installs
|
||||||
|
it from git, and the ref decides which code actually annotates. Point it at a
|
||||||
|
branch/tag/SHA to try unmerged changes remotely.
|
||||||
|
"""
|
||||||
|
|
||||||
|
image: str = DEFAULT_ANNOTATE_JOB_IMAGE
|
||||||
|
# Annotation is a bounded pass over a dataset; a tighter cap than training's
|
||||||
|
# "2d" keeps a wedged vLLM server from burning a day of GPU time.
|
||||||
|
timeout: str | None = "2h"
|
||||||
|
lerobot_ref: str = "main"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PlanConfig:
|
class PlanConfig:
|
||||||
@@ -207,6 +230,11 @@ class AnnotationPipelineConfig:
|
|||||||
vlm: VlmConfig = field(default_factory=VlmConfig)
|
vlm: VlmConfig = field(default_factory=VlmConfig)
|
||||||
executor: ExecutorConfig = field(default_factory=ExecutorConfig)
|
executor: ExecutorConfig = field(default_factory=ExecutorConfig)
|
||||||
|
|
||||||
|
# Where the annotation runs: omitted / "local" annotates on this machine, any
|
||||||
|
# other value is an HF Jobs flavor (e.g. "h200") and submits the run there.
|
||||||
|
# List flavors + pricing with `hf jobs hardware`.
|
||||||
|
job: AnnotationJobConfig = field(default_factory=AnnotationJobConfig)
|
||||||
|
|
||||||
skip_validation: bool = False
|
skip_validation: bool = False
|
||||||
only_episodes: tuple[int, ...] | None = None
|
only_episodes: tuple[int, ...] | None = None
|
||||||
|
|
||||||
|
|||||||
@@ -30,8 +30,8 @@ Phase 3 is why the ``plan`` module must be re-entered after the
|
|||||||
timestamps.
|
timestamps.
|
||||||
|
|
||||||
Distributed execution is provided by Hugging Face Jobs (see
|
Distributed execution is provided by Hugging Face Jobs (see
|
||||||
``examples/annotations/run_hf_job.py``); the runner inside the job
|
``lerobot.jobs.annotate``, reached via ``--job.target=<flavor>``); the pod
|
||||||
invokes ``lerobot-annotate`` which uses this in-process executor.
|
inside the job invokes ``lerobot-annotate`` which uses this in-process executor.
|
||||||
Episode-level concurrency is controlled by
|
Episode-level concurrency is controlled by
|
||||||
``ExecutorConfig.episode_parallelism``.
|
``ExecutorConfig.episode_parallelism``.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -194,12 +194,13 @@ def make_vlm_client(config: VlmConfig) -> VlmClient:
|
|||||||
"""Build the shared VLM client.
|
"""Build the shared VLM client.
|
||||||
|
|
||||||
Only the ``openai`` backend is supported for now. The shipped workflow
|
Only the ``openai`` backend is supported for now. The shipped workflow
|
||||||
is Hugging Face Jobs (``examples/annotations/run_hf_job.py``): it boots
|
is Hugging Face Jobs (``lerobot-annotate --job.target=<flavor>``): it
|
||||||
a vLLM server inside the ``vllm/vllm-openai`` image and the pipeline
|
boots a vLLM server inside the ``vllm/vllm-openai`` image and the
|
||||||
talks to it over the OpenAI-compatible API (``--vlm.backend=openai``,
|
pipeline talks to it over the OpenAI-compatible API
|
||||||
optionally auto-spawning the server via ``auto_serve`` /
|
(``--vlm.backend=openai``, optionally auto-spawning the server via
|
||||||
``serve_command``). The former in-process ``vllm`` / ``transformers``
|
``auto_serve`` / ``serve_command``). The former in-process ``vllm`` /
|
||||||
backends were removed to keep the support surface to the HF Jobs path.
|
``transformers`` backends were removed to keep the support surface to
|
||||||
|
the HF Jobs path.
|
||||||
|
|
||||||
For ``stub``, construct :class:`StubVlmClient` directly with a responder
|
For ``stub``, construct :class:`StubVlmClient` directly with a responder
|
||||||
callable; it is rejected here to make accidental misuse obvious.
|
callable; it is rejected here to make accidental misuse obvious.
|
||||||
@@ -213,8 +214,8 @@ def make_vlm_client(config: VlmConfig) -> VlmClient:
|
|||||||
if config.backend in {"vllm", "transformers"}:
|
if config.backend in {"vllm", "transformers"}:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"backend={config.backend!r} (in-process local model) is not supported for now — "
|
f"backend={config.backend!r} (in-process local model) is not supported for now — "
|
||||||
"only backend='openai' (the Hugging Face Jobs flow) is. Run the pipeline via "
|
"only backend='openai' (the Hugging Face Jobs flow) is. Run the pipeline with "
|
||||||
"examples/annotations/run_hf_job.py, which serves the model with vLLM in the "
|
"`lerobot-annotate --job.target=<flavor>`, which serves the model with vLLM in the "
|
||||||
"vllm/vllm-openai image and talks to it over the OpenAI-compatible API."
|
"vllm/vllm-openai image and talks to it over the OpenAI-compatible API."
|
||||||
)
|
)
|
||||||
raise ValueError(f"Unknown VLM backend: {config.backend!r}")
|
raise ValueError(f"Unknown VLM backend: {config.backend!r}")
|
||||||
|
|||||||
@@ -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]:
|
||||||
|
|||||||
@@ -33,6 +33,8 @@ class DatasetConfig:
|
|||||||
# looked up under $HF_LEROBOT_HOME/repo_id and Hub downloads use a revision-safe cache under $HF_LEROBOT_HOME/hub.
|
# looked up under $HF_LEROBOT_HOME/repo_id and Hub downloads use a revision-safe cache under $HF_LEROBOT_HOME/hub.
|
||||||
root: str | None = None
|
root: str | None = None
|
||||||
episodes: list[int] | None = None
|
episodes: list[int] | None = None
|
||||||
|
# Episode indices to drop (e.g. corrupt or heterogeneous ones). Applied on top of `episodes`.
|
||||||
|
exclude_episodes: list[int] | None = None
|
||||||
image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig)
|
image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig)
|
||||||
revision: str | None = None
|
revision: str | None = None
|
||||||
use_imagenet_stats: bool = True
|
use_imagenet_stats: bool = True
|
||||||
@@ -62,6 +64,10 @@ class DatasetConfig:
|
|||||||
if len(self.episodes) != len(set(self.episodes)):
|
if len(self.episodes) != len(set(self.episodes)):
|
||||||
duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1})
|
duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1})
|
||||||
raise ValueError(f"Episode indices contain duplicates: {duplicates}")
|
raise ValueError(f"Episode indices contain duplicates: {duplicates}")
|
||||||
|
if self.exclude_episodes is not None and any(ep < 0 for ep in self.exclude_episodes):
|
||||||
|
raise ValueError(
|
||||||
|
f"exclude_episodes must be non-negative, got: {[ep for ep in self.exclude_episodes if ep < 0]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ class MessageTurn:
|
|||||||
raise ValueError(f"Unsupported message stream: {self.stream!r}")
|
raise ValueError(f"Unsupported message stream: {self.stream!r}")
|
||||||
if self.content is None and self.tool_calls_from is None:
|
if self.content is None and self.tool_calls_from is None:
|
||||||
raise ValueError("MessageTurn.content is required unless tool_calls_from is set.")
|
raise ValueError("MessageTurn.content is required unless tool_calls_from is set.")
|
||||||
if self.content is not None and not isinstance(self.content, (str, list)):
|
if self.content is not None and not isinstance(self.content, str | list):
|
||||||
raise TypeError("MessageTurn.content must be a string, a list of HF-style blocks, or None.")
|
raise TypeError("MessageTurn.content must be a string, a list of HF-style blocks, or None.")
|
||||||
if isinstance(self.content, list):
|
if isinstance(self.content, list):
|
||||||
for block in self.content:
|
for block in self.content:
|
||||||
@@ -147,7 +147,7 @@ class TrainingRecipe:
|
|||||||
return cls.from_dict(data)
|
return cls.from_dict(data)
|
||||||
|
|
||||||
def _validate_message_recipe(self) -> None:
|
def _validate_message_recipe(self) -> None:
|
||||||
"""Ensure every templated binding is known and at least one turn is a target."""
|
"""Validate bindings and require text or low-level action supervision."""
|
||||||
assert self.messages is not None
|
assert self.messages is not None
|
||||||
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
|
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
|
||||||
|
|
||||||
@@ -156,8 +156,14 @@ class TrainingRecipe:
|
|||||||
if missing:
|
if missing:
|
||||||
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
|
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
|
||||||
|
|
||||||
if not any(turn.target for turn in self.messages):
|
has_target = any(turn.target for turn in self.messages)
|
||||||
raise ValueError("Message recipes must contain at least one target turn.")
|
has_low_level = any(turn.stream == "low_level" for turn in self.messages)
|
||||||
|
if not (has_target or has_low_level):
|
||||||
|
raise ValueError(
|
||||||
|
"Message recipes must contain at least one supervised turn — "
|
||||||
|
"either ``target: true`` (text CE) or ``stream: low_level`` "
|
||||||
|
"(flow/action loss)."
|
||||||
|
)
|
||||||
|
|
||||||
def _validate_blend_recipe(self) -> None:
|
def _validate_blend_recipe(self) -> None:
|
||||||
"""Ensure each blend component is a non-empty, weighted message recipe."""
|
"""Ensure each blend component is a non-empty, weighted message recipe."""
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
# Predicts subtasks from tasks and trains subtask-conditioned action flow without memory or plans.
|
||||||
|
# Requires `subtask` annotations; samples with missing `if_present` bindings do not render.
|
||||||
|
|
||||||
|
blend:
|
||||||
|
|
||||||
|
high_level_subtask:
|
||||||
|
weight: 0.30
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||||
|
|
||||||
|
low_level_execution:
|
||||||
|
weight: 0.70
|
||||||
|
messages:
|
||||||
|
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||||
|
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
# Paper-style joint sequence (pi0.5 §IV-B): one sample supervises the subtask
|
||||||
|
# text with CE and, because the assistant turn is part of the prefix, conditions
|
||||||
|
# the FAST and flow action losses on the same annotated subtask in one forward.
|
||||||
|
# The supervised span is attended causally; the action losses see task + subtask.
|
||||||
|
#
|
||||||
|
# Pair with `--policy.joint_subtask_conditioning=true` at inference so the flow
|
||||||
|
# prefix reproduces this layout (task turn with state + causal generated subtask).
|
||||||
|
# Samples without a `subtask` annotation fall back to a plain task-prompt
|
||||||
|
# low-level sample via `if_present`.
|
||||||
|
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: low_level}
|
||||||
|
- {role: assistant, content: "${subtask}", stream: low_level, target: true, if_present: subtask}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
# Trains subtask prediction, subtask-conditioned action flow, and memory updates without plans.
|
||||||
|
# Requires `subtask` and `memory`; missing `if_present` bindings skip the affected sub-recipe.
|
||||||
|
|
||||||
|
blend:
|
||||||
|
|
||||||
|
high_level_subtask:
|
||||||
|
weight: 0.25
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||||
|
|
||||||
|
low_level_execution:
|
||||||
|
weight: 0.60
|
||||||
|
messages:
|
||||||
|
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||||
|
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||||
|
|
||||||
|
memory_update:
|
||||||
|
# `active_at` densifies sparse boundaries while preserving the prior-memory/subtask mapping.
|
||||||
|
# Inference controls update timing through `subtask_change` events.
|
||||||
|
weight: 0.15
|
||||||
|
bindings:
|
||||||
|
prior_memory: "nth_prev(style=memory, offset=1)"
|
||||||
|
current_memory: "active_at(t, style=memory)"
|
||||||
|
completed_subtask: "nth_prev(style=subtask, offset=1)"
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: assistant, content: "Previous memory: ${prior_memory}", stream: high_level, if_present: prior_memory}
|
||||||
|
- {role: user, content: "Completed subtask: ${completed_subtask}", stream: high_level, if_present: completed_subtask}
|
||||||
|
- {role: assistant, content: "${current_memory}", stream: high_level, target: true, if_present: current_memory}
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
# Adds memory, spoken interjection responses, and camera-grounded VQA to subtask/action training.
|
||||||
|
# Missing optional annotations skip only their sub-recipe; `say` tool calls tokenize as `<say>...</say>`.
|
||||||
|
|
||||||
|
blend:
|
||||||
|
|
||||||
|
high_level_subtask:
|
||||||
|
weight: 0.25
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||||
|
|
||||||
|
low_level_execution:
|
||||||
|
weight: 0.40
|
||||||
|
messages:
|
||||||
|
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||||
|
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||||
|
|
||||||
|
memory_update:
|
||||||
|
# `active_at` densifies sparse boundaries while preserving the prior-memory/subtask mapping.
|
||||||
|
# Inference controls update timing through `subtask_change` events.
|
||||||
|
weight: 0.10
|
||||||
|
bindings:
|
||||||
|
prior_memory: "nth_prev(style=memory, offset=1)"
|
||||||
|
current_memory: "active_at(t, style=memory)"
|
||||||
|
completed_subtask: "nth_prev(style=subtask, offset=1)"
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: assistant, content: "Previous memory: ${prior_memory}", stream: high_level, if_present: prior_memory}
|
||||||
|
- {role: user, content: "Completed subtask: ${completed_subtask}", stream: high_level, if_present: completed_subtask}
|
||||||
|
- {role: assistant, content: "${current_memory}", stream: high_level, target: true, if_present: current_memory}
|
||||||
|
|
||||||
|
user_interjection_response:
|
||||||
|
weight: 0.10
|
||||||
|
bindings:
|
||||||
|
interjection: "emitted_at(t, style=interjection)"
|
||||||
|
speech: "emitted_at(t, role=assistant, tool_name=say)"
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: high_level}
|
||||||
|
- {role: user, content: "${interjection}", stream: high_level, if_present: interjection}
|
||||||
|
# The assistant target is a `say` tool call flattened to a `<say>...</say>` marker.
|
||||||
|
- {role: assistant, stream: high_level, target: true, if_present: speech, tool_calls_from: speech}
|
||||||
|
|
||||||
|
# Each camera uses a separate VQA sub-recipe for view-specific binding.
|
||||||
|
ask_vqa_top:
|
||||||
|
weight: 0.075
|
||||||
|
bindings:
|
||||||
|
vqa_query: "emitted_at(t, style=vqa, role=user, camera=observation.images.front)"
|
||||||
|
vqa: "emitted_at(t, style=vqa, role=assistant, camera=observation.images.front)"
|
||||||
|
messages:
|
||||||
|
- role: user
|
||||||
|
stream: high_level
|
||||||
|
if_present: vqa_query
|
||||||
|
content:
|
||||||
|
- {type: image, feature: observation.images.front}
|
||||||
|
- {type: text, text: "${vqa_query}"}
|
||||||
|
- {role: assistant, content: "${vqa}", stream: high_level, target: true, if_present: vqa}
|
||||||
|
|
||||||
|
ask_vqa_wrist:
|
||||||
|
weight: 0.075
|
||||||
|
bindings:
|
||||||
|
vqa_query: "emitted_at(t, style=vqa, role=user, camera=observation.images.wrist)"
|
||||||
|
vqa: "emitted_at(t, style=vqa, role=assistant, camera=observation.images.wrist)"
|
||||||
|
messages:
|
||||||
|
- role: user
|
||||||
|
stream: high_level
|
||||||
|
if_present: vqa_query
|
||||||
|
content:
|
||||||
|
- {type: image, feature: observation.images.wrist}
|
||||||
|
- {type: text, text: "${vqa_query}"}
|
||||||
|
- {role: assistant, content: "${vqa}", stream: high_level, target: true, if_present: vqa}
|
||||||
@@ -14,6 +14,7 @@
|
|||||||
import builtins
|
import builtins
|
||||||
import datetime as dt
|
import datetime as dt
|
||||||
import json
|
import json
|
||||||
|
import multiprocessing
|
||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
@@ -101,6 +102,12 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
batch_size: int = 8
|
batch_size: int = 8
|
||||||
prefetch_factor: int = 4
|
prefetch_factor: int = 4
|
||||||
persistent_workers: bool = True
|
persistent_workers: bool = True
|
||||||
|
# DataLoader worker start method. "spawn" is safer than "fork" with
|
||||||
|
# non-fork-safe libs (PyAV / torchcodec / ffmpeg), but adds some
|
||||||
|
# worker-startup time per run since workers re-import modules instead
|
||||||
|
# of inheriting parent state. Override with `--dataloader_multiprocessing_context=fork`
|
||||||
|
# when appropriate, or set it to `null` to use Python's platform default.
|
||||||
|
dataloader_multiprocessing_context: str | None = "spawn"
|
||||||
steps: int = 100_000
|
steps: int = 100_000
|
||||||
# Run policy in the simulation environment every N steps to measure reward/success (0 = disabled).
|
# Run policy in the simulation environment every N steps to measure reward/success (0 = disabled).
|
||||||
env_eval_freq: int = 20_000
|
env_eval_freq: int = 20_000
|
||||||
@@ -212,6 +219,17 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
self.reward_model.pretrained_path = str(policy_dir)
|
self.reward_model.pretrained_path = str(policy_dir)
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
|
available_contexts = multiprocessing.get_all_start_methods()
|
||||||
|
if (
|
||||||
|
self.dataloader_multiprocessing_context is not None
|
||||||
|
and self.dataloader_multiprocessing_context not in available_contexts
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"`dataloader_multiprocessing_context` must be None or one of "
|
||||||
|
f"{available_contexts} on this platform, got "
|
||||||
|
f"{self.dataloader_multiprocessing_context!r}."
|
||||||
|
)
|
||||||
|
|
||||||
self._resolve_pretrained_from_cli()
|
self._resolve_pretrained_from_cli()
|
||||||
|
|
||||||
if self.policy is None and self.reward_model is None:
|
if self.policy is None and self.reward_model is None:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -163,10 +163,40 @@ class DatasetReader:
|
|||||||
def _load_hf_dataset(self) -> datasets.Dataset:
|
def _load_hf_dataset(self) -> datasets.Dataset:
|
||||||
"""hf_dataset contains all the observations, states, actions, rewards, etc."""
|
"""hf_dataset contains all the observations, states, actions, rewards, etc."""
|
||||||
features = get_hf_features_from_features(self._meta.features)
|
features = get_hf_features_from_features(self._meta.features)
|
||||||
|
# Annotated datasets may have language columns absent from metadata.
|
||||||
|
# Extend the schema before the strict Parquet cast.
|
||||||
|
features = self._extend_features_with_language_columns(features)
|
||||||
hf_dataset = load_nested_dataset(self.root / "data", features=features, episodes=self.episodes)
|
hf_dataset = load_nested_dataset(self.root / "data", features=features, episodes=self.episodes)
|
||||||
hf_dataset.set_transform(hf_transform_to_torch)
|
hf_dataset.set_transform(hf_transform_to_torch)
|
||||||
return hf_dataset
|
return hf_dataset
|
||||||
|
|
||||||
|
def _extend_features_with_language_columns(self, features: datasets.Features) -> datasets.Features:
|
||||||
|
"""Register language columns found in Parquet but missing from metadata."""
|
||||||
|
# Leave empty datasets to fail through the normal loading path.
|
||||||
|
try:
|
||||||
|
sample = next((self.root / "data").glob("*/*.parquet"))
|
||||||
|
except StopIteration:
|
||||||
|
return features
|
||||||
|
|
||||||
|
from pyarrow import parquet as _pq # noqa: PLC0415
|
||||||
|
|
||||||
|
schema_names = set(_pq.read_schema(sample).names)
|
||||||
|
from .language import ( # noqa: PLC0415
|
||||||
|
LANGUAGE_EVENTS,
|
||||||
|
LANGUAGE_PERSISTENT,
|
||||||
|
language_events_column_feature,
|
||||||
|
language_persistent_column_feature,
|
||||||
|
)
|
||||||
|
|
||||||
|
extra: dict[str, object] = {}
|
||||||
|
if LANGUAGE_PERSISTENT in schema_names and LANGUAGE_PERSISTENT not in features:
|
||||||
|
extra[LANGUAGE_PERSISTENT] = language_persistent_column_feature()
|
||||||
|
if LANGUAGE_EVENTS in schema_names and LANGUAGE_EVENTS not in features:
|
||||||
|
extra[LANGUAGE_EVENTS] = language_events_column_feature()
|
||||||
|
if not extra:
|
||||||
|
return features
|
||||||
|
return datasets.Features({**features, **extra})
|
||||||
|
|
||||||
def _check_cached_episodes_sufficient(self) -> bool:
|
def _check_cached_episodes_sufficient(self) -> bool:
|
||||||
"""Check if the cached dataset contains all requested episodes and their video files."""
|
"""Check if the cached dataset contains all requested episodes and their video files."""
|
||||||
if self.hf_dataset is None or len(self.hf_dataset) == 0:
|
if self.hf_dataset is None or len(self.hf_dataset) == 0:
|
||||||
|
|||||||
@@ -172,6 +172,23 @@ class DatasetWriter:
|
|||||||
def _get_image_file_dir(self, episode_index: int, image_key: str) -> Path:
|
def _get_image_file_dir(self, episode_index: int, image_key: str) -> Path:
|
||||||
return self._get_image_file_path(episode_index, image_key, frame_index=0).parent
|
return self._get_image_file_path(episode_index, image_key, frame_index=0).parent
|
||||||
|
|
||||||
|
def _get_episode_buffer_index(self) -> int:
|
||||||
|
episode_index = self.episode_buffer["episode_index"]
|
||||||
|
# episode_index is `int` when freshly created, but becomes `np.ndarray` after
|
||||||
|
# save_episode() mutates the buffer. Handle both types here.
|
||||||
|
if isinstance(episode_index, np.ndarray):
|
||||||
|
episode_index = episode_index.item() if episode_index.size == 1 else episode_index[0]
|
||||||
|
return int(episode_index)
|
||||||
|
|
||||||
|
def _delete_camera_frame_dirs(self, camera_keys: list[str]) -> None:
|
||||||
|
if self.image_writer is not None:
|
||||||
|
self._wait_image_writer()
|
||||||
|
episode_index = self._get_episode_buffer_index()
|
||||||
|
for camera_key in camera_keys:
|
||||||
|
img_dir = self._get_image_file_dir(episode_index, camera_key)
|
||||||
|
if img_dir.is_dir():
|
||||||
|
shutil.rmtree(img_dir)
|
||||||
|
|
||||||
def _save_image(
|
def _save_image(
|
||||||
self, image: torch.Tensor | np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1
|
self, image: torch.Tensor | np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -369,7 +386,9 @@ class DatasetWriter:
|
|||||||
self._episodes_since_last_encoding = 0
|
self._episodes_since_last_encoding = 0
|
||||||
|
|
||||||
if episode_data is None:
|
if episode_data is None:
|
||||||
self.clear_episode_buffer(delete_images=len(self._meta.image_keys) > 0)
|
if len(self._meta.image_keys) > 0:
|
||||||
|
self._delete_camera_frame_dirs(self._meta.image_keys)
|
||||||
|
self.episode_buffer = self._create_episode_buffer()
|
||||||
|
|
||||||
def _batch_save_episode_video(self, start_episode: int, end_episode: int | None = None) -> None:
|
def _batch_save_episode_video(self, start_episode: int, end_episode: int | None = None) -> None:
|
||||||
"""Batch save videos for multiple episodes."""
|
"""Batch save videos for multiple episodes."""
|
||||||
@@ -561,10 +580,10 @@ class DatasetWriter:
|
|||||||
return metadata
|
return metadata
|
||||||
|
|
||||||
def clear_episode_buffer(self, delete_images: bool = True) -> None:
|
def clear_episode_buffer(self, delete_images: bool = True) -> None:
|
||||||
"""Discard the current episode buffer and optionally delete temp images.
|
"""Discard the current episode buffer and optionally delete temp camera frames.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
delete_images: If ``True``, remove temporary image directories
|
delete_images: If ``True``, remove temporary camera frame directories
|
||||||
written for the current episode.
|
written for the current episode.
|
||||||
"""
|
"""
|
||||||
# Cancel streaming encoder if active
|
# Cancel streaming encoder if active
|
||||||
@@ -572,17 +591,7 @@ class DatasetWriter:
|
|||||||
self._streaming_encoder.cancel_episode()
|
self._streaming_encoder.cancel_episode()
|
||||||
|
|
||||||
if delete_images:
|
if delete_images:
|
||||||
if self.image_writer is not None:
|
self._delete_camera_frame_dirs(self._meta.camera_keys)
|
||||||
self._wait_image_writer()
|
|
||||||
episode_index = self.episode_buffer["episode_index"]
|
|
||||||
# episode_index is `int` when freshly created, but becomes `np.ndarray` after
|
|
||||||
# save_episode() mutates the buffer. Handle both types here.
|
|
||||||
if isinstance(episode_index, np.ndarray):
|
|
||||||
episode_index = episode_index.item() if episode_index.size == 1 else episode_index[0]
|
|
||||||
for cam_key in self._meta.image_keys:
|
|
||||||
img_dir = self._get_image_file_dir(episode_index, cam_key)
|
|
||||||
if img_dir.is_dir():
|
|
||||||
shutil.rmtree(img_dir)
|
|
||||||
|
|
||||||
self.episode_buffer = self._create_episode_buffer()
|
self.episode_buffer = self._create_episode_buffer()
|
||||||
|
|
||||||
|
|||||||
@@ -66,6 +66,17 @@ def resolve_delta_timestamps(
|
|||||||
return delta_timestamps
|
return delta_timestamps
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_episodes(
|
||||||
|
episodes: list[int] | None, exclude_episodes: list[int] | None, total_episodes: int
|
||||||
|
) -> list[int] | None:
|
||||||
|
"""Apply an episode exclusion list on top of an optional allowlist."""
|
||||||
|
if not exclude_episodes:
|
||||||
|
return episodes
|
||||||
|
base = episodes if episodes is not None else list(range(total_episodes))
|
||||||
|
excluded = set(exclude_episodes)
|
||||||
|
return [episode for episode in base if episode not in excluded]
|
||||||
|
|
||||||
|
|
||||||
def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset:
|
def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset:
|
||||||
"""Handles the logic of setting up delta timestamps and image transforms before creating a dataset.
|
"""Handles the logic of setting up delta timestamps and image transforms before creating a dataset.
|
||||||
|
|
||||||
@@ -87,11 +98,14 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
|||||||
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
||||||
)
|
)
|
||||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
|
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
|
||||||
|
episodes = _resolve_episodes(
|
||||||
|
cfg.dataset.episodes, cfg.dataset.exclude_episodes, ds_meta.total_episodes
|
||||||
|
)
|
||||||
if not cfg.dataset.streaming:
|
if not cfg.dataset.streaming:
|
||||||
dataset = LeRobotDataset(
|
dataset = LeRobotDataset(
|
||||||
cfg.dataset.repo_id,
|
cfg.dataset.repo_id,
|
||||||
root=cfg.dataset.root,
|
root=cfg.dataset.root,
|
||||||
episodes=cfg.dataset.episodes,
|
episodes=episodes,
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
image_transforms=image_transforms,
|
image_transforms=image_transforms,
|
||||||
revision=cfg.dataset.revision,
|
revision=cfg.dataset.revision,
|
||||||
@@ -104,7 +118,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
|||||||
dataset = StreamingLeRobotDataset(
|
dataset = StreamingLeRobotDataset(
|
||||||
cfg.dataset.repo_id,
|
cfg.dataset.repo_id,
|
||||||
root=cfg.dataset.root,
|
root=cfg.dataset.root,
|
||||||
episodes=cfg.dataset.episodes,
|
episodes=episodes,
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
image_transforms=image_transforms,
|
image_transforms=image_transforms,
|
||||||
revision=cfg.dataset.revision,
|
revision=cfg.dataset.revision,
|
||||||
|
|||||||
@@ -162,14 +162,28 @@ def render_sample(
|
|||||||
task: str | None = None,
|
task: str | None = None,
|
||||||
dataset_ctx: Any | None = None,
|
dataset_ctx: Any | None = None,
|
||||||
) -> RenderedMessages | None:
|
) -> RenderedMessages | None:
|
||||||
"""Render the chat-style messages for a single dataset sample.
|
"""Resolve one sample's bindings and render its message recipe.
|
||||||
|
|
||||||
Resolves the recipe's bindings against ``persistent`` and ``events`` rows
|
Returns ``None`` when no text or low-level action supervision applies.
|
||||||
at frame timestamp ``t``, then expands the recipe's message templates.
|
|
||||||
Returns ``None`` if the resolved sample contains no target message.
|
|
||||||
"""
|
"""
|
||||||
persistent_rows = _normalize_rows(persistent or [])
|
persistent_rows = _normalize_rows(persistent or [])
|
||||||
event_rows = _normalize_rows(events or [])
|
event_rows = _normalize_rows(events or [])
|
||||||
|
|
||||||
|
# Route sparse VQA frames to a matching view-specific component before weighted selection.
|
||||||
|
# This avoids dropping annotated frames or selecting VQA without annotations.
|
||||||
|
if recipe.blend is not None:
|
||||||
|
vqa_rendered = _render_vqa_if_present(
|
||||||
|
recipe,
|
||||||
|
persistent=persistent_rows,
|
||||||
|
events=event_rows,
|
||||||
|
t=t,
|
||||||
|
sample_idx=sample_idx,
|
||||||
|
task=task,
|
||||||
|
dataset_ctx=dataset_ctx,
|
||||||
|
)
|
||||||
|
if vqa_rendered is not None:
|
||||||
|
return vqa_rendered
|
||||||
|
|
||||||
selected_recipe = _select_recipe(recipe, sample_idx)
|
selected_recipe = _select_recipe(recipe, sample_idx)
|
||||||
bindings = _resolve_bindings(
|
bindings = _resolve_bindings(
|
||||||
selected_recipe,
|
selected_recipe,
|
||||||
@@ -183,6 +197,55 @@ def render_sample(
|
|||||||
return _render_message_recipe(selected_recipe, bindings)
|
return _render_message_recipe(selected_recipe, bindings)
|
||||||
|
|
||||||
|
|
||||||
|
def _render_vqa_if_present(
|
||||||
|
recipe: TrainingRecipe,
|
||||||
|
*,
|
||||||
|
persistent: Sequence[LanguageRow],
|
||||||
|
events: Sequence[LanguageRow],
|
||||||
|
t: float,
|
||||||
|
sample_idx: int,
|
||||||
|
task: str | None,
|
||||||
|
dataset_ctx: Any | None,
|
||||||
|
) -> RenderedMessages | None:
|
||||||
|
"""Render a matching VQA component, or return ``None`` for normal selection.
|
||||||
|
|
||||||
|
Multiple matching views are selected deterministically by relative weight.
|
||||||
|
"""
|
||||||
|
assert recipe.blend is not None
|
||||||
|
renderable: list[tuple[float, RenderedMessages]] = []
|
||||||
|
for name, component in recipe.blend.items():
|
||||||
|
if not name.startswith("ask_vqa"):
|
||||||
|
continue
|
||||||
|
bindings = _resolve_bindings(
|
||||||
|
component,
|
||||||
|
persistent=persistent,
|
||||||
|
events=events,
|
||||||
|
t=t,
|
||||||
|
sample_idx=sample_idx,
|
||||||
|
task=task,
|
||||||
|
dataset_ctx=dataset_ctx,
|
||||||
|
)
|
||||||
|
rendered = _render_message_recipe(component, bindings)
|
||||||
|
if rendered is not None:
|
||||||
|
renderable.append((float(component.weight or 0.0), rendered))
|
||||||
|
|
||||||
|
if not renderable:
|
||||||
|
return None
|
||||||
|
if len(renderable) == 1:
|
||||||
|
return renderable[0][1]
|
||||||
|
|
||||||
|
# Choose among matching cameras by relative weight, or uniformly when all weights are zero.
|
||||||
|
total = sum(w for w, _ in renderable) or float(len(renderable))
|
||||||
|
digest = hashlib.blake2b(f"vqa:{sample_idx}".encode(), digest_size=8).digest()
|
||||||
|
draw = int.from_bytes(digest, "big") / 2**64 * total
|
||||||
|
cumulative = 0.0
|
||||||
|
for w, rendered in renderable:
|
||||||
|
cumulative += w or (total / len(renderable))
|
||||||
|
if draw < cumulative:
|
||||||
|
return rendered
|
||||||
|
return renderable[-1][1]
|
||||||
|
|
||||||
|
|
||||||
def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
||||||
"""Pick a deterministic blend component for ``sample_idx`` (or return ``recipe``)."""
|
"""Pick a deterministic blend component for ``sample_idx`` (or return ``recipe``)."""
|
||||||
if recipe.blend is None:
|
if recipe.blend is None:
|
||||||
@@ -346,7 +409,9 @@ def _render_message_recipe(
|
|||||||
if turn.target:
|
if turn.target:
|
||||||
target_indices.append(message_idx)
|
target_indices.append(message_idx)
|
||||||
|
|
||||||
if not target_indices:
|
# Keep samples with either text targets or low-level action supervision.
|
||||||
|
has_low_level = any(stream == "low_level" for stream in streams)
|
||||||
|
if not target_indices and not has_low_level:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
rendered = {
|
rendered = {
|
||||||
@@ -403,14 +468,12 @@ def _validate_rendered(rendered: RenderedMessages) -> None:
|
|||||||
|
|
||||||
if len(streams) != len(messages):
|
if len(streams) != len(messages):
|
||||||
raise ValueError("message_streams must be aligned with messages.")
|
raise ValueError("message_streams must be aligned with messages.")
|
||||||
if not target_indices:
|
# Require text or low-level action supervision.
|
||||||
raise ValueError("Rendered samples must contain at least one target message.")
|
if not target_indices and not any(s == "low_level" for s in streams):
|
||||||
|
raise ValueError("Rendered samples must contain a target message or a low_level-stream message.")
|
||||||
for idx in target_indices:
|
for idx in target_indices:
|
||||||
if idx < 0 or idx >= len(messages):
|
if idx < 0 or idx >= len(messages):
|
||||||
raise ValueError(f"Target message index {idx} is out of bounds.")
|
raise ValueError(f"Target message index {idx} is out of bounds.")
|
||||||
# ``stream`` is enforced non-None at MessageTurn construction time
|
|
||||||
# (see ``MessageTurn.__post_init__``), so a missing stream here would
|
|
||||||
# mean the dataclass invariant was bypassed; no need to re-check.
|
|
||||||
|
|
||||||
|
|
||||||
def _nth_relative(
|
def _nth_relative(
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -478,18 +492,19 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
"""Return the number of frames in the selected episodes."""
|
"""Return the number of frames in the selected episodes."""
|
||||||
return self.num_frames
|
return self.num_frames
|
||||||
|
|
||||||
def __getitem__(self, idx) -> dict:
|
def __getitem__(self, idx: int | slice) -> dict | list[dict]:
|
||||||
"""Return a single frame by index, with all transforms applied.
|
"""Return one frame or a slice of frames, with all transforms applied.
|
||||||
|
|
||||||
Loads the frame from the underlying HF dataset, expands delta-timestamp
|
Loads the frame from the underlying HF dataset, expands delta-timestamp
|
||||||
windows, decodes video frames, and applies image transforms. Delegates
|
windows, decodes video frames, and applies image transforms. Delegates
|
||||||
the core logic to :meth:`DatasetReader.get_item`.
|
the core logic to :class:`DatasetReader`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
idx: Index into the (possibly episode-filtered) dataset.
|
idx: Integer index or slice into the possibly episode-filtered dataset.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict mapping feature names to their tensor values for this frame.
|
A frame dictionary for an integer index, or a list of frame
|
||||||
|
dictionaries for a slice.
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
RuntimeError: If the dataset is currently being recorded and
|
RuntimeError: If the dataset is currently being recorded and
|
||||||
@@ -499,6 +514,9 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Cannot read from a dataset that is being recorded. Call finalize() first, then access items."
|
"Cannot read from a dataset that is being recorded. Call finalize() first, then access items."
|
||||||
)
|
)
|
||||||
|
if isinstance(idx, slice):
|
||||||
|
return [self[item_idx] for item_idx in range(*idx.indices(len(self)))]
|
||||||
|
|
||||||
reader = self._ensure_reader()
|
reader = self._ensure_reader()
|
||||||
if reader.hf_dataset is None:
|
if reader.hf_dataset is None:
|
||||||
# One-shot load after finalize()
|
# One-shot load after finalize()
|
||||||
@@ -622,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()
|
||||||
@@ -639,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:
|
||||||
@@ -650,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
|
||||||
|
|
||||||
@@ -789,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.
|
||||||
|
|
||||||
@@ -822,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.
|
||||||
@@ -850,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
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -256,6 +256,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
|
|||||||
shuffle: bool = True,
|
shuffle: bool = True,
|
||||||
return_uint8: bool = False,
|
return_uint8: bool = False,
|
||||||
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
|
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
|
||||||
|
*,
|
||||||
|
token: str | bool | None = None,
|
||||||
):
|
):
|
||||||
"""Initialize a StreamingLeRobotDataset.
|
"""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.
|
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".
|
||||||
|
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
|
||||||
@@ -306,7 +313,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
|
||||||
@@ -334,12 +345,14 @@ 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}
|
||||||
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",
|
||||||
streaming=self.streaming,
|
streaming=self.streaming,
|
||||||
data_files="data/*/*.parquet",
|
data_files="data/*/*.parquet",
|
||||||
revision=self.revision,
|
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)
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -322,7 +322,7 @@ class HILSerlRobotEnvConfig(EnvConfig):
|
|||||||
class LiberoEnv(EnvConfig):
|
class LiberoEnv(EnvConfig):
|
||||||
task: str = "libero_10" # can also choose libero_spatial, libero_object, etc.
|
task: str = "libero_10" # can also choose libero_spatial, libero_object, etc.
|
||||||
task_ids: list[int] | None = None
|
task_ids: list[int] | None = None
|
||||||
fps: int = 30
|
fps: int = 20 # Must match robosuite's default control_freq (20 Hz)
|
||||||
episode_length: int | None = None
|
episode_length: int | None = None
|
||||||
obs_type: str = "pixels_agent_pos"
|
obs_type: str = "pixels_agent_pos"
|
||||||
render_mode: str = "rgb_array"
|
render_mode: str = "rgb_array"
|
||||||
@@ -354,6 +354,9 @@ class LiberoEnv(EnvConfig):
|
|||||||
control_mode: str = "relative" # or "absolute"
|
control_mode: str = "relative" # or "absolute"
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
if self.fps <= 0:
|
||||||
|
raise ValueError(f"fps must be positive, got {self.fps}")
|
||||||
|
|
||||||
if self.obs_type == "pixels":
|
if self.obs_type == "pixels":
|
||||||
self.features[LIBERO_KEY_PIXELS_AGENTVIEW] = PolicyFeature(
|
self.features[LIBERO_KEY_PIXELS_AGENTVIEW] = PolicyFeature(
|
||||||
type=FeatureType.VISUAL, shape=(self.observation_height, self.observation_width, 3)
|
type=FeatureType.VISUAL, shape=(self.observation_height, self.observation_width, 3)
|
||||||
@@ -412,6 +415,7 @@ class LiberoEnv(EnvConfig):
|
|||||||
"render_mode": self.render_mode,
|
"render_mode": self.render_mode,
|
||||||
"observation_height": self.observation_height,
|
"observation_height": self.observation_height,
|
||||||
"observation_width": self.observation_width,
|
"observation_width": self.observation_width,
|
||||||
|
"control_freq": self.fps,
|
||||||
}
|
}
|
||||||
if self.task_ids is not None:
|
if self.task_ids is not None:
|
||||||
kwargs["task_ids"] = self.task_ids
|
kwargs["task_ids"] = self.task_ids
|
||||||
|
|||||||
@@ -125,10 +125,13 @@ class LiberoEnv(gym.Env):
|
|||||||
n_envs: int = 1,
|
n_envs: int = 1,
|
||||||
camera_name_mapping: dict[str, str] | None = None,
|
camera_name_mapping: dict[str, str] | None = None,
|
||||||
num_steps_wait: int = 10,
|
num_steps_wait: int = 10,
|
||||||
|
control_freq: int = 20,
|
||||||
control_mode: str = "relative",
|
control_mode: str = "relative",
|
||||||
is_libero_plus: bool = False,
|
is_libero_plus: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
if control_freq <= 0:
|
||||||
|
raise ValueError(f"control_freq must be positive, got {control_freq}")
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.is_libero_plus = is_libero_plus
|
self.is_libero_plus = is_libero_plus
|
||||||
self.obs_type = obs_type
|
self.obs_type = obs_type
|
||||||
@@ -154,6 +157,7 @@ class LiberoEnv(gym.Env):
|
|||||||
}
|
}
|
||||||
self.camera_name_mapping = camera_name_mapping
|
self.camera_name_mapping = camera_name_mapping
|
||||||
self.num_steps_wait = num_steps_wait
|
self.num_steps_wait = num_steps_wait
|
||||||
|
self.control_freq = control_freq
|
||||||
self.episode_index = episode_index
|
self.episode_index = episode_index
|
||||||
self.episode_length = episode_length
|
self.episode_length = episode_length
|
||||||
# Load once and keep
|
# Load once and keep
|
||||||
@@ -260,6 +264,7 @@ class LiberoEnv(gym.Env):
|
|||||||
bddl_file_name=self._task_bddl_file,
|
bddl_file_name=self._task_bddl_file,
|
||||||
camera_heights=self.observation_height,
|
camera_heights=self.observation_height,
|
||||||
camera_widths=self.observation_width,
|
camera_widths=self.observation_width,
|
||||||
|
control_freq=self.control_freq,
|
||||||
)
|
)
|
||||||
env.reset()
|
env.reset()
|
||||||
self._env = env
|
self._env = env
|
||||||
|
|||||||
@@ -155,6 +155,7 @@ class MetaworldEnv(gym.Env):
|
|||||||
env.model.cam_pos[2] = [0.75, 0.075, 0.7]
|
env.model.cam_pos[2] = [0.75, 0.075, 0.7]
|
||||||
env.reset()
|
env.reset()
|
||||||
env._freeze_rand_vec = False # otherwise no randomization
|
env._freeze_rand_vec = False # otherwise no randomization
|
||||||
|
env.seeded_rand_vec = True # use seeded RNG so reset(seed=X) controls object positions
|
||||||
self._env = env
|
self._env = env
|
||||||
|
|
||||||
def render(self) -> np.ndarray:
|
def render(self) -> np.ndarray:
|
||||||
@@ -220,6 +221,8 @@ class MetaworldEnv(gym.Env):
|
|||||||
self._ensure_env()
|
self._ensure_env()
|
||||||
super().reset(seed=seed)
|
super().reset(seed=seed)
|
||||||
|
|
||||||
|
if seed is not None:
|
||||||
|
self._env.seed(seed)
|
||||||
raw_obs, info = self._env.reset(seed=seed)
|
raw_obs, info = self._env.reset(seed=seed)
|
||||||
|
|
||||||
observation = self._format_raw_obs(raw_obs)
|
observation = self._format_raw_obs(raw_obs)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from lerobot.utils.import_utils import require_package
|
|||||||
# guard the optional dependency here so importing this package fails loudly if it's missing.
|
# guard the optional dependency here so importing this package fails loudly if it's missing.
|
||||||
require_package("datasets", extra="dataset")
|
require_package("datasets", extra="dataset")
|
||||||
|
|
||||||
|
from .annotate import submit_annotate_to_hf
|
||||||
from .hf import submit_to_hf
|
from .hf import submit_to_hf
|
||||||
|
|
||||||
__all__ = ["submit_to_hf"]
|
__all__ = ["submit_annotate_to_hf", "submit_to_hf"]
|
||||||
|
|||||||
@@ -0,0 +1,176 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
"""Run ``lerobot-annotate`` on HF Jobs (HuggingFace GPUs).
|
||||||
|
|
||||||
|
Same shape as the training submitter in ``hf.py``, with one difference: the
|
||||||
|
annotation pipeline serves its own VLM, so the pod starts from the official
|
||||||
|
``vllm/vllm-openai`` image (which has no lerobot) instead of the prebuilt
|
||||||
|
``lerobot-gpu`` image, and installs lerobot on top before running.
|
||||||
|
|
||||||
|
Because there is no config repo to stage, the pod replays the user's own CLI
|
||||||
|
flags — everything except the client-only ``--job.*`` and the host-local
|
||||||
|
``--root``, which is replaced by ``--repo_id`` so the pod pulls the dataset
|
||||||
|
from the Hub.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import shlex
|
||||||
|
import sys
|
||||||
|
from dataclasses import is_dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from huggingface_hub import HfApi, get_token, run_job
|
||||||
|
|
||||||
|
from .dataset import ensure_dataset_available
|
||||||
|
|
||||||
|
# Package-internal reuse of the training submitter's job plumbing: following a
|
||||||
|
# submitted job and forwarding argv are identical for annotation runs.
|
||||||
|
from .hf import _pod_forwarded_args, follow_job, resolve_job_tags
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from lerobot.annotations.steerable_pipeline.config import AnnotationPipelineConfig
|
||||||
|
|
||||||
|
LEROBOT_GIT_URL = "https://github.com/huggingface/lerobot.git"
|
||||||
|
|
||||||
|
# Mirrors the pins in pyproject.toml. The vLLM image resolves dependencies on its
|
||||||
|
# own otherwise, and pulls av 18 / datasets 5 / draccus 0.11 — each of which breaks
|
||||||
|
# lerobot at import time. `--upgrade-strategy only-if-needed` keeps vLLM's own
|
||||||
|
# (torch, transformers, ...) pins intact.
|
||||||
|
_RUNTIME_REQUIREMENTS = (
|
||||||
|
"'datasets>=4.7.0,<5.0.0' 'pyarrow>=21.0.0,<30.0.0' 'av>=15.0.0,<16.0.0' 'draccus==0.10.0' "
|
||||||
|
"'pandas>=2.0.0,<3.0.0' jsonlines gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
||||||
|
"openai"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Flags the submitter resolves itself instead of forwarding verbatim: `--root`
|
||||||
|
# names a directory only this machine has, `--repo_id` is re-emitted from the
|
||||||
|
# config, and the config-file args name local files (rejected up front by
|
||||||
|
# `submit_annotate_to_hf`). `--job.*` is dropped separately, by prefix; bare
|
||||||
|
# `--job` is not, hence its entry here — it is the one arg that could smuggle a
|
||||||
|
# remote `target` onto the pod and have the job recursively submit itself.
|
||||||
|
_SUBMITTER_OWNED_ARGS = ("--root", "--repo_id", "--config_path", "--job")
|
||||||
|
|
||||||
|
|
||||||
|
def _local_config_file_args(cfg: AnnotationPipelineConfig) -> list[str]:
|
||||||
|
"""The CLI args that name a config file on the client's disk.
|
||||||
|
|
||||||
|
draccus exposes ``--config_path`` for the whole config plus a ``--<field>``
|
||||||
|
for every nested dataclass (``--vlm``, ``--plan``, ``--job``, ...). The pod has
|
||||||
|
none of those files, so a remote run has to reject them rather than silently
|
||||||
|
drop the settings they carry.
|
||||||
|
"""
|
||||||
|
return ["--config_path", *(f"--{name}" for name in vars(cfg) if is_dataclass(getattr(cfg, name)))]
|
||||||
|
|
||||||
|
|
||||||
|
def build_pod_setup(lerobot_ref: str) -> str:
|
||||||
|
"""Shell prelude that turns the vLLM image into a ``lerobot-annotate`` runtime."""
|
||||||
|
spec = f"lerobot @ git+{LEROBOT_GIT_URL}@{lerobot_ref}"
|
||||||
|
return (
|
||||||
|
# git to install from the repo, ffmpeg to decode the dataset's videos.
|
||||||
|
"apt-get update -qq && apt-get install -y -qq git ffmpeg && "
|
||||||
|
f"pip install --no-deps {shlex.quote(spec)} && "
|
||||||
|
f"pip install --upgrade-strategy only-if-needed {_RUNTIME_REQUIREMENTS} && "
|
||||||
|
# vLLM's cudagraph memory estimate over-reserves and starves the KV cache;
|
||||||
|
# PyAV is the video backend the server can decode our frames with.
|
||||||
|
"export VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS=0 && "
|
||||||
|
"export VLLM_VIDEO_BACKEND=pyav"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_pod_command(repo_id: str, lerobot_ref: str, argv: list[str]) -> list[str]:
|
||||||
|
"""Build the ``bash -c`` command the pod runs: setup prelude, then annotation.
|
||||||
|
|
||||||
|
``argv`` is the user's CLI (``sys.argv[1:]``) minus the flags in
|
||||||
|
``_SUBMITTER_OWNED_ARGS``; ``--repo_id`` is re-added from the config so the pod
|
||||||
|
always annotates the dataset we just made sure is reachable on the Hub.
|
||||||
|
``--job.target=local`` stops the pod from re-dispatching to itself.
|
||||||
|
"""
|
||||||
|
forwarded = _pod_forwarded_args(argv, drop_names=_SUBMITTER_OWNED_ARGS, drop_prefixes=("--job.",))
|
||||||
|
annotate = shlex.join(["lerobot-annotate", f"--repo_id={repo_id}", *forwarded, "--job.target=local"])
|
||||||
|
return ["bash", "-c", f"{build_pod_setup(lerobot_ref)} && {annotate}"]
|
||||||
|
|
||||||
|
|
||||||
|
def submit_annotate_to_hf(cfg: AnnotationPipelineConfig) -> None:
|
||||||
|
"""Submit an annotation run to HF Jobs infrastructure.
|
||||||
|
|
||||||
|
Resolves credentials, makes sure the source dataset is reachable from the pod,
|
||||||
|
submits the job, then tails its logs until the job reaches a terminal stage —
|
||||||
|
or returns immediately with ``--job.detach``. Ctrl-C detaches without
|
||||||
|
cancelling the remote job.
|
||||||
|
"""
|
||||||
|
token = get_token()
|
||||||
|
if not token:
|
||||||
|
raise RuntimeError("Not logged in to Hugging Face. Run `hf auth login` first.")
|
||||||
|
|
||||||
|
if cfg.repo_id is None:
|
||||||
|
raise ValueError(
|
||||||
|
"Remote annotation requires --repo_id: the pod downloads the dataset from the Hub, "
|
||||||
|
"and --root only names a directory on this machine."
|
||||||
|
)
|
||||||
|
|
||||||
|
argv = sys.argv[1:]
|
||||||
|
passed = {tok.split("=", 1)[0] for tok in argv}
|
||||||
|
used_config_files = sorted(passed.intersection(_local_config_file_args(cfg)))
|
||||||
|
if used_config_files:
|
||||||
|
raise ValueError(
|
||||||
|
f"{', '.join(used_config_files)} cannot be used with a remote --job.target: the pod "
|
||||||
|
"cannot read config files from this machine. Pass the settings as CLI flags instead."
|
||||||
|
)
|
||||||
|
|
||||||
|
if not cfg.push_to_hub:
|
||||||
|
# The pod's filesystem is discarded when the job ends, so without a push the
|
||||||
|
# run produces nothing. Warn rather than fail: a smoke test over
|
||||||
|
# --only_episodes that only inspects the logs is a legitimate use.
|
||||||
|
print(
|
||||||
|
"WARNING: --push_to_hub is off. The annotated dataset lives only on the pod and is "
|
||||||
|
"discarded when the job ends. Pass --push_to_hub=true to keep the result."
|
||||||
|
)
|
||||||
|
|
||||||
|
api = HfApi(token=token)
|
||||||
|
tags = resolve_job_tags(cfg.job.tags)
|
||||||
|
ensure_dataset_available(cfg.repo_id, api=api, tags=tags)
|
||||||
|
|
||||||
|
command = build_pod_command(cfg.repo_id, cfg.job.lerobot_ref, argv)
|
||||||
|
|
||||||
|
print(f"Submitting job to HF Jobs (flavor={cfg.job.target}, image={cfg.job.image}) ...")
|
||||||
|
job_info = run_job(
|
||||||
|
image=cfg.job.image,
|
||||||
|
command=command,
|
||||||
|
flavor=cfg.job.target,
|
||||||
|
secrets={"HF_TOKEN": token},
|
||||||
|
timeout=cfg.job.timeout,
|
||||||
|
# HF Jobs labels are key/value; expose each tag as a queryable label.
|
||||||
|
labels=dict.fromkeys(tags, "true"),
|
||||||
|
)
|
||||||
|
job_id = job_info.id
|
||||||
|
job_url = getattr(job_info, "url", None)
|
||||||
|
print(f"Job submitted: {job_id}")
|
||||||
|
if job_url:
|
||||||
|
print(f" Job page: {job_url}")
|
||||||
|
target_repo_id = cfg.new_repo_id or cfg.repo_id
|
||||||
|
if cfg.push_to_hub:
|
||||||
|
print(f" Dataset repo: https://huggingface.co/datasets/{target_repo_id}")
|
||||||
|
print(f" Monitor: hf jobs logs {job_id}")
|
||||||
|
print(f" Cancel: hf jobs cancel {job_id}")
|
||||||
|
|
||||||
|
# No success marker: `lerobot-annotate` keeps working after the upload log line
|
||||||
|
# (dataset card, version tag), so completion has to be stage-based.
|
||||||
|
if not follow_job(job_id, detach=cfg.job.detach):
|
||||||
|
return
|
||||||
|
|
||||||
|
if cfg.push_to_hub:
|
||||||
|
print(f"\nAnnotation complete — dataset pushed to https://huggingface.co/datasets/{target_repo_id}")
|
||||||
|
else:
|
||||||
|
print("\nAnnotation complete. Note: --push_to_hub was off, so the result stayed on the pod.")
|
||||||
+69
-54
@@ -223,6 +223,74 @@ def _poll_until_done(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def follow_job(job_id: str, *, detach: bool = False, success_marker: str | None = None) -> bool:
|
||||||
|
"""Watch a submitted job to the end, streaming its logs to stdout.
|
||||||
|
|
||||||
|
Returns True when the job finished successfully and False when we stopped watching
|
||||||
|
without a verdict — `detach`, or the user pressing Ctrl-C, which detaches rather than
|
||||||
|
cancelling the remote job. Raises RuntimeError when the job reaches a terminal stage
|
||||||
|
other than COMPLETED.
|
||||||
|
|
||||||
|
`success_marker` finishes as soon as that string appears in the logs instead of waiting
|
||||||
|
out the platform's post-run finalization (~30s). Callers that have a log line meaning
|
||||||
|
"the artifact is on the Hub" should pass it; without one, completion is stage-based.
|
||||||
|
"""
|
||||||
|
if detach:
|
||||||
|
return False
|
||||||
|
|
||||||
|
done = threading.Event()
|
||||||
|
detached = threading.Event()
|
||||||
|
marker_seen = threading.Event()
|
||||||
|
stage_holder: dict[str, str | None] = {}
|
||||||
|
|
||||||
|
def _poll() -> None:
|
||||||
|
stage_holder["stage"] = _poll_until_done(job_id, done, status_holder=stage_holder)
|
||||||
|
|
||||||
|
poll_thread = threading.Thread(target=_poll, daemon=True)
|
||||||
|
poll_thread.start()
|
||||||
|
log_thread = threading.Thread(
|
||||||
|
target=_tail_logs, args=(job_id, done, success_marker, marker_seen), daemon=True
|
||||||
|
)
|
||||||
|
log_thread.start()
|
||||||
|
|
||||||
|
def _detach(sig, frame):
|
||||||
|
detached.set()
|
||||||
|
done.set()
|
||||||
|
print("\nDetached. Job is still running.")
|
||||||
|
print(f" Monitor: hf jobs logs {job_id}")
|
||||||
|
print(f" Cancel: hf jobs cancel {job_id}")
|
||||||
|
|
||||||
|
# signal.signal only works on the main thread; when called from a worker thread
|
||||||
|
# (e.g. an orchestration framework) skip the Ctrl-C-detaches-instead-of-cancels
|
||||||
|
# handler rather than crashing with ValueError.
|
||||||
|
install_sigint = threading.current_thread() is threading.main_thread()
|
||||||
|
original_sigint = signal.getsignal(signal.SIGINT) if install_sigint else None
|
||||||
|
if install_sigint:
|
||||||
|
signal.signal(signal.SIGINT, _detach)
|
||||||
|
try:
|
||||||
|
# Timeout-based join so SIGINT is delivered to the main thread promptly.
|
||||||
|
while poll_thread.is_alive():
|
||||||
|
poll_thread.join(timeout=0.5)
|
||||||
|
log_thread.join(timeout=5)
|
||||||
|
finally:
|
||||||
|
if install_sigint:
|
||||||
|
signal.signal(signal.SIGINT, original_sigint)
|
||||||
|
|
||||||
|
if detached.is_set():
|
||||||
|
return False
|
||||||
|
if marker_seen.is_set():
|
||||||
|
return True
|
||||||
|
|
||||||
|
stage = stage_holder.get("stage")
|
||||||
|
if stage != "COMPLETED":
|
||||||
|
message = stage_holder.get("message")
|
||||||
|
detail = f" ({message})" if message else ""
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Job {job_id} ended with stage={stage}{detail}. Check logs: hf jobs logs {job_id}"
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _pod_forwarded_args(
|
def _pod_forwarded_args(
|
||||||
argv: list[str], drop_names: tuple[str, ...] = (), drop_prefixes: tuple[str, ...] = ()
|
argv: list[str], drop_names: tuple[str, ...] = (), drop_prefixes: tuple[str, ...] = ()
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
@@ -362,64 +430,11 @@ def submit_to_hf(cfg: TrainPipelineConfig) -> None:
|
|||||||
print(f" Monitor: hf jobs logs {job_id}")
|
print(f" Monitor: hf jobs logs {job_id}")
|
||||||
print(f" Cancel: hf jobs cancel {job_id}")
|
print(f" Cancel: hf jobs cancel {job_id}")
|
||||||
|
|
||||||
if cfg.job.detach:
|
|
||||||
return
|
|
||||||
|
|
||||||
done = threading.Event()
|
|
||||||
detached = threading.Event()
|
|
||||||
pushed_ok = threading.Event()
|
|
||||||
stage_holder: dict[str, str | None] = {}
|
|
||||||
|
|
||||||
def _poll() -> None:
|
|
||||||
stage_holder["stage"] = _poll_until_done(job_id, done, status_holder=stage_holder)
|
|
||||||
|
|
||||||
poll_thread = threading.Thread(target=_poll, daemon=True)
|
|
||||||
poll_thread.start()
|
|
||||||
# Finish as soon as the model is pushed, rather than waiting out the platform's
|
# Finish as soon as the model is pushed, rather than waiting out the platform's
|
||||||
# post-run finalization before the job stage flips to COMPLETED. This matches the
|
# post-run finalization before the job stage flips to COMPLETED. This matches the
|
||||||
# exact log line emitted by PreTrainedPolicy.push_model_to_hub — the two must stay
|
# exact log line emitted by PreTrainedPolicy.push_model_to_hub — the two must stay
|
||||||
# in sync. If it ever stops matching we just fall back to stage-based completion
|
# in sync. If it ever stops matching we just fall back to stage-based completion
|
||||||
# (~30s slower), so the contract is an optimization, not a correctness requirement.
|
# (~30s slower), so the contract is an optimization, not a correctness requirement.
|
||||||
success_marker = f"Model pushed to https://huggingface.co/{repo_id}"
|
success_marker = f"Model pushed to https://huggingface.co/{repo_id}"
|
||||||
log_thread = threading.Thread(
|
if follow_job(job_id, detach=cfg.job.detach, success_marker=success_marker):
|
||||||
target=_tail_logs, args=(job_id, done, success_marker, pushed_ok), daemon=True
|
|
||||||
)
|
|
||||||
log_thread.start()
|
|
||||||
|
|
||||||
def _detach(sig, frame):
|
|
||||||
detached.set()
|
|
||||||
done.set()
|
|
||||||
print("\nDetached. Job is still running.")
|
|
||||||
print(f" Monitor: hf jobs logs {job_id}")
|
|
||||||
print(f" Cancel: hf jobs cancel {job_id}")
|
|
||||||
|
|
||||||
# signal.signal only works on the main thread; when called from a worker thread
|
|
||||||
# (e.g. an orchestration framework) skip the Ctrl-C-detaches-instead-of-cancels
|
|
||||||
# handler rather than crashing with ValueError.
|
|
||||||
install_sigint = threading.current_thread() is threading.main_thread()
|
|
||||||
original_sigint = signal.getsignal(signal.SIGINT) if install_sigint else None
|
|
||||||
if install_sigint:
|
|
||||||
signal.signal(signal.SIGINT, _detach)
|
|
||||||
try:
|
|
||||||
# Timeout-based join so SIGINT is delivered to the main thread promptly.
|
|
||||||
while poll_thread.is_alive():
|
|
||||||
poll_thread.join(timeout=0.5)
|
|
||||||
log_thread.join(timeout=5)
|
|
||||||
finally:
|
|
||||||
if install_sigint:
|
|
||||||
signal.signal(signal.SIGINT, original_sigint)
|
|
||||||
|
|
||||||
if detached.is_set():
|
|
||||||
return
|
|
||||||
|
|
||||||
if pushed_ok.is_set():
|
|
||||||
print(f"\nTraining complete — model pushed to https://huggingface.co/{repo_id}")
|
print(f"\nTraining complete — model pushed to https://huggingface.co/{repo_id}")
|
||||||
return
|
|
||||||
|
|
||||||
stage = stage_holder.get("stage")
|
|
||||||
if stage != "COMPLETED":
|
|
||||||
message = stage_holder.get("message")
|
|
||||||
detail = f" ({message})" if message else ""
|
|
||||||
raise RuntimeError(
|
|
||||||
f"Job {job_id} ended with stage={stage}{detail}. Check logs: hf jobs logs {job_id}"
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import logging
|
|||||||
import time
|
import time
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from functools import cached_property
|
|
||||||
from typing import TYPE_CHECKING, Any, TypedDict
|
from typing import TYPE_CHECKING, Any, TypedDict
|
||||||
|
|
||||||
from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected
|
from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected
|
||||||
@@ -854,7 +853,7 @@ class DamiaoMotorsBus(MotorsBusBase):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Motor {motor_obj} doesn't have a valid recv_id (None).")
|
raise ValueError(f"Motor {motor_obj} doesn't have a valid recv_id (None).")
|
||||||
|
|
||||||
@cached_property
|
@property
|
||||||
def is_calibrated(self) -> bool:
|
def is_calibrated(self) -> bool:
|
||||||
"""Check if motors are calibrated."""
|
"""Check if motors are calibrated."""
|
||||||
return bool(self.calibration)
|
return bool(self.calibration)
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import abc
|
import abc
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
@@ -818,13 +819,13 @@ class SerialMotorsBus(MotorsBusBase):
|
|||||||
"""
|
"""
|
||||||
motor_names = self._get_motors_list(motors)
|
motor_names = self._get_motors_list(motors)
|
||||||
|
|
||||||
start_positions = self.sync_read("Present_Position", motor_names, normalize=False)
|
start_positions = self.sync_read("Present_Position", motor_names, normalize=False, num_retry=5)
|
||||||
mins = start_positions.copy()
|
mins = start_positions.copy()
|
||||||
maxes = start_positions.copy()
|
maxes = start_positions.copy()
|
||||||
|
|
||||||
user_pressed_enter = False
|
user_pressed_enter = False
|
||||||
while not user_pressed_enter:
|
while not user_pressed_enter:
|
||||||
positions = self.sync_read("Present_Position", motor_names, normalize=False)
|
positions = self.sync_read("Present_Position", motor_names, normalize=False, num_retry=5)
|
||||||
mins = {motor: min(positions[motor], min_) for motor, min_ in mins.items()}
|
mins = {motor: min(positions[motor], min_) for motor, min_ in mins.items()}
|
||||||
maxes = {motor: max(positions[motor], max_) for motor, max_ in maxes.items()}
|
maxes = {motor: max(positions[motor], max_) for motor, max_ in maxes.items()}
|
||||||
|
|
||||||
@@ -837,9 +838,12 @@ class SerialMotorsBus(MotorsBusBase):
|
|||||||
if enter_pressed():
|
if enter_pressed():
|
||||||
user_pressed_enter = True
|
user_pressed_enter = True
|
||||||
|
|
||||||
if display_values and not user_pressed_enter:
|
if not user_pressed_enter:
|
||||||
# Move cursor up to overwrite the previous output
|
if display_values:
|
||||||
move_cursor_up(len(motor_names) + 3)
|
# Move cursor up to overwrite the previous output
|
||||||
|
move_cursor_up(len(motor_names) + 3)
|
||||||
|
# Throttle reads even when the live table is disabled.
|
||||||
|
time.sleep(0.02)
|
||||||
|
|
||||||
same_min_max = [motor for motor in motor_names if mins[motor] == maxes[motor]]
|
same_min_max = [motor for motor in motor_names if mins[motor] == maxes[motor]]
|
||||||
if same_min_max:
|
if same_min_max:
|
||||||
|
|||||||
@@ -79,6 +79,8 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
use_film_scale_modulation: FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning.
|
use_film_scale_modulation: FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning.
|
||||||
Bias modulation is used be default, while this parameter indicates whether to also use scale
|
Bias modulation is used be default, while this parameter indicates whether to also use scale
|
||||||
modulation.
|
modulation.
|
||||||
|
gradient_checkpointing: Whether to checkpoint the Unet residual blocks during training. This reduces
|
||||||
|
activation memory at the cost of recomputing those blocks during the backward pass.
|
||||||
noise_scheduler_type: Name of the noise scheduler to use. Supported options: ["DDPM", "DDIM"].
|
noise_scheduler_type: Name of the noise scheduler to use. Supported options: ["DDPM", "DDIM"].
|
||||||
num_train_timesteps: Number of diffusion steps for the forward diffusion schedule.
|
num_train_timesteps: Number of diffusion steps for the forward diffusion schedule.
|
||||||
beta_schedule: Name of the diffusion beta schedule as per DDPMScheduler from Hugging Face diffusers.
|
beta_schedule: Name of the diffusion beta schedule as per DDPMScheduler from Hugging Face diffusers.
|
||||||
@@ -132,6 +134,7 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
n_groups: int = 8
|
n_groups: int = 8
|
||||||
diffusion_step_embed_dim: int = 128
|
diffusion_step_embed_dim: int = 128
|
||||||
use_film_scale_modulation: bool = True
|
use_film_scale_modulation: bool = True
|
||||||
|
gradient_checkpointing: bool = False
|
||||||
# Noise scheduler.
|
# Noise scheduler.
|
||||||
noise_scheduler_type: str = "DDPM"
|
noise_scheduler_type: str = "DDPM"
|
||||||
num_train_timesteps: int = 100
|
num_train_timesteps: int = 100
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ import torch
|
|||||||
import torch.nn.functional as F # noqa: N812
|
import torch.nn.functional as F # noqa: N812
|
||||||
import torchvision
|
import torchvision
|
||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
|
from torch.utils.checkpoint import checkpoint
|
||||||
|
|
||||||
from lerobot.utils.constants import ACTION, OBS_ENV_STATE, OBS_IMAGES, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_ENV_STATE, OBS_IMAGES, OBS_STATE
|
||||||
from lerobot.utils.import_utils import _diffusers_available, require_package
|
from lerobot.utils.import_utils import _diffusers_available, require_package
|
||||||
@@ -727,22 +728,35 @@ class DiffusionConditionalUnet1d(nn.Module):
|
|||||||
else:
|
else:
|
||||||
global_feature = timesteps_embed
|
global_feature = timesteps_embed
|
||||||
|
|
||||||
|
use_gc = self.config.gradient_checkpointing and self.training
|
||||||
|
|
||||||
# Run encoder, keeping track of skip features to pass to the decoder.
|
# Run encoder, keeping track of skip features to pass to the decoder.
|
||||||
encoder_skip_features: list[Tensor] = []
|
encoder_skip_features: list[Tensor] = []
|
||||||
for resnet, resnet2, downsample in self.down_modules:
|
for resnet, resnet2, downsample in self.down_modules:
|
||||||
x = resnet(x, global_feature)
|
if use_gc:
|
||||||
x = resnet2(x, global_feature)
|
x = checkpoint(resnet, x, global_feature, use_reentrant=False)
|
||||||
|
x = checkpoint(resnet2, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = resnet(x, global_feature)
|
||||||
|
x = resnet2(x, global_feature)
|
||||||
encoder_skip_features.append(x)
|
encoder_skip_features.append(x)
|
||||||
x = downsample(x)
|
x = downsample(x)
|
||||||
|
|
||||||
for mid_module in self.mid_modules:
|
for mid_module in self.mid_modules:
|
||||||
x = mid_module(x, global_feature)
|
if use_gc:
|
||||||
|
x = checkpoint(mid_module, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = mid_module(x, global_feature)
|
||||||
|
|
||||||
# Run decoder, using the skip features from the encoder.
|
# Run decoder, using the skip features from the encoder.
|
||||||
for resnet, resnet2, upsample in self.up_modules:
|
for resnet, resnet2, upsample in self.up_modules:
|
||||||
x = torch.cat((x, encoder_skip_features.pop()), dim=1)
|
x = torch.cat((x, encoder_skip_features.pop()), dim=1)
|
||||||
x = resnet(x, global_feature)
|
if use_gc:
|
||||||
x = resnet2(x, global_feature)
|
x = checkpoint(resnet, x, global_feature, use_reentrant=False)
|
||||||
|
x = checkpoint(resnet2, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = resnet(x, global_feature)
|
||||||
|
x = resnet2(x, global_feature)
|
||||||
x = upsample(x)
|
x = upsample(x)
|
||||||
|
|
||||||
x = self.final_conv(x)
|
x = self.final_conv(x)
|
||||||
|
|||||||
@@ -177,6 +177,7 @@ def make_pre_post_processors(
|
|||||||
return make_groot_pre_post_processors_from_pretrained(
|
return make_groot_pre_post_processors_from_pretrained(
|
||||||
config=policy_cfg,
|
config=policy_cfg,
|
||||||
pretrained_path=pretrained_path,
|
pretrained_path=pretrained_path,
|
||||||
|
revision=pretrained_revision,
|
||||||
dataset_stats=kwargs.get("dataset_stats"),
|
dataset_stats=kwargs.get("dataset_stats"),
|
||||||
dataset_meta=kwargs.get("dataset_meta"),
|
dataset_meta=kwargs.get("dataset_meta"),
|
||||||
preprocessor_overrides=kwargs.get("preprocessor_overrides"),
|
preprocessor_overrides=kwargs.get("preprocessor_overrides"),
|
||||||
|
|||||||
@@ -37,13 +37,19 @@ def is_image_feature(key: str) -> bool:
|
|||||||
@dataclass
|
@dataclass
|
||||||
class ConcurrencyConfig:
|
class ConcurrencyConfig:
|
||||||
"""Configuration for the concurrency of the actor and learner.
|
"""Configuration for the concurrency of the actor and learner.
|
||||||
|
|
||||||
Possible values are:
|
Possible values are:
|
||||||
- "threads": Use threads for the actor and learner.
|
- "threads": Use threads for the actor and learner.
|
||||||
- "processes": Use processes for the actor and learner.
|
- "processes": Use processes for the actor and learner.
|
||||||
|
|
||||||
|
``multiprocessing_context`` selects the process-wide start method when
|
||||||
|
processes are used. Set it to ``None`` to preserve Python's default or a
|
||||||
|
method already selected by the embedding application.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
actor: str = "threads"
|
actor: str = "threads"
|
||||||
learner: str = "threads"
|
learner: str = "threads"
|
||||||
|
multiprocessing_context: str | None = "spawn"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -475,6 +475,7 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
config: GrootConfig,
|
config: GrootConfig,
|
||||||
pretrained_path: str,
|
pretrained_path: str,
|
||||||
*,
|
*,
|
||||||
|
revision: str | None = None,
|
||||||
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
||||||
dataset_meta: Any | None = None,
|
dataset_meta: Any | None = None,
|
||||||
preprocessor_overrides: dict[str, Any] | None = None,
|
preprocessor_overrides: dict[str, Any] | None = None,
|
||||||
@@ -511,6 +512,7 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
|
|
||||||
preprocessor, postprocessor = _load_groot_processor_pipelines(
|
preprocessor, postprocessor = _load_groot_processor_pipelines(
|
||||||
pretrained_path,
|
pretrained_path,
|
||||||
|
revision=revision,
|
||||||
preprocessor_overrides=preprocessor_overrides,
|
preprocessor_overrides=preprocessor_overrides,
|
||||||
postprocessor_overrides=postprocessor_overrides,
|
postprocessor_overrides=postprocessor_overrides,
|
||||||
preprocessor_config_filename=preprocessor_config_filename,
|
preprocessor_config_filename=preprocessor_config_filename,
|
||||||
@@ -526,6 +528,7 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
def _load_groot_processor_pipelines(
|
def _load_groot_processor_pipelines(
|
||||||
pretrained_path: str,
|
pretrained_path: str,
|
||||||
*,
|
*,
|
||||||
|
revision: str | None,
|
||||||
preprocessor_overrides: dict[str, Any],
|
preprocessor_overrides: dict[str, Any],
|
||||||
postprocessor_overrides: dict[str, Any],
|
postprocessor_overrides: dict[str, Any],
|
||||||
preprocessor_config_filename: str,
|
preprocessor_config_filename: str,
|
||||||
@@ -540,6 +543,7 @@ def _load_groot_processor_pipelines(
|
|||||||
preprocessor = PolicyProcessorPipeline.from_pretrained(
|
preprocessor = PolicyProcessorPipeline.from_pretrained(
|
||||||
pretrained_model_name_or_path=pretrained_path,
|
pretrained_model_name_or_path=pretrained_path,
|
||||||
config_filename=preprocessor_config_filename,
|
config_filename=preprocessor_config_filename,
|
||||||
|
revision=revision,
|
||||||
overrides=preprocessor_overrides,
|
overrides=preprocessor_overrides,
|
||||||
to_transition=batch_to_transition,
|
to_transition=batch_to_transition,
|
||||||
to_output=transition_to_batch,
|
to_output=transition_to_batch,
|
||||||
@@ -547,6 +551,7 @@ def _load_groot_processor_pipelines(
|
|||||||
postprocessor = PolicyProcessorPipeline.from_pretrained(
|
postprocessor = PolicyProcessorPipeline.from_pretrained(
|
||||||
pretrained_model_name_or_path=pretrained_path,
|
pretrained_model_name_or_path=pretrained_path,
|
||||||
config_filename=postprocessor_config_filename,
|
config_filename=postprocessor_config_filename,
|
||||||
|
revision=revision,
|
||||||
overrides=postprocessor_overrides,
|
overrides=postprocessor_overrides,
|
||||||
to_transition=policy_action_to_transition,
|
to_transition=policy_action_to_transition,
|
||||||
to_output=transition_to_policy_action,
|
to_output=transition_to_policy_action,
|
||||||
|
|||||||
@@ -524,8 +524,6 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
|
|
||||||
def embed_suffix(self, noisy_actions, timestep):
|
def embed_suffix(self, noisy_actions, timestep):
|
||||||
"""Embed noisy_actions, timestep to prepare for Expert Gemma processing."""
|
"""Embed noisy_actions, timestep to prepare for Expert Gemma processing."""
|
||||||
embs = []
|
|
||||||
pad_masks = []
|
|
||||||
att_masks = []
|
att_masks = []
|
||||||
|
|
||||||
# Embed timestep using sine-cosine positional encoding
|
# Embed timestep using sine-cosine positional encoding
|
||||||
@@ -551,23 +549,17 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
return F.silu(x)
|
return F.silu(x)
|
||||||
|
|
||||||
time_emb = self._apply_checkpoint(time_mlp_func, time_emb)
|
time_emb = self._apply_checkpoint(time_mlp_func, time_emb)
|
||||||
action_time_emb = action_emb
|
|
||||||
adarms_cond = time_emb
|
adarms_cond = time_emb
|
||||||
|
|
||||||
embs.append(action_time_emb)
|
bsize, action_time_dim = action_emb.shape[:2]
|
||||||
bsize, action_time_dim = action_time_emb.shape[:2]
|
pad_masks = torch.ones(bsize, action_time_dim, dtype=torch.bool, device=timestep.device)
|
||||||
action_time_mask = torch.ones(bsize, action_time_dim, dtype=torch.bool, device=timestep.device)
|
|
||||||
pad_masks.append(action_time_mask)
|
|
||||||
|
|
||||||
# Set attention masks so that image, language and state inputs do not attend to action tokens
|
# Set attention masks so that image, language and state inputs do not attend to action tokens
|
||||||
att_masks += [1] + ([0] * (self.config.chunk_size - 1))
|
att_masks += [1] + ([0] * (self.config.chunk_size - 1))
|
||||||
|
att_masks = torch.tensor(att_masks, dtype=action_emb.dtype, device=action_emb.device)
|
||||||
embs = torch.cat(embs, dim=1)
|
|
||||||
pad_masks = torch.cat(pad_masks, dim=1)
|
|
||||||
att_masks = torch.tensor(att_masks, dtype=embs.dtype, device=embs.device)
|
|
||||||
att_masks = att_masks[None, :].expand(bsize, len(att_masks))
|
att_masks = att_masks[None, :].expand(bsize, len(att_masks))
|
||||||
|
|
||||||
return embs, pad_masks, att_masks, adarms_cond
|
return action_emb, pad_masks, att_masks, adarms_cond
|
||||||
|
|
||||||
def forward(self, images, img_masks, tokens, masks, actions, noise, time) -> Tensor:
|
def forward(self, images, img_masks, tokens, masks, actions, noise, time) -> Tensor:
|
||||||
"""Do a full training forward pass and compute the loss."""
|
"""Do a full training forward pass and compute the loss."""
|
||||||
|
|||||||
@@ -175,9 +175,6 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
|
|||||||
if isinstance(task_index_value, Tensor) and task_index_value.dim() == 0:
|
if isinstance(task_index_value, Tensor) and task_index_value.dim() == 0:
|
||||||
complementary_data["task_index"] = task_index_value.unsqueeze(0)
|
complementary_data["task_index"] = task_index_value.unsqueeze(0)
|
||||||
|
|
||||||
complementary_data.pop("language_persistent", None)
|
|
||||||
complementary_data.pop("language_events", None)
|
|
||||||
|
|
||||||
if "messages" in complementary_data:
|
if "messages" in complementary_data:
|
||||||
messages = complementary_data["messages"]
|
messages = complementary_data["messages"]
|
||||||
if isinstance(messages, list) and (not messages or isinstance(messages[0], dict)):
|
if isinstance(messages, list) and (not messages or isinstance(messages[0], dict)):
|
||||||
|
|||||||
@@ -132,10 +132,20 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
for axis in ["x", "y", "z", "gripper"]:
|
for axis in ["x", "y", "z"]:
|
||||||
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
|
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
|
||||||
|
features[PipelineFeatureType.ACTION].pop("gripper", None)
|
||||||
|
|
||||||
for feat in ["enabled", "target_x", "target_y", "target_z", "target_wx", "target_wy", "target_wz"]:
|
for feat in [
|
||||||
|
"enabled",
|
||||||
|
"target_x",
|
||||||
|
"target_y",
|
||||||
|
"target_z",
|
||||||
|
"target_wx",
|
||||||
|
"target_wy",
|
||||||
|
"target_wz",
|
||||||
|
"gripper_vel",
|
||||||
|
]:
|
||||||
features[PipelineFeatureType.ACTION][f"{feat}"] = PolicyFeature(
|
features[PipelineFeatureType.ACTION][f"{feat}"] = PolicyFeature(
|
||||||
type=FeatureType.ACTION, shape=(1,)
|
type=FeatureType.ACTION, shape=(1,)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ from pathlib import Path
|
|||||||
from typing import Any, TypedDict, TypeVar, cast
|
from typing import Any, TypedDict, TypeVar, cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from huggingface_hub import hf_hub_download
|
from huggingface_hub import hf_hub_download, snapshot_download
|
||||||
from safetensors.torch import load_file, save_file
|
from safetensors.torch import load_file, save_file
|
||||||
|
|
||||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||||
@@ -205,6 +205,10 @@ class ProcessorStep(ABC):
|
|||||||
"""
|
"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def save_artifacts(self, save_directory: Path) -> dict[str, str]:
|
||||||
|
"""Save non-tensor assets and map constructor arguments to relative paths."""
|
||||||
|
return {}
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
"""Resets the internal state of the processor step, if any."""
|
"""Resets the internal state of the processor step, if any."""
|
||||||
return None
|
return None
|
||||||
@@ -549,6 +553,22 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
pipeline_config = self.get_config()
|
pipeline_config = self.get_config()
|
||||||
pipeline_state_dict = self.state_dict()
|
pipeline_state_dict = self.state_dict()
|
||||||
|
|
||||||
|
for processor_step, step_entry in zip(self.steps, pipeline_config["steps"], strict=True):
|
||||||
|
artifacts = processor_step.save_artifacts(save_directory)
|
||||||
|
if artifacts:
|
||||||
|
for config_key, relative_path in artifacts.items():
|
||||||
|
artifact_path = Path(relative_path)
|
||||||
|
if artifact_path.is_absolute() or ".." in artifact_path.parts:
|
||||||
|
raise ValueError(
|
||||||
|
f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
|
||||||
|
)
|
||||||
|
if not (save_directory / artifact_path).exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Processor step did not save declared artifact '{relative_path}'"
|
||||||
|
)
|
||||||
|
step_entry["config"][config_key] = artifact_path.as_posix()
|
||||||
|
step_entry["artifacts"] = artifacts
|
||||||
|
|
||||||
for state_key, step_state_dict in pipeline_state_dict.items():
|
for state_key, step_state_dict in pipeline_state_dict.items():
|
||||||
state_filename = f"{state_key}.safetensors"
|
state_filename = f"{state_key}.safetensors"
|
||||||
save_file(step_state_dict, save_directory / state_filename)
|
save_file(step_state_dict, save_directory / state_filename)
|
||||||
@@ -713,6 +733,8 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
ProcessorMigrationError: If the model requires migration to processor format.
|
ProcessorMigrationError: If the model requires migration to processor format.
|
||||||
"""
|
"""
|
||||||
model_id = str(pretrained_model_name_or_path)
|
model_id = str(pretrained_model_name_or_path)
|
||||||
|
model_path = Path(model_id)
|
||||||
|
is_local_source = model_path.is_dir() or model_path.is_file()
|
||||||
hub_download_kwargs = {
|
hub_download_kwargs = {
|
||||||
"force_download": force_download,
|
"force_download": force_download,
|
||||||
"resume_download": resume_download,
|
"resume_download": resume_download,
|
||||||
@@ -731,7 +753,13 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
|
|
||||||
# 3. Build steps with overrides
|
# 3. Build steps with overrides
|
||||||
steps, validated_overrides = cls._build_steps_with_overrides(
|
steps, validated_overrides = cls._build_steps_with_overrides(
|
||||||
loaded_config, overrides or {}, model_id, base_path, hub_download_kwargs
|
loaded_config,
|
||||||
|
overrides or {},
|
||||||
|
model_id,
|
||||||
|
base_path,
|
||||||
|
config_filename,
|
||||||
|
hub_download_kwargs,
|
||||||
|
is_local_source,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. Validate that all overrides were used
|
# 4. Validate that all overrides were used
|
||||||
@@ -920,7 +948,9 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
overrides: dict[str, Any],
|
overrides: dict[str, Any],
|
||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
|
config_filename: str,
|
||||||
hub_download_kwargs: dict[str, Any],
|
hub_download_kwargs: dict[str, Any],
|
||||||
|
is_local_source: bool = False,
|
||||||
) -> tuple[list[ProcessorStep], set[str]]:
|
) -> tuple[list[ProcessorStep], set[str]]:
|
||||||
"""Build all processor steps with overrides and state loading.
|
"""Build all processor steps with overrides and state loading.
|
||||||
|
|
||||||
@@ -944,7 +974,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
3. **State Loading** (via _load_step_state):
|
3. **State Loading** (via _load_step_state):
|
||||||
- **If step has "state_file"**: Load tensor state from .safetensors
|
- **If step has "state_file"**: Load tensor state from .safetensors
|
||||||
- **Local first**: Check base_path/state_file.safetensors
|
- **Local first**: Check base_path/state_file.safetensors
|
||||||
- **Hub fallback**: Download state file if not found locally
|
- **Hub fallback**: Download state file if the pipeline was loaded from the Hub
|
||||||
- **Optional**: Only load if step has load_state_dict method
|
- **Optional**: Only load if step has load_state_dict method
|
||||||
|
|
||||||
4. **Override Tracking**:
|
4. **Override Tracking**:
|
||||||
@@ -962,6 +992,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: The model identifier (needed for Hub state file downloads)
|
model_id: The model identifier (needed for Hub state file downloads)
|
||||||
base_path: Local directory path for finding state files
|
base_path: Local directory path for finding state files
|
||||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||||
|
is_local_source: Whether model_id resolved to a local directory or config file.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (instantiated_steps_list, unused_override_keys)
|
Tuple of (instantiated_steps_list, unused_override_keys)
|
||||||
@@ -972,13 +1003,68 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
ImportError: If a step class cannot be imported or found in registry
|
ImportError: If a step class cannot be imported or found in registry
|
||||||
ValueError: If a step cannot be instantiated with its configuration
|
ValueError: If a step cannot be instantiated with its configuration
|
||||||
"""
|
"""
|
||||||
|
loaded_config = deepcopy(loaded_config)
|
||||||
|
cls._resolve_artifact_paths(
|
||||||
|
loaded_config,
|
||||||
|
model_id,
|
||||||
|
base_path,
|
||||||
|
config_filename,
|
||||||
|
hub_download_kwargs,
|
||||||
|
)
|
||||||
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
|
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
|
||||||
|
|
||||||
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
|
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
|
||||||
cls._load_step_state(step_instance, step_entry, model_id, base_path, hub_download_kwargs)
|
cls._load_step_state(
|
||||||
|
step_instance,
|
||||||
|
step_entry,
|
||||||
|
model_id,
|
||||||
|
base_path,
|
||||||
|
config_filename,
|
||||||
|
hub_download_kwargs,
|
||||||
|
is_local_source,
|
||||||
|
)
|
||||||
|
|
||||||
return steps, remaining_override_keys
|
return steps, remaining_override_keys
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _resolve_artifact_paths(
|
||||||
|
cls,
|
||||||
|
loaded_config: dict[str, Any],
|
||||||
|
model_id: str,
|
||||||
|
base_path: Path | None,
|
||||||
|
config_filename: str,
|
||||||
|
hub_download_kwargs: dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
"""Resolve declared relative processor artifacts before step construction."""
|
||||||
|
is_local = Path(model_id).is_dir() or Path(model_id).is_file()
|
||||||
|
|
||||||
|
for step_entry in loaded_config["steps"]:
|
||||||
|
artifacts = step_entry.get("artifacts", {})
|
||||||
|
for config_key, relative_path in artifacts.items():
|
||||||
|
artifact_path = Path(relative_path)
|
||||||
|
if artifact_path.is_absolute() or ".." in artifact_path.parts:
|
||||||
|
raise ValueError(
|
||||||
|
f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
resolved_path = base_path / artifact_path if base_path is not None else artifact_path
|
||||||
|
if not resolved_path.exists() and not is_local:
|
||||||
|
repository_path = Path(config_filename).parent / artifact_path
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=model_id,
|
||||||
|
repo_type="model",
|
||||||
|
allow_patterns=f"{repository_path.as_posix()}/**",
|
||||||
|
**hub_download_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not resolved_path.exists():
|
||||||
|
step_name = step_entry.get("registry_name", step_entry.get("class", "unknown"))
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Missing processor artifact '{relative_path}' for step '{step_name}' "
|
||||||
|
f"next to '{config_filename}'. Checkpoint artifacts are incomplete."
|
||||||
|
)
|
||||||
|
step_entry["config"][config_key] = str(resolved_path)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _build_steps_from_config(
|
def _build_steps_from_config(
|
||||||
cls,
|
cls,
|
||||||
@@ -1138,7 +1224,9 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
step_entry: dict[str, Any],
|
step_entry: dict[str, Any],
|
||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
|
config_filename: str,
|
||||||
hub_download_kwargs: dict[str, Any],
|
hub_download_kwargs: dict[str, Any],
|
||||||
|
is_local_source: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Load state dictionary for a processor step if available.
|
"""Load state dictionary for a processor step if available.
|
||||||
|
|
||||||
@@ -1157,7 +1245,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
- **Use case**: Loading from local saved model directory
|
- **Use case**: Loading from local saved model directory
|
||||||
|
|
||||||
2. **Hub download fallback**: Download state file from repository
|
2. **Hub download fallback**: Download state file from repository
|
||||||
- **When triggered**: Local file not found or base_path is None
|
- **When triggered**: Local file not found and the pipeline source is a Hub repo
|
||||||
- **Process**: Use hf_hub_download with same parameters as config
|
- **Process**: Use hf_hub_download with same parameters as config
|
||||||
- **Example**: Download "normalize_step_0.safetensors" from "user/repo"
|
- **Example**: Download "normalize_step_0.safetensors" from "user/repo"
|
||||||
- **Result**: Downloaded to local cache, path returned
|
- **Result**: Downloaded to local cache, path returned
|
||||||
@@ -1178,6 +1266,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: The model identifier (used for Hub downloads if needed)
|
model_id: The model identifier (used for Hub downloads if needed)
|
||||||
base_path: Local directory path for finding state files (None for Hub-only)
|
base_path: Local directory path for finding state files (None for Hub-only)
|
||||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||||
|
is_local_source: Whether model_id resolved to a local directory or config file.
|
||||||
|
|
||||||
Note:
|
Note:
|
||||||
This method modifies step_instance in-place and returns None.
|
This method modifies step_instance in-place and returns None.
|
||||||
@@ -1191,11 +1280,17 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
# Try local file first
|
# Try local file first
|
||||||
if base_path and (base_path / state_filename).exists():
|
if base_path and (base_path / state_filename).exists():
|
||||||
state_path = str(base_path / state_filename)
|
state_path = str(base_path / state_filename)
|
||||||
|
elif is_local_source:
|
||||||
|
state_path = base_path / state_filename if base_path else Path(state_filename)
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"State file '{state_filename}' was not found for local processor pipeline "
|
||||||
|
f"'{model_id}' at '{state_path}'."
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Download from Hub
|
# Download from Hub
|
||||||
state_path = hf_hub_download(
|
state_path = hf_hub_download(
|
||||||
repo_id=model_id,
|
repo_id=model_id,
|
||||||
filename=state_filename,
|
filename=(Path(config_filename).parent / state_filename).as_posix(),
|
||||||
repo_type="model",
|
repo_type="model",
|
||||||
**hub_download_kwargs,
|
**hub_download_kwargs,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -16,7 +16,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||||
@@ -32,17 +32,18 @@ from .pipeline import ProcessorStep, ProcessorStepRegistry
|
|||||||
@dataclass
|
@dataclass
|
||||||
@ProcessorStepRegistry.register(name="render_messages_processor")
|
@ProcessorStepRegistry.register(name="render_messages_processor")
|
||||||
class RenderMessagesStep(ProcessorStep):
|
class RenderMessagesStep(ProcessorStep):
|
||||||
"""Processor step that turns raw language columns into rendered chat messages.
|
"""Render language columns into recipe-defined messages and supervision metadata."""
|
||||||
|
|
||||||
Reads ``language_persistent`` and ``language_events`` from the transition's
|
|
||||||
complementary data, renders them through ``recipe`` at the sample timestamp,
|
|
||||||
and replaces the raw columns with the resulting ``messages`` /
|
|
||||||
``message_streams`` / ``target_message_indices`` keys.
|
|
||||||
"""
|
|
||||||
|
|
||||||
recipe: TrainingRecipe
|
recipe: TrainingRecipe
|
||||||
dataset_ctx: Any | None = None
|
dataset_ctx: Any | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if isinstance(self.recipe, dict):
|
||||||
|
self.recipe = TrainingRecipe.from_dict(self.recipe)
|
||||||
|
|
||||||
|
def get_config(self) -> dict[str, Any]:
|
||||||
|
return {"recipe": asdict(self.recipe)}
|
||||||
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition | None:
|
def __call__(self, transition: EnvTransition) -> EnvTransition | None:
|
||||||
"""Render messages for a single transition; return ``None`` to drop it."""
|
"""Render messages for a single transition; return ``None`` to drop it."""
|
||||||
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||||
@@ -50,7 +51,17 @@ class RenderMessagesStep(ProcessorStep):
|
|||||||
events = complementary_data.get(LANGUAGE_EVENTS) or []
|
events = complementary_data.get(LANGUAGE_EVENTS) or []
|
||||||
|
|
||||||
if not persistent and not events:
|
if not persistent and not events:
|
||||||
return transition
|
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||||
|
if rendered is None:
|
||||||
|
return transition
|
||||||
|
new_transition = transition.copy()
|
||||||
|
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||||
|
new_complementary_data.update(rendered)
|
||||||
|
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||||
|
return new_transition
|
||||||
|
|
||||||
|
if _is_batched_language(persistent) or _is_batched_language(events):
|
||||||
|
return self._call_batch(transition, complementary_data, persistent, events)
|
||||||
|
|
||||||
timestamp = complementary_data.get("timestamp")
|
timestamp = complementary_data.get("timestamp")
|
||||||
if timestamp is None:
|
if timestamp is None:
|
||||||
@@ -67,18 +78,147 @@ class RenderMessagesStep(ProcessorStep):
|
|||||||
dataset_ctx=self.dataset_ctx,
|
dataset_ctx=self.dataset_ctx,
|
||||||
)
|
)
|
||||||
if rendered is None:
|
if rendered is None:
|
||||||
return None
|
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||||
|
if rendered is None:
|
||||||
|
return None
|
||||||
|
|
||||||
new_transition = transition.copy()
|
new_transition = transition.copy()
|
||||||
new_complementary_data = dict(complementary_data)
|
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||||
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||||
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||||
new_complementary_data.update(rendered)
|
new_complementary_data.update(rendered)
|
||||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||||
return new_transition
|
return new_transition
|
||||||
|
|
||||||
|
def _call_batch(
|
||||||
|
self,
|
||||||
|
transition: EnvTransition,
|
||||||
|
complementary_data: dict[str, Any],
|
||||||
|
persistent_batch: list,
|
||||||
|
events_batch: list,
|
||||||
|
) -> EnvTransition | None:
|
||||||
|
timestamp = complementary_data.get("timestamp")
|
||||||
|
if timestamp is None:
|
||||||
|
raise KeyError("RenderMessagesStep requires sample timestamp in complementary data.")
|
||||||
|
|
||||||
|
batch_size = max(len(persistent_batch), len(events_batch))
|
||||||
|
messages: list[list[dict[str, Any]]] = []
|
||||||
|
message_streams: list[list[str | None]] = []
|
||||||
|
target_message_indices: list[list[int]] = []
|
||||||
|
keep_indices: list[int] = []
|
||||||
|
|
||||||
|
for i in range(batch_size):
|
||||||
|
rendered = render_sample(
|
||||||
|
recipe=self.recipe,
|
||||||
|
persistent=persistent_batch[i] if i < len(persistent_batch) else [],
|
||||||
|
events=events_batch[i] if i < len(events_batch) else [],
|
||||||
|
t=_batch_value(timestamp, i),
|
||||||
|
sample_idx=int(_batch_value(complementary_data.get("index", 0), i)),
|
||||||
|
task=_batch_value(complementary_data.get("task"), i),
|
||||||
|
dataset_ctx=self.dataset_ctx,
|
||||||
|
)
|
||||||
|
if rendered is None:
|
||||||
|
rendered = _fallback_low_level_render(_batch_value(complementary_data.get("task"), i))
|
||||||
|
if rendered is None:
|
||||||
|
continue
|
||||||
|
keep_indices.append(i)
|
||||||
|
messages.append(rendered["messages"])
|
||||||
|
message_streams.append(rendered["message_streams"])
|
||||||
|
target_message_indices.append(rendered["target_message_indices"])
|
||||||
|
|
||||||
|
if not messages:
|
||||||
|
return None
|
||||||
|
|
||||||
|
new_transition = (
|
||||||
|
_select_batch_indices(transition, keep_indices)
|
||||||
|
if len(keep_indices) != batch_size
|
||||||
|
else transition.copy()
|
||||||
|
)
|
||||||
|
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||||
|
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||||
|
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||||
|
new_complementary_data["messages"] = messages
|
||||||
|
new_complementary_data["message_streams"] = message_streams
|
||||||
|
new_complementary_data["target_message_indices"] = target_message_indices
|
||||||
|
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||||
|
return new_transition
|
||||||
|
|
||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
"""Pass features through unchanged; rendering only touches complementary data."""
|
"""Pass features through unchanged; rendering only touches complementary data."""
|
||||||
return features
|
return features
|
||||||
|
|
||||||
|
|
||||||
|
def _scalar(value: Any) -> float | int:
|
||||||
|
"""Unwrap a tensor/array/single-element list into a Python scalar."""
|
||||||
|
if hasattr(value, "item"):
|
||||||
|
return value.item()
|
||||||
|
if isinstance(value, list):
|
||||||
|
if len(value) != 1:
|
||||||
|
raise ValueError(f"Expected a scalar, got list of length {len(value)}: {value!r}")
|
||||||
|
return _scalar(value[0])
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _is_batched_language(value: Any) -> bool:
|
||||||
|
return isinstance(value, list) and bool(value) and isinstance(value[0], list)
|
||||||
|
|
||||||
|
|
||||||
|
def _batch_value(value: Any, index: int) -> Any:
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if isinstance(value, list):
|
||||||
|
return value[index]
|
||||||
|
if hasattr(value, "ndim") and value.ndim > 0:
|
||||||
|
return _scalar(value[index])
|
||||||
|
return _scalar(value)
|
||||||
|
|
||||||
|
|
||||||
|
def _select_batch_indices(transition: EnvTransition, indices: list[int]) -> EnvTransition:
|
||||||
|
selected = transition.copy()
|
||||||
|
for key in (TransitionKey.OBSERVATION, TransitionKey.COMPLEMENTARY_DATA):
|
||||||
|
data = selected.get(key)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
selected[key] = {k: _select_value(v, indices) for k, v in data.items()}
|
||||||
|
action = selected.get(TransitionKey.ACTION)
|
||||||
|
if action is not None:
|
||||||
|
selected[TransitionKey.ACTION] = _select_value(action, indices)
|
||||||
|
return selected
|
||||||
|
|
||||||
|
|
||||||
|
def _select_value(value: Any, indices: list[int]) -> Any:
|
||||||
|
if isinstance(value, list) and len(value) >= len(indices):
|
||||||
|
return [value[i] for i in indices]
|
||||||
|
if hasattr(value, "index_select") and hasattr(value, "new_tensor") and getattr(value, "ndim", 0) > 0:
|
||||||
|
return value.index_select(0, value.new_tensor(indices).long())
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _fallback_low_level_render(task: Any) -> dict[str, Any] | None:
|
||||||
|
"""Keep action-only samples trainable when no recipe branch matches."""
|
||||||
|
if hasattr(task, "item"):
|
||||||
|
task = task.item()
|
||||||
|
if isinstance(task, list):
|
||||||
|
messages = []
|
||||||
|
message_streams = []
|
||||||
|
target_message_indices = []
|
||||||
|
for t in task:
|
||||||
|
rendered = _fallback_low_level_render(t)
|
||||||
|
if rendered is None:
|
||||||
|
return None
|
||||||
|
messages.append(rendered["messages"])
|
||||||
|
message_streams.append(rendered["message_streams"])
|
||||||
|
target_message_indices.append(rendered["target_message_indices"])
|
||||||
|
return {
|
||||||
|
"messages": messages,
|
||||||
|
"message_streams": message_streams,
|
||||||
|
"target_message_indices": target_message_indices,
|
||||||
|
}
|
||||||
|
if not isinstance(task, str) or not task:
|
||||||
|
return None
|
||||||
|
return {
|
||||||
|
"messages": [{"role": "user", "content": task}],
|
||||||
|
"message_streams": ["low_level"],
|
||||||
|
"target_message_indices": [],
|
||||||
|
}
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -32,6 +33,7 @@ import torch
|
|||||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||||
from lerobot.types import EnvTransition, RobotObservation, TransitionKey
|
from lerobot.types import EnvTransition, RobotObservation, TransitionKey
|
||||||
from lerobot.utils.constants import (
|
from lerobot.utils.constants import (
|
||||||
|
ACTION_CODE_TOKEN_MASK,
|
||||||
ACTION_TOKEN_MASK,
|
ACTION_TOKEN_MASK,
|
||||||
ACTION_TOKENS,
|
ACTION_TOKENS,
|
||||||
OBS_LANGUAGE_ATTENTION_MASK,
|
OBS_LANGUAGE_ATTENTION_MASK,
|
||||||
@@ -136,7 +138,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
|
|||||||
# Standardize to a list of strings for the tokenizer
|
# Standardize to a list of strings for the tokenizer
|
||||||
if isinstance(task, str):
|
if isinstance(task, str):
|
||||||
return [task]
|
return [task]
|
||||||
elif isinstance(task, (list, tuple)) and all(isinstance(t, str) for t in task):
|
elif isinstance(task, list | tuple) and all(isinstance(t, str) for t in task):
|
||||||
return list(task)
|
return list(task)
|
||||||
|
|
||||||
return None
|
return None
|
||||||
@@ -349,6 +351,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
max_action_tokens: int = 256
|
max_action_tokens: int = 256
|
||||||
fast_skip_tokens: int = 128
|
fast_skip_tokens: int = 128
|
||||||
paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224"
|
paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224"
|
||||||
|
allow_truncation: bool = True
|
||||||
# Internal tokenizer instance (not part of the config)
|
# Internal tokenizer instance (not part of the config)
|
||||||
action_tokenizer: Any = field(default=None, init=False, repr=False)
|
action_tokenizer: Any = field(default=None, init=False, repr=False)
|
||||||
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
|
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
|
||||||
@@ -412,14 +415,15 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
# During inference, no action is available, skip tokenization
|
# During inference, no action is available, skip tokenization
|
||||||
return new_transition
|
return new_transition
|
||||||
|
|
||||||
# Tokenize and get both tokens and mask
|
# Tokenize and get masks for the full formatted sequence and the discrete action codes.
|
||||||
tokens, mask = self._tokenize_action(action)
|
tokens, mask, code_mask = self._tokenize_action(action)
|
||||||
|
|
||||||
# Store mask in complementary data
|
# Store mask in complementary data
|
||||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||||
if complementary_data is None:
|
if complementary_data is None:
|
||||||
complementary_data = {}
|
complementary_data = {}
|
||||||
complementary_data[ACTION_TOKEN_MASK] = mask
|
complementary_data[ACTION_TOKEN_MASK] = mask
|
||||||
|
complementary_data[ACTION_CODE_TOKEN_MASK] = code_mask
|
||||||
complementary_data[ACTION_TOKENS] = tokens
|
complementary_data[ACTION_TOKENS] = tokens
|
||||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
||||||
return new_transition
|
return new_transition
|
||||||
@@ -430,7 +434,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
"""
|
"""
|
||||||
return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens
|
return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens
|
||||||
|
|
||||||
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
"""
|
"""
|
||||||
Tokenizes the action tensor and creates a mask.
|
Tokenizes the action tensor and creates a mask.
|
||||||
|
|
||||||
@@ -459,6 +463,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
# The fast tokenizer expects action data and returns token IDs
|
# The fast tokenizer expects action data and returns token IDs
|
||||||
tokens_list = []
|
tokens_list = []
|
||||||
masks_list = []
|
masks_list = []
|
||||||
|
code_masks_list = []
|
||||||
|
|
||||||
for i in range(batch_size):
|
for i in range(batch_size):
|
||||||
# Tokenize single action (move to CPU first as tokenizer uses scipy which requires numpy)
|
# Tokenize single action (move to CPU first as tokenizer uses scipy which requires numpy)
|
||||||
@@ -476,65 +481,82 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
if tokens.dim() > 1:
|
if tokens.dim() > 1:
|
||||||
tokens = tokens.flatten()
|
tokens = tokens.flatten()
|
||||||
|
|
||||||
|
action_code_tokens = self._act_tokens_to_paligemma_tokens(tokens)
|
||||||
bos_id = self._paligemma_tokenizer.bos_token_id
|
bos_id = self._paligemma_tokenizer.bos_token_id
|
||||||
# add bos
|
prompt_tokens = torch.tensor(
|
||||||
|
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
||||||
|
device=action.device,
|
||||||
|
)
|
||||||
|
end_tokens = torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device)
|
||||||
|
|
||||||
|
code_start = 1 + len(prompt_tokens)
|
||||||
|
code_end = code_start + len(action_code_tokens)
|
||||||
tokens = torch.cat(
|
tokens = torch.cat(
|
||||||
[
|
[
|
||||||
torch.tensor([bos_id], device=action.device),
|
torch.tensor([bos_id], device=action.device),
|
||||||
torch.tensor(
|
prompt_tokens,
|
||||||
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
action_code_tokens,
|
||||||
device=action.device,
|
end_tokens,
|
||||||
),
|
|
||||||
self._act_tokens_to_paligemma_tokens(tokens),
|
|
||||||
torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device),
|
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
code_mask = torch.zeros(len(tokens), dtype=torch.bool, device=action.device)
|
||||||
|
code_mask[code_start:code_end] = True
|
||||||
|
|
||||||
# Truncate or pad to max_action_tokens
|
# Truncate or pad to max_action_tokens
|
||||||
if len(tokens) > self.max_action_tokens:
|
if len(tokens) > self.max_action_tokens:
|
||||||
|
if not self.allow_truncation:
|
||||||
|
raise ValueError(
|
||||||
|
f"FAST action sequence has {len(tokens)} tokens, exceeding "
|
||||||
|
f"max_action_tokens={self.max_action_tokens}."
|
||||||
|
)
|
||||||
logging.warning(
|
logging.warning(
|
||||||
f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. "
|
f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. "
|
||||||
"Consider increasing the `max_action_tokens` in your model config if this happens frequently."
|
"Consider increasing the `max_action_tokens` in your model config if this happens frequently."
|
||||||
)
|
)
|
||||||
tokens = tokens[: self.max_action_tokens]
|
tokens = tokens[: self.max_action_tokens]
|
||||||
|
code_mask = code_mask[: self.max_action_tokens]
|
||||||
mask = torch.ones(self.max_action_tokens, dtype=torch.bool, device=action.device)
|
mask = torch.ones(self.max_action_tokens, dtype=torch.bool, device=action.device)
|
||||||
else:
|
else:
|
||||||
|
pad_len = self.max_action_tokens - len(tokens)
|
||||||
mask = torch.cat(
|
mask = torch.cat(
|
||||||
[
|
[
|
||||||
torch.ones(len(tokens), dtype=torch.bool, device=action.device),
|
torch.ones(len(tokens), dtype=torch.bool, device=action.device),
|
||||||
torch.zeros(
|
torch.zeros(pad_len, dtype=torch.bool, device=action.device),
|
||||||
self.max_action_tokens - len(tokens), dtype=torch.bool, device=action.device
|
|
||||||
),
|
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
code_mask = torch.nn.functional.pad(code_mask, (0, pad_len), value=False)
|
||||||
# Pad tokens with zeros
|
# Pad tokens with zeros
|
||||||
tokens = torch.nn.functional.pad(tokens, (0, self.max_action_tokens - len(tokens)), value=0)
|
tokens = torch.nn.functional.pad(tokens, (0, pad_len), value=0)
|
||||||
|
|
||||||
tokens_list.append(tokens)
|
tokens_list.append(tokens)
|
||||||
masks_list.append(mask)
|
masks_list.append(mask)
|
||||||
|
code_masks_list.append(code_mask)
|
||||||
|
|
||||||
# Stack into batched tensors
|
# Stack into batched tensors
|
||||||
tokens_batch = torch.stack(tokens_list, dim=0) # (B, max_action_tokens)
|
tokens_batch = torch.stack(tokens_list, dim=0) # (B, max_action_tokens)
|
||||||
masks_batch = torch.stack(masks_list, dim=0) # (B, max_action_tokens)
|
masks_batch = torch.stack(masks_list, dim=0) # (B, max_action_tokens)
|
||||||
|
code_masks_batch = torch.stack(code_masks_list, dim=0) # (B, max_action_tokens)
|
||||||
|
|
||||||
# Remove batch dimension if input was single sample
|
# Remove batch dimension if input was single sample
|
||||||
if single_sample:
|
if single_sample:
|
||||||
tokens_batch = tokens_batch.squeeze(0)
|
tokens_batch = tokens_batch.squeeze(0)
|
||||||
masks_batch = masks_batch.squeeze(0)
|
masks_batch = masks_batch.squeeze(0)
|
||||||
|
code_masks_batch = code_masks_batch.squeeze(0)
|
||||||
|
|
||||||
# Move to the same device as the input
|
# Move to the same device as the input
|
||||||
if device is not None:
|
if device is not None:
|
||||||
tokens_batch = tokens_batch.to(device)
|
tokens_batch = tokens_batch.to(device)
|
||||||
masks_batch = masks_batch.to(device)
|
masks_batch = masks_batch.to(device)
|
||||||
|
code_masks_batch = code_masks_batch.to(device)
|
||||||
|
|
||||||
return tokens_batch, masks_batch
|
return tokens_batch, masks_batch, code_masks_batch
|
||||||
|
|
||||||
def action(self, action: torch.Tensor) -> torch.Tensor:
|
def action(self, action: torch.Tensor) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
This method is not used since we override __call__.
|
This method is not used since we override __call__.
|
||||||
Required by ActionProcessorStep ABC.
|
Required by ActionProcessorStep ABC.
|
||||||
"""
|
"""
|
||||||
tokens, _ = self._tokenize_action(action)
|
tokens, _, _ = self._tokenize_action(action)
|
||||||
return tokens
|
return tokens
|
||||||
|
|
||||||
def get_config(self) -> dict[str, Any]:
|
def get_config(self) -> dict[str, Any]:
|
||||||
@@ -550,6 +572,9 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
config = {
|
config = {
|
||||||
"trust_remote_code": self.trust_remote_code,
|
"trust_remote_code": self.trust_remote_code,
|
||||||
"max_action_tokens": self.max_action_tokens,
|
"max_action_tokens": self.max_action_tokens,
|
||||||
|
"fast_skip_tokens": self.fast_skip_tokens,
|
||||||
|
"paligemma_tokenizer_name": self.paligemma_tokenizer_name,
|
||||||
|
"allow_truncation": self.allow_truncation,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Only save tokenizer_name if it was used to create the tokenizer
|
# Only save tokenizer_name if it was used to create the tokenizer
|
||||||
@@ -558,6 +583,14 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
|||||||
|
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
def save_artifacts(self, save_directory: Path) -> dict[str, str]:
|
||||||
|
artifact_path = Path("action_tokenizer")
|
||||||
|
save_pretrained = getattr(self.action_tokenizer, "save_pretrained", None)
|
||||||
|
if save_pretrained is None:
|
||||||
|
raise TypeError("Action tokenizer must implement save_pretrained() to save a portable pipeline.")
|
||||||
|
save_pretrained(save_directory / artifact_path)
|
||||||
|
return {"action_tokenizer_name": artifact_path.as_posix()}
|
||||||
|
|
||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
|
|||||||
@@ -91,7 +91,7 @@ from lerobot.robots import so_follower # noqa: F401
|
|||||||
from lerobot.teleoperators import gamepad, so_leader # noqa: F401
|
from lerobot.teleoperators import gamepad, so_leader # noqa: F401
|
||||||
from lerobot.teleoperators.utils import TeleopEvents
|
from lerobot.teleoperators.utils import TeleopEvents
|
||||||
from lerobot.utils.device_utils import get_safe_torch_device
|
from lerobot.utils.device_utils import get_safe_torch_device
|
||||||
from lerobot.utils.process import ProcessSignalHandler
|
from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.robot_utils import precise_sleep
|
from lerobot.utils.robot_utils import precise_sleep
|
||||||
from lerobot.utils.transition import (
|
from lerobot.utils.transition import (
|
||||||
@@ -124,9 +124,7 @@ def actor_cli(cfg: TrainRLServerPipelineConfig):
|
|||||||
cfg.validate()
|
cfg.validate()
|
||||||
display_pid = False
|
display_pid = False
|
||||||
if not use_threads(cfg):
|
if not use_threads(cfg):
|
||||||
import torch.multiprocessing as mp
|
ensure_multiprocessing_start_method(cfg.policy.concurrency.multiprocessing_context)
|
||||||
|
|
||||||
mp.set_start_method("spawn")
|
|
||||||
display_pid = True
|
display_pid = True
|
||||||
|
|
||||||
# Create logs directory to ensure it exists
|
# Create logs directory to ensure it exists
|
||||||
|
|||||||
@@ -102,7 +102,7 @@ from lerobot.utils.constants import (
|
|||||||
)
|
)
|
||||||
from lerobot.utils.device_utils import get_safe_torch_device
|
from lerobot.utils.device_utils import get_safe_torch_device
|
||||||
from lerobot.utils.io_utils import load_json, write_json
|
from lerobot.utils.io_utils import load_json, write_json
|
||||||
from lerobot.utils.process import ProcessSignalHandler
|
from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.utils import (
|
from lerobot.utils.utils import (
|
||||||
format_big_number,
|
format_big_number,
|
||||||
@@ -123,9 +123,7 @@ def train_cli(cfg: TrainRLServerPipelineConfig):
|
|||||||
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
||||||
require_package("grpcio", extra="hilserl", import_name="grpc")
|
require_package("grpcio", extra="hilserl", import_name="grpc")
|
||||||
if not use_threads(cfg):
|
if not use_threads(cfg):
|
||||||
import torch.multiprocessing as mp
|
ensure_multiprocessing_start_method(cfg.policy.concurrency.multiprocessing_context)
|
||||||
|
|
||||||
mp.set_start_method("spawn")
|
|
||||||
|
|
||||||
# Use the job_name from the config
|
# Use the job_name from the config
|
||||||
train(
|
train(
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -323,6 +323,10 @@ class LeKiwiClient(Robot):
|
|||||||
np.ndarray: the action sent to the motors, potentially clipped.
|
np.ndarray: the action sent to the motors, potentially clipped.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# Action values may be torch tensors (e.g. replayed from a dataset) or numpy
|
||||||
|
# scalars; json.dumps only serializes Python primitives, so coerce each value to a
|
||||||
|
# plain float before sending.
|
||||||
|
action = {key: float(value) for key, value in action.items()}
|
||||||
self.zmq_cmd_socket.send_string(json.dumps(action)) # action is in motor space
|
self.zmq_cmd_socket.send_string(json.dumps(action)) # action is in motor space
|
||||||
|
|
||||||
# TODO(Steven): Remove the np conversion when it is possible to record a non-numpy array value
|
# TODO(Steven): Remove the np conversion when it is possible to record a non-numpy array value
|
||||||
|
|||||||
@@ -150,9 +150,6 @@ class OpenArmFollower(Robot):
|
|||||||
|
|
||||||
self.configure()
|
self.configure()
|
||||||
|
|
||||||
if self.is_calibrated:
|
|
||||||
self.bus.set_zero_position()
|
|
||||||
|
|
||||||
self.bus.enable_torque()
|
self.bus.enable_torque()
|
||||||
|
|
||||||
logger.info(f"{self} connected.")
|
logger.info(f"{self} connected.")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -326,8 +326,17 @@ class RolloutConfig:
|
|||||||
|
|
||||||
policy_path = parser.get_path_arg("policy")
|
policy_path = parser.get_path_arg("policy")
|
||||||
if policy_path:
|
if policy_path:
|
||||||
cli_overrides = parser.get_cli_overrides("policy")
|
yaml_overrides = parser.get_yaml_overrides("policy")
|
||||||
self.policy = PreTrainedConfig.from_pretrained(policy_path, cli_overrides=cli_overrides)
|
cli_overrides = parser.get_cli_overrides("policy") or []
|
||||||
|
policy_overrides = yaml_overrides + cli_overrides
|
||||||
|
pretrained_revision = parser.parse_arg("pretrained_revision", cli_overrides)
|
||||||
|
if pretrained_revision is None:
|
||||||
|
pretrained_revision = parser.parse_arg("pretrained_revision", yaml_overrides)
|
||||||
|
self.policy = PreTrainedConfig.from_pretrained(
|
||||||
|
policy_path,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
cli_overrides=policy_overrides,
|
||||||
|
)
|
||||||
self.policy.pretrained_path = policy_path
|
self.policy.pretrained_path = policy_path
|
||||||
if self.policy is None:
|
if self.policy is None:
|
||||||
raise ValueError("--policy.path is required for rollout")
|
raise ValueError("--policy.path is required for rollout")
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from threading import Event
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.configs import FeatureType
|
from lerobot.configs import FeatureType, PreTrainedConfig
|
||||||
from lerobot.datasets import (
|
from lerobot.datasets import (
|
||||||
LeRobotDataset,
|
LeRobotDataset,
|
||||||
aggregate_pipeline_dataset_features,
|
aggregate_pipeline_dataset_features,
|
||||||
@@ -159,6 +159,35 @@ class RolloutContext:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy:
|
||||||
|
"""Load policy weights, keeping adapter and base-model revisions independent."""
|
||||||
|
pretrained_revision = policy_config.pretrained_revision
|
||||||
|
policy_class = get_policy_class(policy_config.type)
|
||||||
|
|
||||||
|
if not policy_config.use_peft:
|
||||||
|
return policy_class.from_pretrained(
|
||||||
|
policy_config.pretrained_path,
|
||||||
|
config=policy_config,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
)
|
||||||
|
|
||||||
|
from peft import PeftConfig, PeftModel
|
||||||
|
|
||||||
|
peft_path = policy_config.pretrained_path
|
||||||
|
peft_config = PeftConfig.from_pretrained(peft_path, revision=pretrained_revision)
|
||||||
|
policy = policy_class.from_pretrained(
|
||||||
|
pretrained_name_or_path=peft_config.base_model_name_or_path,
|
||||||
|
config=policy_config,
|
||||||
|
revision=peft_config.revision,
|
||||||
|
)
|
||||||
|
return PeftModel.from_pretrained(
|
||||||
|
policy,
|
||||||
|
peft_path,
|
||||||
|
config=peft_config,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_rollout_context(
|
def build_rollout_context(
|
||||||
cfg: RolloutConfig,
|
cfg: RolloutConfig,
|
||||||
shutdown_event: Event,
|
shutdown_event: Event,
|
||||||
@@ -176,7 +205,6 @@ def build_rollout_context(
|
|||||||
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
||||||
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
||||||
policy_config = cfg.policy
|
policy_config = cfg.policy
|
||||||
policy_class = get_policy_class(policy_config.type)
|
|
||||||
|
|
||||||
if hasattr(policy_config, "compile_model"):
|
if hasattr(policy_config, "compile_model"):
|
||||||
policy_config.compile_model = cfg.use_torch_compile
|
policy_config.compile_model = cfg.use_torch_compile
|
||||||
@@ -187,17 +215,7 @@ def build_rollout_context(
|
|||||||
"Please use `cpu` or `cuda` backend."
|
"Please use `cpu` or `cuda` backend."
|
||||||
)
|
)
|
||||||
|
|
||||||
if policy_config.use_peft:
|
policy = _load_pretrained_policy(policy_config)
|
||||||
from peft import PeftConfig, PeftModel
|
|
||||||
|
|
||||||
peft_path = policy_config.pretrained_path
|
|
||||||
peft_config = PeftConfig.from_pretrained(peft_path)
|
|
||||||
policy = policy_class.from_pretrained(
|
|
||||||
pretrained_name_or_path=peft_config.base_model_name_or_path, config=policy_config
|
|
||||||
)
|
|
||||||
policy = PeftModel.from_pretrained(policy, peft_path, config=peft_config)
|
|
||||||
else:
|
|
||||||
policy = policy_class.from_pretrained(policy_config.pretrained_path, config=policy_config)
|
|
||||||
|
|
||||||
if is_rtc:
|
if is_rtc:
|
||||||
policy.config.rtc_config = cfg.inference.rtc
|
policy.config.rtc_config = cfg.inference.rtc
|
||||||
@@ -392,6 +410,7 @@ def build_rollout_context(
|
|||||||
preprocessor, postprocessor = make_pre_post_processors(
|
preprocessor, postprocessor = make_pre_post_processors(
|
||||||
policy_cfg=policy_config,
|
policy_cfg=policy_config,
|
||||||
pretrained_path=cfg.policy.pretrained_path,
|
pretrained_path=cfg.policy.pretrained_path,
|
||||||
|
pretrained_revision=policy_config.pretrained_revision,
|
||||||
dataset_stats=dataset_stats,
|
dataset_stats=dataset_stats,
|
||||||
preprocessor_overrides={
|
preprocessor_overrides={
|
||||||
"device_processor": {"device": cfg.device},
|
"device_processor": {"device": cfg.device},
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ import pyarrow as pa
|
|||||||
import tqdm
|
import tqdm
|
||||||
from datasets import Dataset, Features, Image
|
from datasets import Dataset, Features, Image
|
||||||
from huggingface_hub import HfApi, snapshot_download
|
from huggingface_hub import HfApi, snapshot_download
|
||||||
|
from huggingface_hub.errors import RevisionNotFoundError
|
||||||
from requests import HTTPError
|
from requests import HTTPError
|
||||||
|
|
||||||
from lerobot.datasets import CODEBASE_VERSION, LeRobotDataset, aggregate_stats
|
from lerobot.datasets import CODEBASE_VERSION, LeRobotDataset, aggregate_stats
|
||||||
@@ -521,7 +522,7 @@ def convert_dataset(
|
|||||||
hub_api = HfApi()
|
hub_api = HfApi()
|
||||||
try:
|
try:
|
||||||
hub_api.delete_tag(repo_id, tag=CODEBASE_VERSION, repo_type="dataset")
|
hub_api.delete_tag(repo_id, tag=CODEBASE_VERSION, repo_type="dataset")
|
||||||
except HTTPError as e:
|
except (HTTPError, RevisionNotFoundError) as e:
|
||||||
print(f"tag={CODEBASE_VERSION} probably doesn't exist. Skipping exception ({e})")
|
print(f"tag={CODEBASE_VERSION} probably doesn't exist. Skipping exception ({e})")
|
||||||
pass
|
pass
|
||||||
hub_api.delete_files(
|
hub_api.delete_files(
|
||||||
|
|||||||
@@ -24,7 +24,14 @@ Example:
|
|||||||
--root=/path/to/dataset \\
|
--root=/path/to/dataset \\
|
||||||
--vlm.model_id=Qwen/Qwen2.5-VL-7B-Instruct
|
--vlm.model_id=Qwen/Qwen2.5-VL-7B-Instruct
|
||||||
|
|
||||||
For distributed runs, see ``examples/annotations/run_hf_job.py``.
|
Pass ``--job.target=<flavor>`` to run the same command on a Hugging Face
|
||||||
|
Jobs GPU instead of this machine (see ``lerobot.jobs.annotate``):
|
||||||
|
|
||||||
|
uv run lerobot-annotate \\
|
||||||
|
--repo_id=user/dataset \\
|
||||||
|
--new_repo_id=user/dataset_annotated \\
|
||||||
|
--push_to_hub=true \\
|
||||||
|
--job.target=h200
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -69,6 +76,14 @@ def _resolve_root(cfg: AnnotationPipelineConfig) -> Path:
|
|||||||
def annotate(cfg: AnnotationPipelineConfig) -> None:
|
def annotate(cfg: AnnotationPipelineConfig) -> None:
|
||||||
"""Run the steerable annotation pipeline against a dataset."""
|
"""Run the steerable annotation pipeline against a dataset."""
|
||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||||
|
|
||||||
|
if cfg.job.is_remote:
|
||||||
|
# Imported lazily: the submitter pulls in LeRobotDataset (the `dataset`
|
||||||
|
# extra), which a local annotation run over --root doesn't need.
|
||||||
|
from lerobot.jobs.annotate import submit_annotate_to_hf
|
||||||
|
|
||||||
|
return submit_annotate_to_hf(cfg)
|
||||||
|
|
||||||
root = _resolve_root(cfg)
|
root = _resolve_root(cfg)
|
||||||
logger.info("annotate: root=%s", root)
|
logger.info("annotate: root=%s", root)
|
||||||
|
|
||||||
|
|||||||
@@ -453,6 +453,9 @@ def eval_policy(
|
|||||||
raise exc from None
|
raise exc from None
|
||||||
|
|
||||||
start = time.time()
|
start = time.time()
|
||||||
|
# Preserve the mode for direct callers. eval_policy_all scopes the mode
|
||||||
|
# around all tasks so parallel evaluations cannot race with each other.
|
||||||
|
was_training = policy.training
|
||||||
policy.eval()
|
policy.eval()
|
||||||
|
|
||||||
# Determine how many batched rollouts we need to get n_episodes. Note that if n_episodes is not evenly
|
# Determine how many batched rollouts we need to get n_episodes. Note that if n_episodes is not evenly
|
||||||
@@ -674,6 +677,8 @@ def eval_policy(
|
|||||||
if save_predicted_video:
|
if save_predicted_video:
|
||||||
info["predicted_video_paths"] = predicted_video_paths
|
info["predicted_video_paths"] = predicted_video_paths
|
||||||
|
|
||||||
|
policy.train(was_training)
|
||||||
|
|
||||||
return info
|
return info
|
||||||
|
|
||||||
|
|
||||||
@@ -1010,40 +1015,48 @@ def eval_policy_all(
|
|||||||
recording_private=recording_private,
|
recording_private=recording_private,
|
||||||
)
|
)
|
||||||
|
|
||||||
if max_parallel_tasks <= 1:
|
# Set the shared policy's mode before launching any workers. Restoring it
|
||||||
prefetch_thread: threading.Thread | None = None
|
# inside individual tasks would let one task enable training mode while
|
||||||
for i, (task_group, task_id, env) in enumerate(tasks):
|
# another task is still evaluating.
|
||||||
if prefetch_thread is not None:
|
was_training = policy.training
|
||||||
prefetch_thread.join()
|
policy.eval()
|
||||||
prefetch_thread = None
|
try:
|
||||||
|
if max_parallel_tasks <= 1:
|
||||||
|
prefetch_thread: threading.Thread | None = None
|
||||||
|
for i, (task_group, task_id, env) in enumerate(tasks):
|
||||||
|
if prefetch_thread is not None:
|
||||||
|
prefetch_thread.join()
|
||||||
|
prefetch_thread = None
|
||||||
|
|
||||||
try:
|
|
||||||
tg, tid, metrics = task_runner(task_group, task_id, env)
|
|
||||||
_accumulate_to(tg, metrics)
|
|
||||||
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
|
||||||
finally:
|
|
||||||
env.close()
|
|
||||||
# Prefetch next task's workers *after* closing current env to prevent
|
|
||||||
# GPU memory overlap between consecutive tasks.
|
|
||||||
if i + 1 < len(tasks):
|
|
||||||
next_env = tasks[i + 1][2]
|
|
||||||
if hasattr(next_env, "_ensure"):
|
|
||||||
prefetch_thread = threading.Thread(target=next_env._ensure, daemon=True)
|
|
||||||
prefetch_thread.start()
|
|
||||||
else:
|
|
||||||
with cf.ThreadPoolExecutor(max_workers=max_parallel_tasks) as executor:
|
|
||||||
fut2meta = {}
|
|
||||||
for task_group, task_id, env in tasks:
|
|
||||||
fut = executor.submit(task_runner, task_group, task_id, env)
|
|
||||||
fut2meta[fut] = (task_group, task_id, env)
|
|
||||||
for fut in cf.as_completed(fut2meta):
|
|
||||||
tg, tid, env = fut2meta[fut]
|
|
||||||
try:
|
try:
|
||||||
tg, tid, metrics = fut.result()
|
tg, tid, metrics = task_runner(task_group, task_id, env)
|
||||||
_accumulate_to(tg, metrics)
|
_accumulate_to(tg, metrics)
|
||||||
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
||||||
finally:
|
finally:
|
||||||
env.close()
|
env.close()
|
||||||
|
# Prefetch next task's workers *after* closing current env to prevent
|
||||||
|
# GPU memory overlap between consecutive tasks.
|
||||||
|
if i + 1 < len(tasks):
|
||||||
|
next_env = tasks[i + 1][2]
|
||||||
|
if hasattr(next_env, "_ensure"):
|
||||||
|
prefetch_thread = threading.Thread(target=next_env._ensure, daemon=True)
|
||||||
|
prefetch_thread.start()
|
||||||
|
else:
|
||||||
|
with cf.ThreadPoolExecutor(max_workers=max_parallel_tasks) as executor:
|
||||||
|
fut2meta = {}
|
||||||
|
for task_group, task_id, env in tasks:
|
||||||
|
fut = executor.submit(task_runner, task_group, task_id, env)
|
||||||
|
fut2meta[fut] = (task_group, task_id, env)
|
||||||
|
for fut in cf.as_completed(fut2meta):
|
||||||
|
tg, tid, env = fut2meta[fut]
|
||||||
|
try:
|
||||||
|
tg, tid, metrics = fut.result()
|
||||||
|
_accumulate_to(tg, metrics)
|
||||||
|
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
||||||
|
finally:
|
||||||
|
env.close()
|
||||||
|
finally:
|
||||||
|
policy.train(was_training)
|
||||||
|
|
||||||
# compute aggregated metrics helper (robust to lists/scalars)
|
# compute aggregated metrics helper (robust to lists/scalars)
|
||||||
def _agg_from_list(xs):
|
def _agg_from_list(xs):
|
||||||
|
|||||||
@@ -453,9 +453,11 @@ def record(
|
|||||||
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
||||||
)
|
)
|
||||||
|
|
||||||
robot.connect()
|
# Connect the teleoperator before the robot so the robot isn't left idle (and possibly
|
||||||
|
# tripping a firmware watchdog) during teleop init. Matches lerobot_teleoperate.py.
|
||||||
if teleop is not None:
|
if teleop is not None:
|
||||||
teleop.connect()
|
teleop.connect()
|
||||||
|
robot.connect()
|
||||||
|
|
||||||
listener, events = init_keyboard_listener()
|
listener, events = init_keyboard_listener()
|
||||||
|
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ from lerobot.robots import ( # noqa: F401
|
|||||||
earthrover_mini_plus,
|
earthrover_mini_plus,
|
||||||
hope_jr,
|
hope_jr,
|
||||||
koch_follower,
|
koch_follower,
|
||||||
|
lekiwi,
|
||||||
make_robot_from_config,
|
make_robot_from_config,
|
||||||
omx_follower,
|
omx_follower,
|
||||||
openarm_follower,
|
openarm_follower,
|
||||||
|
|||||||
@@ -51,19 +51,7 @@ from lerobot.teleoperators import ( # noqa: F401
|
|||||||
rebot_102_leader,
|
rebot_102_leader,
|
||||||
so_leader,
|
so_leader,
|
||||||
)
|
)
|
||||||
|
from lerobot.utils.import_utils import register_third_party_plugins
|
||||||
COMPATIBLE_DEVICES = [
|
|
||||||
"koch_follower",
|
|
||||||
"koch_leader",
|
|
||||||
"omx_follower",
|
|
||||||
"omx_leader",
|
|
||||||
"openarm_mini",
|
|
||||||
"so100_follower",
|
|
||||||
"so100_leader",
|
|
||||||
"so101_follower",
|
|
||||||
"so101_leader",
|
|
||||||
"lekiwi",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -80,18 +68,19 @@ class SetupConfig:
|
|||||||
|
|
||||||
@draccus.wrap()
|
@draccus.wrap()
|
||||||
def setup_motors(cfg: SetupConfig):
|
def setup_motors(cfg: SetupConfig):
|
||||||
if cfg.device.type not in COMPATIBLE_DEVICES:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
if isinstance(cfg.device, RobotConfig):
|
if isinstance(cfg.device, RobotConfig):
|
||||||
device = make_robot_from_config(cfg.device)
|
device = make_robot_from_config(cfg.device)
|
||||||
else:
|
else:
|
||||||
device = make_teleoperator_from_config(cfg.device)
|
device = make_teleoperator_from_config(cfg.device)
|
||||||
|
|
||||||
device.setup_motors()
|
setup = getattr(device, "setup_motors", None)
|
||||||
|
if not callable(setup):
|
||||||
|
raise NotImplementedError(f"Device type '{cfg.device.type}' does not support motor setup.")
|
||||||
|
setup()
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
register_third_party_plugins()
|
||||||
setup_motors()
|
setup_motors()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -71,6 +71,16 @@ from lerobot.utils.utils import (
|
|||||||
from .lerobot_eval import eval_policy_all
|
from .lerobot_eval import eval_policy_all
|
||||||
|
|
||||||
|
|
||||||
|
def _dataloader_worker_kwargs(cfg: TrainPipelineConfig) -> dict[str, Any]:
|
||||||
|
"""Return worker-only DataLoader options, disabling them for single-process loading."""
|
||||||
|
workers_enabled = cfg.num_workers > 0
|
||||||
|
return {
|
||||||
|
"prefetch_factor": cfg.prefetch_factor if workers_enabled else None,
|
||||||
|
"persistent_workers": cfg.persistent_workers and workers_enabled,
|
||||||
|
"multiprocessing_context": cfg.dataloader_multiprocessing_context if workers_enabled else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def update_policy(
|
def update_policy(
|
||||||
train_metrics: MetricsTracker,
|
train_metrics: MetricsTracker,
|
||||||
policy: PreTrainedPolicy,
|
policy: PreTrainedPolicy,
|
||||||
@@ -473,8 +483,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
pin_memory=device.type == "cuda",
|
pin_memory=device.type == "cuda",
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
collate_fn=collate_fn,
|
collate_fn=collate_fn,
|
||||||
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None,
|
**_dataloader_worker_kwargs(cfg),
|
||||||
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build eval dataloader if a held-out split exists
|
# Build eval dataloader if a held-out split exists
|
||||||
@@ -500,8 +509,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
pin_memory=device.type == "cuda",
|
pin_memory=device.type == "cuda",
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
collate_fn=eval_collate_fn,
|
collate_fn=eval_collate_fn,
|
||||||
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None,
|
**_dataloader_worker_kwargs(cfg),
|
||||||
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Prepare everything with accelerator
|
# Prepare everything with accelerator
|
||||||
|
|||||||
@@ -23,3 +23,5 @@ from ..config import TeleoperatorConfig
|
|||||||
@dataclass
|
@dataclass
|
||||||
class GamepadTeleopConfig(TeleoperatorConfig):
|
class GamepadTeleopConfig(TeleoperatorConfig):
|
||||||
use_gripper: bool = True
|
use_gripper: bool = True
|
||||||
|
# Use hidapi instead of pygame for controllers that pygame cannot detect reliably.
|
||||||
|
hidapi_fallback: bool = False
|
||||||
|
|||||||
@@ -14,6 +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.
|
||||||
|
|
||||||
|
import logging
|
||||||
import sys
|
import sys
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -27,6 +28,8 @@ from ..teleoperator import Teleoperator
|
|||||||
from ..utils import TeleopEvents
|
from ..utils import TeleopEvents
|
||||||
from .configuration_gamepad import GamepadTeleopConfig
|
from .configuration_gamepad import GamepadTeleopConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class GripperAction(IntEnum):
|
class GripperAction(IntEnum):
|
||||||
CLOSE = 0
|
CLOSE = 0
|
||||||
@@ -56,6 +59,13 @@ class GamepadTeleop(Teleoperator):
|
|||||||
|
|
||||||
self.gamepad = None
|
self.gamepad = None
|
||||||
|
|
||||||
|
self.hidapi_fallback = config.hidapi_fallback
|
||||||
|
if sys.platform == "darwin" and not self.hidapi_fallback:
|
||||||
|
logger.warning(
|
||||||
|
"On macOS, pygame may not reliably detect input from some controllers. "
|
||||||
|
"If you experience issues, set `hidapi_fallback=true`."
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_features(self) -> dict:
|
def action_features(self) -> dict:
|
||||||
if self.config.use_gripper:
|
if self.config.use_gripper:
|
||||||
@@ -76,9 +86,7 @@ class GamepadTeleop(Teleoperator):
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
def connect(self) -> None:
|
def connect(self) -> None:
|
||||||
# use HidApi for macos
|
if self.hidapi_fallback:
|
||||||
if sys.platform == "darwin":
|
|
||||||
# NOTE: On macOS, pygame doesn’t reliably detect input from some controllers so we fall back to hidapi
|
|
||||||
from .gamepad_utils import GamepadControllerHID as Gamepad
|
from .gamepad_utils import GamepadControllerHID as Gamepad
|
||||||
else:
|
else:
|
||||||
from .gamepad_utils import GamepadController as Gamepad
|
from .gamepad_utils import GamepadController as Gamepad
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from torch.utils.data._utils.collate import default_collate
|
|||||||
|
|
||||||
from lerobot.datasets.language import LANGUAGE_COLUMNS
|
from lerobot.datasets.language import LANGUAGE_COLUMNS
|
||||||
|
|
||||||
_PYTHON_LIST_KEYS = {"messages", "message_streams", "target_message_indices"}
|
_PYTHON_LIST_KEYS = {"messages", "message_streams", "target_message_indices", *LANGUAGE_COLUMNS}
|
||||||
|
|
||||||
|
|
||||||
def lerobot_collate_fn(batch: list[dict[str, Any] | None]) -> dict[str, Any] | None:
|
def lerobot_collate_fn(batch: list[dict[str, Any] | None]) -> dict[str, Any] | None:
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ OBS_IMAGES = OBS_IMAGE + "s"
|
|||||||
OBS_LANGUAGE = OBS_STR + ".language"
|
OBS_LANGUAGE = OBS_STR + ".language"
|
||||||
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
||||||
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
||||||
|
OBS_LANGUAGE_CAUSAL_MARKS = OBS_LANGUAGE + ".causal_marks"
|
||||||
OBS_LANGUAGE_SUBTASK = OBS_STR + ".subtask"
|
OBS_LANGUAGE_SUBTASK = OBS_STR + ".subtask"
|
||||||
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
|
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
|
||||||
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
|
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
|
||||||
@@ -34,6 +35,7 @@ ACTION = "action"
|
|||||||
ACTION_PREFIX = ACTION + "."
|
ACTION_PREFIX = ACTION + "."
|
||||||
ACTION_TOKENS = ACTION + ".tokens"
|
ACTION_TOKENS = ACTION + ".tokens"
|
||||||
ACTION_TOKEN_MASK = ACTION + ".token_mask"
|
ACTION_TOKEN_MASK = ACTION + ".token_mask"
|
||||||
|
ACTION_CODE_TOKEN_MASK = ACTION + ".code_token_mask"
|
||||||
REWARD = "next.reward"
|
REWARD = "next.reward"
|
||||||
TRUNCATED = "next.truncated"
|
TRUNCATED = "next.truncated"
|
||||||
DONE = "next.done"
|
DONE = "next.done"
|
||||||
|
|||||||
@@ -16,11 +16,39 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import multiprocessing
|
||||||
import os
|
import os
|
||||||
import signal
|
import signal
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_multiprocessing_start_method(start_method: str | None) -> None:
|
||||||
|
"""Set a multiprocessing start method once, or verify the existing method matches.
|
||||||
|
|
||||||
|
Passing ``None`` leaves Python's process-wide default untouched. This is useful
|
||||||
|
when LeRobot is embedded in an application that owns multiprocessing setup.
|
||||||
|
"""
|
||||||
|
if start_method is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
available_methods = multiprocessing.get_all_start_methods()
|
||||||
|
if start_method not in available_methods:
|
||||||
|
raise ValueError(
|
||||||
|
f"Multiprocessing start method must be one of {available_methods} on this platform, "
|
||||||
|
f"got {start_method!r}."
|
||||||
|
)
|
||||||
|
|
||||||
|
current_method = multiprocessing.get_start_method(allow_none=True)
|
||||||
|
if current_method is None:
|
||||||
|
multiprocessing.set_start_method(start_method)
|
||||||
|
elif current_method != start_method:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Multiprocessing start method is already {current_method!r}; cannot change it to "
|
||||||
|
f"{start_method!r}. Set the configured multiprocessing context to null to keep the "
|
||||||
|
"application's existing method, or launch LeRobot in a fresh process."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ProcessSignalHandler:
|
class ProcessSignalHandler:
|
||||||
"""Utility class to attach graceful shutdown signal handlers.
|
"""Utility class to attach graceful shutdown signal handlers.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -29,6 +29,13 @@ def test_message_recipe_validates_unknown_binding():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_canonical_recipe_loads():
|
||||||
|
"""The canonical PI052 blend YAML loads + validates."""
|
||||||
|
recipe = TrainingRecipe.from_yaml(Path("src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml"))
|
||||||
|
assert recipe.blend is not None
|
||||||
|
assert sum(c.weight for c in recipe.blend.values()) == pytest.approx(1.0)
|
||||||
|
|
||||||
|
|
||||||
def test_message_turn_requires_a_stream():
|
def test_message_turn_requires_a_stream():
|
||||||
"""Every turn must declare a stream — None is rejected at construction.
|
"""Every turn must declare a stream — None is rejected at construction.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -204,6 +204,38 @@ def test_clear_resets_buffer(tmp_path):
|
|||||||
assert dataset.writer.episode_buffer["size"] == 0
|
assert dataset.writer.episode_buffer["size"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_clear_removes_video_frame_staging_dir(tmp_path):
|
||||||
|
"""clear_episode_buffer() removes PNG staging dirs for video features."""
|
||||||
|
video_key = "observation.images.cam"
|
||||||
|
features = {
|
||||||
|
video_key: {
|
||||||
|
"dtype": "video",
|
||||||
|
"shape": (64, 96, 3),
|
||||||
|
"names": ["height", "width", "channels"],
|
||||||
|
},
|
||||||
|
"action": {"dtype": "float32", "shape": (2,), "names": None},
|
||||||
|
}
|
||||||
|
dataset = LeRobotDataset.create(
|
||||||
|
repo_id=DUMMY_REPO_ID,
|
||||||
|
fps=DEFAULT_FPS,
|
||||||
|
features=features,
|
||||||
|
root=tmp_path / "ds",
|
||||||
|
use_videos=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
dataset.add_frame(_make_frame(features))
|
||||||
|
video_staging_dir = (
|
||||||
|
dataset.root
|
||||||
|
/ Path(DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0)).parent
|
||||||
|
)
|
||||||
|
assert video_staging_dir.is_dir()
|
||||||
|
|
||||||
|
dataset.clear_episode_buffer()
|
||||||
|
|
||||||
|
assert dataset.writer.episode_buffer["size"] == 0
|
||||||
|
assert not video_staging_dir.exists()
|
||||||
|
|
||||||
|
|
||||||
def test_finalize_is_idempotent(tmp_path):
|
def test_finalize_is_idempotent(tmp_path):
|
||||||
"""Calling finalize() twice does not raise."""
|
"""Calling finalize() twice does not raise."""
|
||||||
dataset = LeRobotDataset.create(
|
dataset = LeRobotDataset.create(
|
||||||
|
|||||||
@@ -114,6 +114,20 @@ def test_dataset_initialization(tmp_path, lerobot_dataset_factory):
|
|||||||
assert dataset.num_frames == len(dataset)
|
assert dataset.num_frames == len(dataset)
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_slice(tmp_path, lerobot_dataset_factory):
|
||||||
|
dataset = lerobot_dataset_factory(
|
||||||
|
root=tmp_path / "test", total_episodes=3, total_frames=30, use_videos=False
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(dataset[:5]) == 5
|
||||||
|
assert len(dataset[::2]) == (len(dataset) + 1) // 2
|
||||||
|
assert [item["index"].item() for item in dataset[4::-1]] == [4, 3, 2, 1, 0]
|
||||||
|
assert [item["index"].item() for item in dataset[-3:]] == list(range(len(dataset) - 3, len(dataset)))
|
||||||
|
assert dataset[len(dataset) :] == []
|
||||||
|
assert isinstance(dataset[0], dict)
|
||||||
|
assert dataset[:1][0].keys() == dataset[0].keys()
|
||||||
|
|
||||||
|
|
||||||
# TODO(rcadene, aliberts): do not run LeRobotDataset.create, instead refactor LeRobotDatasetMetadata.create
|
# TODO(rcadene, aliberts): do not run LeRobotDataset.create, instead refactor LeRobotDatasetMetadata.create
|
||||||
# and test the small resulting function that validates the features
|
# and test the small resulting function that validates the features
|
||||||
def test_dataset_feature_with_forward_slash_raises_error():
|
def test_dataset_feature_with_forward_slash_raises_error():
|
||||||
@@ -1741,6 +1755,38 @@ def test_delta_timestamps_query_returns_correct_values(tmp_path, empty_lerobot_d
|
|||||||
assert is_pad == [True, False], f"Expected [True, False], got {is_pad}"
|
assert is_pad == [True, False], f"Expected [True, False], got {is_pad}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_slice_with_delta_timestamps(tmp_path, empty_lerobot_dataset_factory):
|
||||||
|
features = {
|
||||||
|
"observation.state": {"dtype": "float32", "shape": (1,), "names": ["x"]},
|
||||||
|
}
|
||||||
|
dataset = empty_lerobot_dataset_factory(
|
||||||
|
root=tmp_path / "test_slice_delta", features=features, use_videos=False, fps=10
|
||||||
|
)
|
||||||
|
|
||||||
|
for frame_idx in range(5):
|
||||||
|
dataset.add_frame(
|
||||||
|
{
|
||||||
|
"observation.state": torch.tensor([frame_idx], dtype=torch.float32),
|
||||||
|
"task": "task_0",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
dataset.save_episode()
|
||||||
|
dataset.finalize()
|
||||||
|
|
||||||
|
sliced_dataset = LeRobotDataset(
|
||||||
|
dataset.repo_id,
|
||||||
|
root=dataset.root,
|
||||||
|
delta_timestamps={"observation.state": [-0.1, 0.0]},
|
||||||
|
tolerance_s=0.04,
|
||||||
|
)
|
||||||
|
|
||||||
|
items = sliced_dataset[:2]
|
||||||
|
|
||||||
|
assert items[0]["observation.state"].tolist() == [0.0, 0.0]
|
||||||
|
assert items[0]["observation.state_is_pad"].tolist() == [True, False]
|
||||||
|
assert items[1]["observation.state"].tolist() == [0.0, 1.0]
|
||||||
|
|
||||||
|
|
||||||
def test_episode_filter_filters_dataset(tmp_path, lerobot_dataset_factory):
|
def test_episode_filter_filters_dataset(tmp_path, lerobot_dataset_factory):
|
||||||
"""episode_filter on LeRobotDataset narrows the loaded dataset to matching episodes."""
|
"""episode_filter on LeRobotDataset narrows the loaded dataset to matching episodes."""
|
||||||
dataset = lerobot_dataset_factory(root=tmp_path / "test", total_episodes=8, total_frames=200)
|
dataset = lerobot_dataset_factory(root=tmp_path / "test", total_episodes=8, total_frames=200)
|
||||||
|
|||||||
@@ -343,6 +343,84 @@ def test_resolve_task_explicit_override_beats_rephrasings():
|
|||||||
assert rendered["messages"][0]["content"] == "explicit override wins"
|
assert rendered["messages"][0]["content"] == "explicit override wins"
|
||||||
|
|
||||||
|
|
||||||
|
def test_flow_only_low_level_recipe_renders_without_target():
|
||||||
|
"""Regression: a flow-only ``low_level`` recipe has no ``target`` turn —
|
||||||
|
its supervision is the action-expert flow loss, not text-CE. It must
|
||||||
|
still render (not ``None``), otherwise every blend draw of it is dropped
|
||||||
|
and the action expert never receives a flow loss."""
|
||||||
|
recipe = TrainingRecipe(
|
||||||
|
messages=[
|
||||||
|
MessageTurn(
|
||||||
|
role="user",
|
||||||
|
content="${subtask}",
|
||||||
|
stream="low_level",
|
||||||
|
if_present="subtask",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
bindings={"subtask": "active_at(t, style=subtask)"},
|
||||||
|
)
|
||||||
|
|
||||||
|
rendered = render_sample(
|
||||||
|
recipe=recipe,
|
||||||
|
persistent=PERSISTENT,
|
||||||
|
events=[],
|
||||||
|
t=0.5,
|
||||||
|
sample_idx=0,
|
||||||
|
task="clean kitchen",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert rendered is not None
|
||||||
|
assert rendered["messages"] == [{"role": "user", "content": "subtask 0"}]
|
||||||
|
assert rendered["message_streams"] == ["low_level"]
|
||||||
|
assert rendered["target_message_indices"] == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_vqa_frame_is_consumed_over_the_weighted_blend():
|
||||||
|
"""A frame carrying a VQA annotation renders the ``ask_vqa*`` sub-recipe
|
||||||
|
even when its blend weight is tiny — VQA annotations are sparse and must
|
||||||
|
never be wasted on a subtask/action draw."""
|
||||||
|
recipe = TrainingRecipe(
|
||||||
|
blend={
|
||||||
|
"high_level_subtask": TrainingRecipe(
|
||||||
|
weight=0.99,
|
||||||
|
messages=[
|
||||||
|
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||||
|
MessageTurn(role="assistant", content="a subtask", stream="high_level", target=True),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"ask_vqa_top": TrainingRecipe(
|
||||||
|
weight=0.01,
|
||||||
|
bindings={
|
||||||
|
"vqa_query": "emitted_at(t, style=vqa, role=user, camera=observation.images.top)",
|
||||||
|
"vqa": "emitted_at(t, style=vqa, role=assistant, camera=observation.images.top)",
|
||||||
|
},
|
||||||
|
messages=[
|
||||||
|
MessageTurn(
|
||||||
|
role="user", content="${vqa_query}", stream="high_level", if_present="vqa_query"
|
||||||
|
),
|
||||||
|
MessageTurn(
|
||||||
|
role="assistant",
|
||||||
|
content="${vqa}",
|
||||||
|
stream="high_level",
|
||||||
|
target=True,
|
||||||
|
if_present="vqa",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
# A frame WITH a vqa event renders VQA on every sample_idx, despite the
|
||||||
|
# ask_vqa weight being only 0.01.
|
||||||
|
for sample_idx in range(20):
|
||||||
|
rendered = render_sample(
|
||||||
|
recipe=recipe, persistent=PERSISTENT, events=EVENTS_AT_1, t=1.0, sample_idx=sample_idx, task="x"
|
||||||
|
)
|
||||||
|
assert rendered["messages"][-1]["content"] == '{"count": 2}', sample_idx
|
||||||
|
# A frame WITHOUT a vqa event falls back to the normal weighted blend.
|
||||||
|
rendered = render_sample(recipe=recipe, persistent=PERSISTENT, events=[], t=1.0, sample_idx=0, task="x")
|
||||||
|
assert rendered["messages"][-1]["content"] == "a subtask"
|
||||||
|
|
||||||
|
|
||||||
def test_emitted_at_persistent_tolerates_small_timestamp_drift():
|
def test_emitted_at_persistent_tolerates_small_timestamp_drift():
|
||||||
"""Persistent ``emitted_at`` should match within EMITTED_AT_TOLERANCE_S
|
"""Persistent ``emitted_at`` should match within EMITTED_AT_TOLERANCE_S
|
||||||
so callers that derive ``t`` arithmetically (``frame_idx / fps``) still
|
so callers that derive ``t`` arithmetically (``frame_idx / fps``) still
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from datasets import Dataset # noqa: E402
|
|||||||
from lerobot.datasets.io_utils import (
|
from lerobot.datasets.io_utils import (
|
||||||
hf_transform_to_torch,
|
hf_transform_to_torch,
|
||||||
)
|
)
|
||||||
from lerobot.datasets.sampler import EpisodeAwareSampler
|
from lerobot.datasets.sampler import EpisodeAwareSampler, compute_sampler_state
|
||||||
|
|
||||||
|
|
||||||
def calculate_episode_data_index(hf_dataset: Dataset) -> dict[str, torch.Tensor]:
|
def calculate_episode_data_index(hf_dataset: Dataset) -> dict[str, torch.Tensor]:
|
||||||
@@ -154,8 +154,6 @@ def test_partial_episode_drop_warns(caplog):
|
|||||||
|
|
||||||
# --- seeded (seed, epoch) shuffling, resume, and state ---
|
# --- seeded (seed, epoch) shuffling, resume, and state ---
|
||||||
|
|
||||||
from lerobot.datasets.sampler import compute_sampler_state # noqa: E402
|
|
||||||
|
|
||||||
EPISODE_BOUNDS = ([0, 2, 3], [2, 3, 6]) # episodes of 2, 1 and 3 frames
|
EPISODE_BOUNDS = ([0, 2, 3], [2, 3, 6]) # episodes of 2, 1 and 3 frames
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -13,12 +13,16 @@
|
|||||||
# 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 numpy as np
|
import numpy as np
|
||||||
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.datasets.utils import safe_shard
|
from lerobot.datasets.utils import safe_shard
|
||||||
from lerobot.utils.constants import ACTION
|
from lerobot.utils.constants import ACTION
|
||||||
@@ -71,6 +75,40 @@ def get_frames_expected_order(streaming_ds: StreamingLeRobotDataset) -> list[int
|
|||||||
return expected_indices
|
return expected_indices
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("token", ["hf_test_token", True, False])
|
||||||
|
@pytest.mark.parametrize("from_local", [False, True])
|
||||||
|
def test_streaming_dataset_forwards_hub_token_only_for_remote_data(tmp_path, monkeypatch, token, from_local):
|
||||||
|
requested_root = tmp_path / "local" if from_local else None
|
||||||
|
metadata = SimpleNamespace(
|
||||||
|
root=requested_root or tmp_path / "snapshot",
|
||||||
|
revision=streaming_dataset_module.CODEBASE_VERSION,
|
||||||
|
_version=streaming_dataset_module.CODEBASE_VERSION,
|
||||||
|
features={},
|
||||||
|
depth_keys=[],
|
||||||
|
image_keys=[],
|
||||||
|
rescale_depth_stats=Mock(),
|
||||||
|
)
|
||||||
|
metadata_cls = Mock(return_value=metadata)
|
||||||
|
load_dataset = Mock(return_value=SimpleNamespace(num_shards=1))
|
||||||
|
monkeypatch.setattr(streaming_dataset_module, "LeRobotDatasetMetadata", metadata_cls)
|
||||||
|
monkeypatch.setattr(streaming_dataset_module, "load_dataset", load_dataset)
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
if from_local:
|
||||||
|
assert "token" not in load_dataset.call_args.kwargs
|
||||||
|
else:
|
||||||
|
assert load_dataset.call_args.kwargs["token"] is 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
|
||||||
|
|||||||
@@ -35,6 +35,17 @@ def test_unknown_type():
|
|||||||
make_env_config("nonexistent")
|
make_env_config("nonexistent")
|
||||||
|
|
||||||
|
|
||||||
|
def test_libero_fps_controls_simulator_frequency():
|
||||||
|
cfg = LiberoEnv(fps=17)
|
||||||
|
|
||||||
|
assert cfg.gym_kwargs["control_freq"] == 17
|
||||||
|
|
||||||
|
|
||||||
|
def test_libero_rejects_nonpositive_fps():
|
||||||
|
with pytest.raises(ValueError, match="fps must be positive"):
|
||||||
|
LiberoEnv(fps=0)
|
||||||
|
|
||||||
|
|
||||||
def test_identity_processors():
|
def test_identity_processors():
|
||||||
"""Base class get_env_processors() returns identity pipelines."""
|
"""Base class get_env_processors() returns identity pipelines."""
|
||||||
cfg = make_env_config("aloha")
|
cfg = make_env_config("aloha")
|
||||||
|
|||||||
@@ -0,0 +1,245 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
import shlex
|
||||||
|
import sys
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import draccus
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
||||||
|
|
||||||
|
from lerobot.annotations.steerable_pipeline.config import (
|
||||||
|
DEFAULT_ANNOTATE_JOB_IMAGE,
|
||||||
|
AnnotationJobConfig,
|
||||||
|
AnnotationPipelineConfig,
|
||||||
|
)
|
||||||
|
from lerobot.jobs.annotate import build_pod_command, build_pod_setup, submit_annotate_to_hf
|
||||||
|
|
||||||
|
|
||||||
|
def _parse(*args):
|
||||||
|
return draccus.parse(AnnotationPipelineConfig, args=list(args))
|
||||||
|
|
||||||
|
|
||||||
|
def _set_argv(monkeypatch, *args):
|
||||||
|
monkeypatch.setattr(sys, "argv", ["lerobot-annotate", *args])
|
||||||
|
|
||||||
|
|
||||||
|
# --- config ----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_annotation_job_defaults_are_local_with_vllm_image():
|
||||||
|
cfg = AnnotationJobConfig()
|
||||||
|
assert cfg.target is None
|
||||||
|
assert cfg.is_remote is False
|
||||||
|
assert cfg.image == DEFAULT_ANNOTATE_JOB_IMAGE
|
||||||
|
assert cfg.timeout == "2h"
|
||||||
|
assert cfg.lerobot_ref == "main"
|
||||||
|
|
||||||
|
|
||||||
|
def test_annotation_config_parses_job_target():
|
||||||
|
cfg = _parse("--repo_id", "u/d", "--job.target", "h200")
|
||||||
|
assert cfg.job.target == "h200"
|
||||||
|
assert cfg.job.is_remote is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_annotation_config_defaults_to_local():
|
||||||
|
assert _parse("--repo_id", "u/d").job.is_remote is False
|
||||||
|
|
||||||
|
|
||||||
|
# --- pod command -----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_setup_installs_requested_ref():
|
||||||
|
setup = build_pod_setup("my-branch")
|
||||||
|
assert "git+https://github.com/huggingface/lerobot.git@my-branch" in setup
|
||||||
|
# The vLLM image has neither ffmpeg (video decode) nor lerobot's pinned deps.
|
||||||
|
assert "ffmpeg" in setup
|
||||||
|
assert "'draccus==0.10.0'" in setup
|
||||||
|
|
||||||
|
|
||||||
|
def _annotate_argv(command):
|
||||||
|
"""Extract the `lerobot-annotate ...` argv from a `bash -c` pod command."""
|
||||||
|
assert command[:2] == ["bash", "-c"]
|
||||||
|
_setup, _, annotate = command[2].rpartition(" && ")
|
||||||
|
return shlex.split(annotate)
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_command_forwards_user_flags_and_pins_local_target():
|
||||||
|
command = build_pod_command(
|
||||||
|
"u/d",
|
||||||
|
"main",
|
||||||
|
["--repo_id=u/d", "--new_repo_id=u/d_annotated", "--push_to_hub=true", "--job.target=h200"],
|
||||||
|
)
|
||||||
|
argv = _annotate_argv(command)
|
||||||
|
assert argv[0] == "lerobot-annotate"
|
||||||
|
# --job.* is client-side orchestration; the pod must not re-dispatch itself.
|
||||||
|
assert not any(a.startswith("--job.") for a in argv[1:-1])
|
||||||
|
assert argv[-1] == "--job.target=local"
|
||||||
|
assert "--new_repo_id=u/d_annotated" in argv
|
||||||
|
assert "--push_to_hub=true" in argv
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_command_replaces_host_local_root_with_repo_id():
|
||||||
|
"""--root points at a directory only the client has; the pod resolves by repo_id."""
|
||||||
|
command = build_pod_command("u/d", "main", ["--root", "/home/me/datasets/d", "--seed=7"])
|
||||||
|
argv = _annotate_argv(command)
|
||||||
|
assert "--root" not in argv
|
||||||
|
assert "/home/me/datasets/d" not in argv
|
||||||
|
assert argv.count("--repo_id=u/d") == 1
|
||||||
|
assert "--seed=7" in argv
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_command_does_not_duplicate_repo_id():
|
||||||
|
command = build_pod_command("u/d", "main", ["--repo_id", "u/d"])
|
||||||
|
assert _annotate_argv(command).count("--repo_id=u/d") == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_command_quotes_flags_containing_spaces_and_json():
|
||||||
|
"""serve_command and chat_template_kwargs must survive the trip through `bash -c`."""
|
||||||
|
serve = "--vlm.serve_command=vllm serve Qwen/Qwen3.6-27B --max-model-len 32768 --port {port}"
|
||||||
|
kwargs = '--vlm.chat_template_kwargs={"enable_thinking": false}'
|
||||||
|
command = build_pod_command("u/d", "main", [serve, kwargs])
|
||||||
|
argv = _annotate_argv(command)
|
||||||
|
assert serve in argv
|
||||||
|
assert kwargs in argv
|
||||||
|
|
||||||
|
|
||||||
|
# --- submission ------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_submit_requires_login(monkeypatch):
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: None)
|
||||||
|
with pytest.raises(RuntimeError, match="hf auth login"):
|
||||||
|
submit_annotate_to_hf(_parse("--repo_id", "u/d", "--job.target", "h200"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_submit_requires_repo_id(monkeypatch):
|
||||||
|
"""A remote run over --root alone can't work: the pod can't see the client's disk."""
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
cfg = _parse("--root", "/tmp/d", "--job.target", "h200")
|
||||||
|
with pytest.raises(ValueError, match="--repo_id"):
|
||||||
|
submit_annotate_to_hf(cfg)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("arg", ["--config_path=annotate.yaml", "--vlm=vlm.yaml", "--job=job.yaml"])
|
||||||
|
def test_submit_rejects_local_config_files(monkeypatch, arg):
|
||||||
|
"""draccus takes a config file for the whole config and for each nested one; the
|
||||||
|
pod can read none of them, so a remote run must refuse rather than drop them."""
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
_set_argv(monkeypatch, arg, "--job.target=h200")
|
||||||
|
cfg = _parse("--repo_id", "u/d", "--job.target", "h200")
|
||||||
|
with pytest.raises(ValueError, match="cannot read config files"):
|
||||||
|
submit_annotate_to_hf(cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def test_pod_command_drops_bare_job_config_file_arg():
|
||||||
|
"""`--job` isn't caught by the `--job.` prefix, and could carry a remote target
|
||||||
|
that would make the pod submit a job of its own — recursively."""
|
||||||
|
argv = _annotate_argv(build_pod_command("u/d", "main", ["--job", "job.yaml", "--seed=7"]))
|
||||||
|
assert "--job" not in argv
|
||||||
|
assert "job.yaml" not in argv
|
||||||
|
assert argv[-1] == "--job.target=local"
|
||||||
|
|
||||||
|
|
||||||
|
def test_submit_dispatches_job(monkeypatch):
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.HfApi", lambda token=None: MagicMock())
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.ensure_dataset_available", lambda *a, **kw: None)
|
||||||
|
|
||||||
|
run_job_calls = []
|
||||||
|
|
||||||
|
def fake_run_job(**kwargs):
|
||||||
|
run_job_calls.append(kwargs)
|
||||||
|
return MagicMock(id="job-123")
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.run_job", fake_run_job)
|
||||||
|
_set_argv(monkeypatch, "--repo_id=u/d", "--push_to_hub=true", "--job.target=h200", "--job.detach=true")
|
||||||
|
|
||||||
|
cfg = _parse("--repo_id", "u/d", "--push_to_hub", "true", "--job.target", "h200", "--job.detach", "true")
|
||||||
|
submit_annotate_to_hf(cfg)
|
||||||
|
|
||||||
|
assert len(run_job_calls) == 1
|
||||||
|
call = run_job_calls[0]
|
||||||
|
assert call["flavor"] == "h200"
|
||||||
|
assert call["image"] == DEFAULT_ANNOTATE_JOB_IMAGE
|
||||||
|
assert call["timeout"] == "2h"
|
||||||
|
# The Hub token is forwarded so the pod can pull a private dataset and push the result.
|
||||||
|
assert call["secrets"]["HF_TOKEN"] == "tok"
|
||||||
|
assert call["labels"].get("lerobot") == "true"
|
||||||
|
argv = _annotate_argv(call["command"])
|
||||||
|
assert argv[0] == "lerobot-annotate"
|
||||||
|
assert "--push_to_hub=true" in argv
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.timeout(15)
|
||||||
|
def test_submit_follows_job_to_completion(monkeypatch, capsys):
|
||||||
|
"""Non-detach path must stream logs and RETURN (not hang) once the job is terminal.
|
||||||
|
|
||||||
|
Exercises the `follow_job` helper shared with the training submitter from the
|
||||||
|
annotation side, which is why the job-state patches target `lerobot.jobs.hf`.
|
||||||
|
Asserting on the completion message and not merely on "didn't hang" is what makes
|
||||||
|
this fail if `follow_job` ever reports detached-without-a-verdict instead.
|
||||||
|
"""
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.HfApi", lambda token=None: MagicMock())
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.ensure_dataset_available", lambda *a, **kw: None)
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.run_job", lambda **kw: MagicMock(id="job-1", url="http://x"))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"lerobot.jobs.hf.inspect_job",
|
||||||
|
lambda job_id: MagicMock(status=MagicMock(stage=MagicMock(value="COMPLETED"), message=None)),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("lerobot.jobs.hf.fetch_job_logs", lambda job_id, follow=True: iter(()))
|
||||||
|
_set_argv(monkeypatch, "--repo_id=u/d", "--job.target=h200")
|
||||||
|
|
||||||
|
submit_annotate_to_hf(_parse("--repo_id", "u/d", "--push_to_hub", "true", "--job.target", "h200"))
|
||||||
|
assert "Annotation complete" in capsys.readouterr().out
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.timeout(15)
|
||||||
|
def test_submit_raises_when_job_fails(monkeypatch):
|
||||||
|
"""A job that ends in a non-COMPLETED stage must surface as an error, not a silent return."""
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.HfApi", lambda token=None: MagicMock())
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.ensure_dataset_available", lambda *a, **kw: None)
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.run_job", lambda **kw: MagicMock(id="job-1", url=None))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"lerobot.jobs.hf.inspect_job",
|
||||||
|
lambda job_id: MagicMock(status=MagicMock(stage=MagicMock(value="ERROR"), message="Job timeout")),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("lerobot.jobs.hf.fetch_job_logs", lambda job_id, follow=True: iter(()))
|
||||||
|
_set_argv(monkeypatch, "--repo_id=u/d", "--job.target=h200")
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="stage=ERROR .Job timeout."):
|
||||||
|
submit_annotate_to_hf(_parse("--repo_id", "u/d", "--job.target", "h200"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_submit_ensures_dataset_is_on_the_hub(monkeypatch):
|
||||||
|
"""A local-only dataset is pushed (privately) before the job can reach it by repo_id."""
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.get_token", lambda: "tok")
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.HfApi", lambda token=None: MagicMock())
|
||||||
|
monkeypatch.setattr("lerobot.jobs.annotate.run_job", lambda **kw: MagicMock(id="job-1"))
|
||||||
|
|
||||||
|
seen = []
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"lerobot.jobs.annotate.ensure_dataset_available",
|
||||||
|
lambda repo_id, *, api, tags=None: seen.append((repo_id, tags)),
|
||||||
|
)
|
||||||
|
_set_argv(monkeypatch, "--repo_id=u/d", "--job.target=h200", "--job.detach=true")
|
||||||
|
|
||||||
|
submit_annotate_to_hf(
|
||||||
|
_parse("--repo_id", "u/d", "--job.target", "h200", "--job.detach", "true", "--job.tags", '["lelab"]')
|
||||||
|
)
|
||||||
|
assert seen == [("u/d", ["lerobot", "lelab"])]
|
||||||
@@ -29,12 +29,26 @@ from lerobot.jobs.hf import (
|
|||||||
_poll_until_done,
|
_poll_until_done,
|
||||||
build_remote_config_file,
|
build_remote_config_file,
|
||||||
build_repo_id,
|
build_repo_id,
|
||||||
|
follow_job,
|
||||||
resolve_job_tags,
|
resolve_job_tags,
|
||||||
resolve_wandb_api_key,
|
resolve_wandb_api_key,
|
||||||
submit_to_hf,
|
submit_to_hf,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_follow_job_detach_returns_without_watching(monkeypatch):
|
||||||
|
"""`detach` must short-circuit before any polling or log streaming starts."""
|
||||||
|
|
||||||
|
def _boom(*a, **kw):
|
||||||
|
raise AssertionError("detach must not touch the job")
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.jobs.hf.inspect_job", _boom)
|
||||||
|
monkeypatch.setattr("lerobot.jobs.hf.fetch_job_logs", _boom)
|
||||||
|
# False = "stopped watching without a verdict", so callers stay quiet rather than
|
||||||
|
# claiming success for a job that is still running.
|
||||||
|
assert follow_job("job-1", detach=True) is False
|
||||||
|
|
||||||
|
|
||||||
def test_resolve_job_tags_always_includes_lerobot_and_dedups():
|
def test_resolve_job_tags_always_includes_lerobot_and_dedups():
|
||||||
assert resolve_job_tags(None) == ["lerobot"]
|
assert resolve_job_tags(None) == ["lerobot"]
|
||||||
assert resolve_job_tags([]) == ["lerobot"]
|
assert resolve_job_tags([]) == ["lerobot"]
|
||||||
|
|||||||
@@ -405,12 +405,18 @@ def test_record_ranges_of_motion(mock_motors, dummy_motors):
|
|||||||
read_pos_stub = mock_motors.build_sequential_sync_read_stub(
|
read_pos_stub = mock_motors.build_sequential_sync_read_stub(
|
||||||
*X_SERIES_CONTROL_TABLE["Present_Position"], positions
|
*X_SERIES_CONTROL_TABLE["Present_Position"], positions
|
||||||
)
|
)
|
||||||
with patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]):
|
bus = DynamixelMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
||||||
bus = DynamixelMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
bus.connect(handshake=False)
|
||||||
bus.connect(handshake=False)
|
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]),
|
||||||
|
patch("lerobot.motors.motors_bus.time.sleep") as mock_sleep,
|
||||||
|
patch.object(bus, "sync_read", wraps=bus.sync_read) as mock_sync_read,
|
||||||
|
):
|
||||||
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
||||||
|
|
||||||
assert mock_motors.stubs[read_pos_stub].calls == 3
|
assert mock_motors.stubs[read_pos_stub].calls == 3
|
||||||
|
assert all(call.kwargs["num_retry"] == 5 for call in mock_sync_read.call_args_list)
|
||||||
|
mock_sleep.assert_called_once_with(0.02)
|
||||||
assert mins == expected_mins
|
assert mins == expected_mins
|
||||||
assert maxes == expected_maxes
|
assert maxes == expected_maxes
|
||||||
|
|||||||
@@ -509,12 +509,18 @@ def test_record_ranges_of_motion(mock_motors, dummy_motors):
|
|||||||
stub = mock_motors.build_sequential_sync_read_stub(
|
stub = mock_motors.build_sequential_sync_read_stub(
|
||||||
*STS_SMS_SERIES_CONTROL_TABLE["Present_Position"], positions
|
*STS_SMS_SERIES_CONTROL_TABLE["Present_Position"], positions
|
||||||
)
|
)
|
||||||
with patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]):
|
bus = FeetechMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
||||||
bus = FeetechMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
bus.connect(handshake=False)
|
||||||
bus.connect(handshake=False)
|
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]),
|
||||||
|
patch("lerobot.motors.motors_bus.time.sleep") as mock_sleep,
|
||||||
|
patch.object(bus, "sync_read", wraps=bus.sync_read) as mock_sync_read,
|
||||||
|
):
|
||||||
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
||||||
|
|
||||||
assert mock_motors.stubs[stub].calls == 3
|
assert mock_motors.stubs[stub].calls == 3
|
||||||
|
assert all(call.kwargs["num_retry"] == 5 for call in mock_sync_read.call_args_list)
|
||||||
|
mock_sleep.assert_called_once_with(0.02)
|
||||||
assert mins == expected_mins
|
assert mins == expected_mins
|
||||||
assert maxes == expected_maxes
|
assert maxes == expected_maxes
|
||||||
|
|||||||
@@ -113,6 +113,7 @@ def test_gaussian_actor_config_default_initialization():
|
|||||||
# Concurrency configuration
|
# Concurrency configuration
|
||||||
assert config.concurrency.actor == "threads"
|
assert config.concurrency.actor == "threads"
|
||||||
assert config.concurrency.learner == "threads"
|
assert config.concurrency.learner == "threads"
|
||||||
|
assert config.concurrency.multiprocessing_context == "spawn"
|
||||||
|
|
||||||
assert isinstance(config.actor_network_kwargs, ActorNetworkConfig)
|
assert isinstance(config.actor_network_kwargs, ActorNetworkConfig)
|
||||||
assert isinstance(config.policy_kwargs, PolicyConfig)
|
assert isinstance(config.policy_kwargs, PolicyConfig)
|
||||||
@@ -152,6 +153,7 @@ def test_concurrency_config():
|
|||||||
config = ConcurrencyConfig()
|
config = ConcurrencyConfig()
|
||||||
assert config.actor == "threads"
|
assert config.actor == "threads"
|
||||||
assert config.learner == "threads"
|
assert config.learner == "threads"
|
||||||
|
assert config.multiprocessing_context == "spawn"
|
||||||
|
|
||||||
|
|
||||||
def test_gaussian_actor_config_custom_initialization():
|
def test_gaussian_actor_config_custom_initialization():
|
||||||
|
|||||||
@@ -26,8 +26,17 @@ import tempfile
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import save_file
|
||||||
|
|
||||||
from lerobot.processor.pipeline import DataProcessorPipeline, ProcessorMigrationError
|
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||||
|
from lerobot.processor.pipeline import (
|
||||||
|
DataProcessorPipeline,
|
||||||
|
ProcessorMigrationError,
|
||||||
|
ProcessorStep,
|
||||||
|
ProcessorStepRegistry,
|
||||||
|
)
|
||||||
|
from lerobot.types import EnvTransition
|
||||||
|
|
||||||
# Simplified Config Loading Tests
|
# Simplified Config Loading Tests
|
||||||
|
|
||||||
@@ -98,6 +107,140 @@ def test_load_config_nonexistent_path_tries_hub():
|
|||||||
DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {})
|
DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {})
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_pretrained_local_directory_missing_state_does_not_call_hub(monkeypatch):
|
||||||
|
"""Local processor dirs must fail locally when a state file is missing."""
|
||||||
|
|
||||||
|
@ProcessorStepRegistry.register("local_missing_state_step")
|
||||||
|
class LocalMissingStateStep(ProcessorStep):
|
||||||
|
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||||
|
return transition
|
||||||
|
|
||||||
|
def transform_features(
|
||||||
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
|
return features
|
||||||
|
|
||||||
|
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
|
tmp_path = Path(tmp_dir)
|
||||||
|
config = {
|
||||||
|
"name": "LocalMissingStatePipeline",
|
||||||
|
"steps": [{"registry_name": "local_missing_state_step", "state_file": "missing.safetensors"}],
|
||||||
|
}
|
||||||
|
(tmp_path / "processor.json").write_text(json.dumps(config))
|
||||||
|
|
||||||
|
def fail_hub_download(*args, **kwargs):
|
||||||
|
pytest.fail("local missing processor state should not call hf_hub_download")
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fail_hub_download)
|
||||||
|
|
||||||
|
with pytest.raises(FileNotFoundError, match="missing.safetensors.*local processor pipeline"):
|
||||||
|
DataProcessorPipeline.from_pretrained(tmp_path, config_filename="processor.json")
|
||||||
|
finally:
|
||||||
|
ProcessorStepRegistry.unregister("local_missing_state_step")
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_pretrained_local_config_file_missing_state_does_not_call_hub(monkeypatch):
|
||||||
|
"""Local single-file processor configs must also keep missing state resolution local."""
|
||||||
|
|
||||||
|
@ProcessorStepRegistry.register("local_file_missing_state_step")
|
||||||
|
class LocalFileMissingStateStep(ProcessorStep):
|
||||||
|
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||||
|
return transition
|
||||||
|
|
||||||
|
def transform_features(
|
||||||
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
|
return features
|
||||||
|
|
||||||
|
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
|
tmp_path = Path(tmp_dir)
|
||||||
|
config_path = tmp_path / "processor.json"
|
||||||
|
config = {
|
||||||
|
"name": "LocalFileMissingStatePipeline",
|
||||||
|
"steps": [
|
||||||
|
{"registry_name": "local_file_missing_state_step", "state_file": "missing.safetensors"}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
config_path.write_text(json.dumps(config))
|
||||||
|
|
||||||
|
def fail_hub_download(*args, **kwargs):
|
||||||
|
pytest.fail("local missing processor state should not call hf_hub_download")
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fail_hub_download)
|
||||||
|
|
||||||
|
with pytest.raises(FileNotFoundError, match="missing.safetensors.*local processor pipeline"):
|
||||||
|
DataProcessorPipeline.from_pretrained(config_path, config_filename="ignored.json")
|
||||||
|
finally:
|
||||||
|
ProcessorStepRegistry.unregister("local_file_missing_state_step")
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_pretrained_hub_source_missing_local_state_still_calls_hub(monkeypatch, tmp_path):
|
||||||
|
"""Hub sources still fall back to hf_hub_download for state files."""
|
||||||
|
|
||||||
|
@ProcessorStepRegistry.register("hub_state_step")
|
||||||
|
class HubStateStep(ProcessorStep):
|
||||||
|
def __init__(self):
|
||||||
|
self.value = torch.tensor(0)
|
||||||
|
|
||||||
|
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||||
|
return transition
|
||||||
|
|
||||||
|
def transform_features(
|
||||||
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
|
return features
|
||||||
|
|
||||||
|
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||||
|
self.value = state["value"]
|
||||||
|
|
||||||
|
try:
|
||||||
|
state_path = tmp_path / "downloaded.safetensors"
|
||||||
|
save_file({"value": torch.tensor(7)}, state_path)
|
||||||
|
loaded_config = {
|
||||||
|
"name": "HubStatePipeline",
|
||||||
|
"steps": [{"registry_name": "hub_state_step", "state_file": "hub_state.safetensors"}],
|
||||||
|
}
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_load_config(cls, model_id, config_filename, hub_download_kwargs):
|
||||||
|
return loaded_config, tmp_path / "hub_cache"
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
return str(state_path)
|
||||||
|
|
||||||
|
monkeypatch.setattr(DataProcessorPipeline, "_load_config", classmethod(fake_load_config))
|
||||||
|
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fake_hub_download)
|
||||||
|
|
||||||
|
pipeline = DataProcessorPipeline.from_pretrained("user/repo", config_filename="processor.json")
|
||||||
|
|
||||||
|
assert calls == [
|
||||||
|
{
|
||||||
|
"repo_id": "user/repo",
|
||||||
|
"filename": "hub_state.safetensors",
|
||||||
|
"repo_type": "model",
|
||||||
|
"force_download": False,
|
||||||
|
"resume_download": None,
|
||||||
|
"proxies": None,
|
||||||
|
"token": None,
|
||||||
|
"cache_dir": None,
|
||||||
|
"local_files_only": False,
|
||||||
|
"revision": None,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
assert pipeline.steps[0].value.item() == 7
|
||||||
|
finally:
|
||||||
|
ProcessorStepRegistry.unregister("hub_state_step")
|
||||||
|
|
||||||
|
|
||||||
# Config Validation Tests
|
# Config Validation Tests
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,9 @@ from lerobot.processor.render_messages_processor import RenderMessagesStep # no
|
|||||||
from lerobot.types import TransitionKey # noqa: E402
|
from lerobot.types import TransitionKey # noqa: E402
|
||||||
|
|
||||||
|
|
||||||
def test_render_messages_step_noops_without_language_columns():
|
def test_render_messages_step_renders_task_fallback_without_language_columns():
|
||||||
|
"""No language columns + a task string → low-level task fallback render,
|
||||||
|
matching what the policy sees at eval time on unannotated observations."""
|
||||||
recipe = TrainingRecipe(
|
recipe = TrainingRecipe(
|
||||||
messages=[
|
messages=[
|
||||||
MessageTurn(role="user", content="${task}", stream="high_level"),
|
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||||
@@ -21,6 +23,24 @@ def test_render_messages_step_noops_without_language_columns():
|
|||||||
)
|
)
|
||||||
transition = create_transition(complementary_data={"task": "do it"})
|
transition = create_transition(complementary_data={"task": "do it"})
|
||||||
|
|
||||||
|
out = RenderMessagesStep(recipe)(transition)
|
||||||
|
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||||
|
|
||||||
|
assert data["messages"] == [{"role": "user", "content": "do it"}]
|
||||||
|
assert data["message_streams"] == ["low_level"]
|
||||||
|
assert data["target_message_indices"] == []
|
||||||
|
assert data["task"] == "do it"
|
||||||
|
|
||||||
|
|
||||||
|
def test_render_messages_step_noops_without_language_columns_or_task():
|
||||||
|
recipe = TrainingRecipe(
|
||||||
|
messages=[
|
||||||
|
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||||
|
MessageTurn(role="assistant", content="${subtask}", stream="low_level", target=True),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
transition = create_transition(complementary_data={})
|
||||||
|
|
||||||
assert RenderMessagesStep(recipe)(transition) == transition
|
assert RenderMessagesStep(recipe)(transition) == transition
|
||||||
|
|
||||||
|
|
||||||
@@ -58,3 +78,70 @@ def test_render_messages_step_renders_and_drops_raw_language():
|
|||||||
assert data["messages"][-1]["content"] == "reach carefully"
|
assert data["messages"][-1]["content"] == "reach carefully"
|
||||||
assert data["message_streams"] == ["high_level", "low_level"]
|
assert data["message_streams"] == ["high_level", "low_level"]
|
||||||
assert data["target_message_indices"] == [1]
|
assert data["target_message_indices"] == [1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_render_messages_step_falls_back_to_low_level_task_when_recipe_misses():
|
||||||
|
recipe = TrainingRecipe(
|
||||||
|
messages=[
|
||||||
|
MessageTurn(
|
||||||
|
role="assistant",
|
||||||
|
content="${subtask}",
|
||||||
|
stream="high_level",
|
||||||
|
target=True,
|
||||||
|
if_present="subtask",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
transition = create_transition(
|
||||||
|
complementary_data={
|
||||||
|
"task": "pick the cube",
|
||||||
|
"timestamp": torch.tensor(0.0),
|
||||||
|
"index": torch.tensor(7),
|
||||||
|
"language_persistent": [],
|
||||||
|
"language_events": [{"style": "unmatched", "timestamp": 0.0}],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
out = RenderMessagesStep(recipe)(transition)
|
||||||
|
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||||
|
|
||||||
|
assert data["messages"] == [{"role": "user", "content": "pick the cube"}]
|
||||||
|
assert data["message_streams"] == ["low_level"]
|
||||||
|
assert data["target_message_indices"] == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_render_messages_step_falls_back_per_sample_in_batched_language():
|
||||||
|
recipe = TrainingRecipe(
|
||||||
|
messages=[
|
||||||
|
MessageTurn(
|
||||||
|
role="assistant",
|
||||||
|
content="${subtask}",
|
||||||
|
stream="high_level",
|
||||||
|
target=True,
|
||||||
|
if_present="subtask",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
transition = create_transition(
|
||||||
|
action=torch.arange(4).reshape(2, 2),
|
||||||
|
complementary_data={
|
||||||
|
"task": ["pick the cube", "open the drawer"],
|
||||||
|
"timestamp": torch.tensor([0.0, 1.0]),
|
||||||
|
"index": torch.tensor([7, 8]),
|
||||||
|
"language_persistent": [[], []],
|
||||||
|
"language_events": [
|
||||||
|
[{"style": "unmatched", "timestamp": 0.0}],
|
||||||
|
[{"style": "unmatched", "timestamp": 1.0}],
|
||||||
|
],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
out = RenderMessagesStep(recipe)(transition)
|
||||||
|
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||||
|
|
||||||
|
assert data["messages"] == [
|
||||||
|
[{"role": "user", "content": "pick the cube"}],
|
||||||
|
[{"role": "user", "content": "open the drawer"}],
|
||||||
|
]
|
||||||
|
assert data["message_streams"] == [["low_level"], ["low_level"]]
|
||||||
|
assert data["target_message_indices"] == [[], []]
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ import pytest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
|
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
|
||||||
from lerobot.processor import DataProcessorPipeline, TokenizerProcessorStep
|
from lerobot.processor import ActionTokenizerProcessorStep, DataProcessorPipeline, TokenizerProcessorStep
|
||||||
from lerobot.processor.converters import create_transition, identity_transition
|
from lerobot.processor.converters import create_transition, identity_transition
|
||||||
from lerobot.types import TransitionKey
|
from lerobot.types import TransitionKey
|
||||||
from lerobot.utils.constants import (
|
from lerobot.utils.constants import (
|
||||||
@@ -88,6 +88,46 @@ class MockTokenizer:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def test_action_tokenizer_config_preserves_token_mapping():
|
||||||
|
processor = object.__new__(ActionTokenizerProcessorStep)
|
||||||
|
processor.trust_remote_code = True
|
||||||
|
processor.max_action_tokens = 384
|
||||||
|
processor.fast_skip_tokens = 64
|
||||||
|
processor.paligemma_tokenizer_name = "custom/paligemma"
|
||||||
|
processor.allow_truncation = False
|
||||||
|
processor.action_tokenizer_name = "custom/fast"
|
||||||
|
processor.action_tokenizer_input_object = None
|
||||||
|
|
||||||
|
assert processor.get_config() == {
|
||||||
|
"trust_remote_code": True,
|
||||||
|
"max_action_tokens": 384,
|
||||||
|
"fast_skip_tokens": 64,
|
||||||
|
"paligemma_tokenizer_name": "custom/paligemma",
|
||||||
|
"allow_truncation": False,
|
||||||
|
"action_tokenizer_name": "custom/fast",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_action_tokenizer_can_reject_truncated_sequences():
|
||||||
|
processor = object.__new__(ActionTokenizerProcessorStep)
|
||||||
|
processor.max_action_tokens = 4
|
||||||
|
processor.fast_skip_tokens = 128
|
||||||
|
processor.allow_truncation = False
|
||||||
|
processor.action_tokenizer = lambda _actions: [1, 2, 3]
|
||||||
|
processor._paligemma_tokenizer = type(
|
||||||
|
"Tokenizer",
|
||||||
|
(),
|
||||||
|
{
|
||||||
|
"vocab_size": 1000,
|
||||||
|
"bos_token_id": 2,
|
||||||
|
"encode": lambda _self, text, **_kwargs: [10, 11] if text == "Action: " else [12, 1],
|
||||||
|
},
|
||||||
|
)()
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="max_action_tokens=4"):
|
||||||
|
processor._tokenize_action(torch.zeros(1, 2, 1))
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_tokenizer():
|
def mock_tokenizer():
|
||||||
"""Provide a mock tokenizer for testing."""
|
"""Provide a mock tokenizer for testing."""
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
import lerobot.scripts.lerobot_setup_motors as motors_module
|
||||||
|
|
||||||
|
|
||||||
|
def test_main_registers_plugins_before_parsing(monkeypatch):
|
||||||
|
calls = []
|
||||||
|
monkeypatch.setattr(motors_module, "register_third_party_plugins", lambda: calls.append("register"))
|
||||||
|
monkeypatch.setattr(motors_module, "setup_motors", lambda: calls.append("setup"))
|
||||||
|
|
||||||
|
motors_module.main()
|
||||||
|
|
||||||
|
assert calls == ["register", "setup"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_setup_motors_accepts_third_party_device(monkeypatch):
|
||||||
|
device = MagicMock()
|
||||||
|
monkeypatch.setattr(motors_module, "make_teleoperator_from_config", lambda _: device)
|
||||||
|
cfg = SimpleNamespace(device=SimpleNamespace(type="third_party"))
|
||||||
|
|
||||||
|
motors_module.setup_motors.__wrapped__(cfg)
|
||||||
|
|
||||||
|
device.setup_motors.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
|
def test_setup_motors_reports_unsupported_device(monkeypatch):
|
||||||
|
device = object()
|
||||||
|
monkeypatch.setattr(motors_module, "make_teleoperator_from_config", lambda _: device)
|
||||||
|
cfg = SimpleNamespace(device=SimpleNamespace(type="third_party"))
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError, match="third_party"):
|
||||||
|
motors_module.setup_motors.__wrapped__(cfg)
|
||||||
@@ -17,6 +17,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import sys
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -106,6 +108,109 @@ def test_sentry_config_defaults():
|
|||||||
assert cfg.target_video_file_size_mb is None
|
assert cfg.target_video_file_size_mb is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_config_passes_policy_pretrained_revision(monkeypatch):
|
||||||
|
from lerobot.configs import PreTrainedConfig, parser
|
||||||
|
from lerobot.rollout import RolloutConfig
|
||||||
|
from tests.mocks.mock_robot import MockRobotConfig
|
||||||
|
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
def fake_from_pretrained(cls, pretrained_name_or_path, **kwargs):
|
||||||
|
captured["pretrained_name_or_path"] = pretrained_name_or_path
|
||||||
|
captured.update(kwargs)
|
||||||
|
return SimpleNamespace(device="cpu", pretrained_revision=kwargs["revision"])
|
||||||
|
|
||||||
|
monkeypatch.setattr(parser, "get_yaml_overrides", lambda _: ["--pretrained_revision=yaml-sha"])
|
||||||
|
monkeypatch.setattr(
|
||||||
|
sys,
|
||||||
|
"argv",
|
||||||
|
["lerobot-rollout", "--policy.path=user/policy", "--policy.pretrained_revision=cli-sha"],
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(PreTrainedConfig, "from_pretrained", classmethod(fake_from_pretrained))
|
||||||
|
|
||||||
|
cfg = RolloutConfig(robot=MockRobotConfig())
|
||||||
|
|
||||||
|
assert captured["pretrained_name_or_path"] == "user/policy"
|
||||||
|
assert captured["revision"] == "cli-sha"
|
||||||
|
assert captured["cli_overrides"] == [
|
||||||
|
"--pretrained_revision=yaml-sha",
|
||||||
|
"--pretrained_revision=cli-sha",
|
||||||
|
]
|
||||||
|
assert cfg.policy.pretrained_path == "user/policy"
|
||||||
|
assert cfg.policy.pretrained_revision == "cli-sha"
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_pretrained_policy_passes_revision(monkeypatch):
|
||||||
|
import lerobot.rollout.context as rollout_context
|
||||||
|
|
||||||
|
policy_config = SimpleNamespace(
|
||||||
|
type="mock",
|
||||||
|
use_peft=False,
|
||||||
|
pretrained_path="user/policy",
|
||||||
|
pretrained_revision="policy-sha",
|
||||||
|
)
|
||||||
|
policy_class = MagicMock()
|
||||||
|
loaded_policy = MagicMock()
|
||||||
|
policy_class.from_pretrained.return_value = loaded_policy
|
||||||
|
monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class)
|
||||||
|
|
||||||
|
policy = rollout_context._load_pretrained_policy(policy_config)
|
||||||
|
|
||||||
|
assert policy is loaded_policy
|
||||||
|
policy_class.from_pretrained.assert_called_once_with(
|
||||||
|
"user/policy",
|
||||||
|
config=policy_config,
|
||||||
|
revision="policy-sha",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_pretrained_peft_policy_keeps_adapter_and_base_revisions_separate(monkeypatch):
|
||||||
|
import lerobot.rollout.context as rollout_context
|
||||||
|
|
||||||
|
policy_config = SimpleNamespace(
|
||||||
|
type="mock",
|
||||||
|
use_peft=True,
|
||||||
|
pretrained_path="user/adapter",
|
||||||
|
pretrained_revision="adapter-sha",
|
||||||
|
)
|
||||||
|
policy_class = MagicMock()
|
||||||
|
base_policy = MagicMock()
|
||||||
|
policy_class.from_pretrained.return_value = base_policy
|
||||||
|
monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class)
|
||||||
|
|
||||||
|
peft_config = SimpleNamespace(
|
||||||
|
base_model_name_or_path="user/base-policy",
|
||||||
|
revision="base-sha",
|
||||||
|
)
|
||||||
|
peft_config_from_pretrained = MagicMock(return_value=peft_config)
|
||||||
|
adapted_policy = MagicMock()
|
||||||
|
peft_model_from_pretrained = MagicMock(return_value=adapted_policy)
|
||||||
|
monkeypatch.setitem(
|
||||||
|
sys.modules,
|
||||||
|
"peft",
|
||||||
|
SimpleNamespace(
|
||||||
|
PeftConfig=SimpleNamespace(from_pretrained=peft_config_from_pretrained),
|
||||||
|
PeftModel=SimpleNamespace(from_pretrained=peft_model_from_pretrained),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
policy = rollout_context._load_pretrained_policy(policy_config)
|
||||||
|
|
||||||
|
assert policy is adapted_policy
|
||||||
|
peft_config_from_pretrained.assert_called_once_with("user/adapter", revision="adapter-sha")
|
||||||
|
policy_class.from_pretrained.assert_called_once_with(
|
||||||
|
pretrained_name_or_path="user/base-policy",
|
||||||
|
config=policy_config,
|
||||||
|
revision="base-sha",
|
||||||
|
)
|
||||||
|
peft_model_from_pretrained.assert_called_once_with(
|
||||||
|
base_policy,
|
||||||
|
"user/adapter",
|
||||||
|
config=peft_config,
|
||||||
|
revision="adapter-sha",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# RolloutRingBuffer
|
# RolloutRingBuffer
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ import torch # noqa: E402
|
|||||||
from lerobot.utils.collate import lerobot_collate_fn # noqa: E402
|
from lerobot.utils.collate import lerobot_collate_fn # noqa: E402
|
||||||
|
|
||||||
|
|
||||||
def test_lerobot_collate_preserves_messages_and_drops_raw_language():
|
def test_lerobot_collate_preserves_messages_and_raw_language():
|
||||||
batch = [
|
batch = [
|
||||||
{
|
{
|
||||||
"index": torch.tensor(0),
|
"index": torch.tensor(0),
|
||||||
@@ -17,14 +17,14 @@ def test_lerobot_collate_preserves_messages_and_drops_raw_language():
|
|||||||
"message_streams": ["low_level"],
|
"message_streams": ["low_level"],
|
||||||
"target_message_indices": [0],
|
"target_message_indices": [0],
|
||||||
"language_persistent": [{"content": "raw"}],
|
"language_persistent": [{"content": "raw"}],
|
||||||
"language_events": [],
|
"language_events": [{"content": "event a"}],
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"index": torch.tensor(1),
|
"index": torch.tensor(1),
|
||||||
"messages": [{"role": "assistant", "content": "b"}],
|
"messages": [{"role": "assistant", "content": "b"}],
|
||||||
"message_streams": ["low_level"],
|
"message_streams": ["low_level"],
|
||||||
"target_message_indices": [0],
|
"target_message_indices": [0],
|
||||||
"language_persistent": [{"content": "raw"}],
|
"language_persistent": [{"content": "raw b"}],
|
||||||
"language_events": [],
|
"language_events": [],
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
@@ -36,8 +36,8 @@ def test_lerobot_collate_preserves_messages_and_drops_raw_language():
|
|||||||
assert out["messages"][1][0]["content"] == "b"
|
assert out["messages"][1][0]["content"] == "b"
|
||||||
assert out["message_streams"] == [["low_level"], ["low_level"]]
|
assert out["message_streams"] == [["low_level"], ["low_level"]]
|
||||||
assert out["target_message_indices"] == [[0], [0]]
|
assert out["target_message_indices"] == [[0], [0]]
|
||||||
assert "language_persistent" not in out
|
assert out["language_persistent"] == [[{"content": "raw"}], [{"content": "raw b"}]]
|
||||||
assert "language_events" not in out
|
assert out["language_events"] == [[{"content": "event a"}], []]
|
||||||
|
|
||||||
|
|
||||||
def test_lerobot_collate_passes_through_standard_batch():
|
def test_lerobot_collate_passes_through_standard_batch():
|
||||||
|
|||||||
Reference in New Issue
Block a user