Compare commits

...

32 Commits

Author SHA1 Message Date
Pepijn c266fe4f1a Merge branch 'main' into codex/language-recipes 2026-07-28 12:50:11 +02:00
Pepijn 913664c320 test(collate): expect preserved language columns 2026-07-28 12:39:15 +02:00
Steven Palma c1b6ea85d6 feat(rl): add multiprocessing option to training pipeline and sets spawn as default + guard (#4140) 2026-07-28 12:19:53 +02:00
Steven Palma ffe25afb8f fix(processors): wrong feature key dropped in delta-action transform_features (#4165) 2026-07-28 11:18:23 +02:00
Pepijn 3f093d8927 feat(data): add recipe-driven language supervision 2026-07-28 10:04:28 +02:00
Steven Palma 95211b98f1 feat(config): add multiprocessing option to DataLoader context and sets spawn as default (#4139)
* Add dataloader_multiprocessing_context, default to spawn

Make the DataLoader multiprocessing start method configurable on
TrainPipelineConfig and default it to 'spawn'.

The previous default (fork on Linux) is unsafe with libraries that hold
non-fork-safe state in the parent process — common ones in this codebase
are PyAV, torchcodec, and the ffmpeg shared libs they wrap. Symptoms
reported in #2488, #2209, and observed locally include:

- multiprocessing.context.AuthenticationError: digest received was wrong
- RuntimeError: Pin memory thread exited unexpectedly
- RuntimeError: DataLoader worker exited unexpectedly
- Random SIGSEGV inside worker processes during video decode

Switching to spawn re-imports modules cleanly in each worker and
eliminates these failure modes. Added the setting as a config field
rather than hard-coding so users on platforms where fork is preferred
can opt back in via --dataloader-multiprocessing-context=fork.

* Address review: shorten config comment, note spawn startup tradeoff

Per @jashshah999, mention that spawn workers re-import modules and so
add some startup time vs fork. Also trim the failure-mode dump from
the inline comment — the linked issue covers the symptoms in detail.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>

* chore(scripts): add multiprocessing_context safeguards

* chore(config): add libs note

---------

Co-authored-by: 0o8o0-blip <0o8o0-blip@users.noreply.github.com>
2026-07-28 00:42:55 +02:00
Xingdong Zuo 95256d766d feat(lekiwi): support LeKiwi in lerobot-replay CLI (#3739)
Register the `lekiwi` robot module in `lerobot_replay.py` so episodes can be
replayed on a LeKiwi via `--robot.type=lekiwi_client`. The module is already
registered in `lerobot_calibrate.py` and `lerobot_setup_motors.py`; this fills
the gap so the replay CLI recognizes the same robot.

Replayed actions are loaded from the dataset as torch tensors, which
`json.dumps` cannot serialize when `LeKiwiClient.send_action` ships them over
ZMQ. Coerce each action value to a plain float before sending. This is scoped
to the LeKiwi network client and does not affect any other robot.

Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-27 19:29:44 +02:00
Thomas Landeg fd53716688 fix(envs): make metaworld seeding reproducible (#3727)
Co-authored-by: Pepijn <138571049+pkooij@users.noreply.github.com>
2026-07-27 18:42:17 +02:00
Steven Palma a96540a2c4 fix rollout policy revision loading (#4161)
Co-authored-by: RaviTeja-Kondeti <rkondet3@asu.edu>
2026-07-27 18:20:39 +02:00
WOLIKIMCHENG acd42b4d85 fix(processor): keep missing local state resolution local (#3715)
Co-authored-by: root <kinsonnee@gmail.com>
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-27 15:53:31 +02:00
Steven Palma bbeacfe57d fix(record): connect teleoperator before robot to avoid watchdog jump (#4166)
* fix(record): connect teleoperator before robot to avoid watchdog jump

lerobot-record connected the robot before the teleoperator. A robot's
connect()/reset() can leave it holding a default pose under a firmware
watchdog (e.g. Unitree G1); if teleop.connect() (model loading, IK init,
network setup) then takes longer than that watchdog, the joints drop to
damping and the first send_action() makes the robot jump.

Swap the order so the teleoperator connects first, matching the ordering
already used in lerobot_teleoperate.py. Pure ordering fix, no API change.

Fixes #3684

* fix(record): trim comment and add connect-order regression test

Address review feedback on #3684:
- Trim the verbose ordering comment down to two lines.
- Add test_record_connects_teleop_before_robot to tests/test_control_robot.py,
  asserting teleop.connect() runs before robot.connect() in record().

* chore(test): remove test

---------

Co-authored-by: Jaimin Patel <jpatel@tuvalabs.com>
Co-authored-by: Martino Russi <77496684+nepyope@users.noreply.github.com>
2026-07-27 14:09:58 +02:00
Steven Palma 801346e18c fix(scripts): restore policy training mode after eval_policy() in lerobot-eval (#4162)
* fix(scripts/eval): restore policy training mode after eval_policy()

`eval_policy` calls `policy.eval()` before the rollout but never restores
the prior mode on return. When called from the training loop
(`lerobot_train.py`'s `eval_policy_all -> run_one -> eval_one ->
eval_policy` chain), the policy is left in eval mode for every subsequent
training step, which silently:

  * disables Dropout (no regularisation),
  * freezes BatchNorm running stats (no further EMA updates).

Under DDP only `is_main_process` runs eval (lerobot_train.py:527), so the
main rank ends up in eval mode while workers stay in train mode — the
all-reduced gradients then combine forward passes computed with different
dropout masks and different BN behaviour, a real DDP-correctness issue.

Scope of impact:
  * Affects every policy with Dropout in its forward path. In-tree, that
    includes the default ACT (6 Dropout layers at p=0.1), Diffusion (vision
    backbone), VQ-BeT, Multi-Task DiT, X-VLA, plus all VLA policies that
    inherit Dropout from their pretrained HF backbone (PI0/PI0.5/PI0-FAST,
    SmolVLA, GR00T-N1.5, EO1, Wall-X).
  * Triggers from the first eval onward. On the default config
    (steps=100k, eval_freq=20k) that's 80% of training; on the LIBERO /
    RoboCasa / VLABench example commands in docs/ (eval_freq=1k–5k)
    it's 95–99% of training.
  * Policies using only LayerNorm/GroupNorm and no Dropout (TDMPC, RTC)
    are unaffected. Policies using `FrozenBatchNorm2d` (ACT's ResNet
    backbone) are immune to the BN-stat half; the Dropout half still bites.

Fix:
  * Snapshot `policy.training` on entry to `eval_policy`.
  * Restore it on normal return.
  * Save-and-restore is a strict no-op for callers that pass an
    already-eval-mode policy (e.g. the standalone `lerobot-eval` script
    loading a frozen checkpoint).
  * Restoration is placed before the normal return only, not in a
    try/finally — exception paths leave the policy in eval mode, same as
    today. A try/finally upgrade would require re-indenting ~165 lines and
    can land as a separate cleanup if desired.

Tests (tests/scripts/test_eval.py, 7 tests total, ~1.6s):
  * Regression gates on the lerobot_eval fix itself: training-mode
    preservation, eval-mode preservation, dropout-active behavioural
    check, non-crash for both entry modes.
  * Quantitative mechanism demonstration
    (`test_missing_mode_restoration_hurts_generalisation`): trains a tiny
    Dropout+BatchNorm MLP under both the bug pattern and the fix pattern
    on identical data and seed, then asserts the buggy variant generalises
    at least 5% worse on a held-out val set. In repeated runs we see
    10-25% deltas on this toy problem; real policies (more layers, more
    Dropout, longer training) generally see larger gaps. Lives alongside
    the regression tests so the empirical proof is reproducible from the
    repo without adding a separate benchmarks/ directory.


* fix(scripts): keep policy train/eval

---------

Co-authored-by: ModeEric <ericjm4@illinois.edu>
2026-07-27 14:08:11 +02:00
MihaiAnca13 ab87fd9764 fix(datasets): clear video frame staging on episode reset (#3683)
* fix video frame staging cleanup on episode reset

* linting

---------

Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-27 13:44:32 +02:00
hf-dependantbot-rollout[bot] 6c57dfd2ee chore: enable Dependabot weekly GitHub Actions bumps (#3677)
Co-authored-by: hf-dependantbot-rollout[bot] <285970069+hf-dependantbot-rollout[bot]@users.noreply.github.com>
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-27 13:22:04 +02:00
Kohei SENDAI d63e6e67a5 fix convverstion err (#3656)
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-27 11:54:09 +02:00
Steven Palma 0d383d09f2 feat(dataset): accept token argument for private HF Hub datasets (#4136) 2026-07-24 18:51:35 +02:00
Caroline Pascal ab2b5b04dd (depth image processing): excluding depth frames from the RGB to BGR image processing (#4135)
* (depth image processing): excluding depth frames from the RGB to BGR image processing

* test(update): updating tests to include RGB/BGR conversion checks
2026-07-24 17:43:17 +02:00
Steven Palma ac5c7b8600 chore(deps): bump diffusers to >=0.38.0,<0.40.0 (#4145)
* fix(deps): bump diffusers cap to <0.39.0 (security)

Diffusers 0.35.x is affected by GHSA-98h9-4798-4q5v (HIGH, CVSS 8.8):
'trust_remote_code bypass via custom_pipeline and local custom components'.
Fixed in diffusers 0.38.0.

Current cap 'diffusers<0.36.0' blocks downstream consumers (e.g.
strands-labs/robots) from picking up the security fix.

The lerobot diffusers surface area is narrow and stable across 0.36-0.38:
- diffusers.schedulers.scheduling_ddim.DDIMScheduler
- diffusers.schedulers.scheduling_ddpm.DDPMScheduler
- diffusers.optimization.get_scheduler
- diffusers.ConfigMixin / ModelMixin / register_to_config
- diffusers.models.attention.{Attention,FeedForward}
- diffusers.models.embeddings.*

None of these were removed, renamed, or had breaking changes in 0.36, 0.37,
or 0.38 release notes. Bumping the cap to <0.39.0 unblocks the security
fix while keeping a major-version safety bound.

* chore(dependecies): bump diffusers

* chore(deps): update uv.lock

---------

Co-authored-by: Cagatay Cali <cagataycali@users.noreply.github.com>
2026-07-24 17:13:53 +02:00
Steven Palma a6befef0ba chore(dependencies): update uv.lock (#3963) 2026-07-24 16:30:36 +02:00
Steven Palma 53843007ea feat(robot): Make SO follower P coefficient configurable (#4142)
* Make SO follower P coefficient configurable

* chore(test): minimize tests

* feat(robots): expose PID coeff in SO arms

---------

Co-authored-by: taivu1998 <46636857+taivu1998@users.noreply.github.com>
2026-07-24 16:03:04 +02:00
Maxime Ellerbach d3bed0feee chore(agents): adding additional infos to AGENTS.md and bring-your-own-policies.mdx (#3904)
* chore(agents): adding additional infos to AGENTS.md

* adding `lerobot-train` requirement inside PR checklist

* prefer using code already implemented from transformers / diffusers instead of re-implementing in tree

---------

Signed-off-by: Maxime Ellerbach <maxime.ellerbach@huggingface.co>
2026-07-24 14:58:43 +02:00
Steven Palma a0eb860d1e feat(dataset): add slice support to LeRobotDataset.__getitem__ (#4129)
* feat(dataset): add efficient slice support

* fix(dataset): handle empty dataset slices

* refactor(dataset): reuse scalar path for slices

---------

Co-authored-by: Francesco Capuano <fc.francescocapuano@gmail.com>
2026-07-23 22:05:29 +02:00
Steven Palma cfd9ff969c fix(envs): set LiberoEnvConfig.fps default to 20 to match robosuite (#4124)
* fix(envs): set LiberoEnvConfig.fps default to 20 to match robosuite

LiberoEnvConfig.fps was set to 30, but the underlying robosuite
OffScreenRenderEnv always runs at its default control_freq of 20 Hz
since fps is never passed through. This mismatch silently decouples
the dataset/eval loop rate from the actual simulation step rate.

Set the default to 20 to match the real sim rate and avoid the
footgun.

Fixes #3368

* fix(libero): apply configured control frequency

---------

Co-authored-by: xinmotlanthua <275663218+xinmotlanthua@users.noreply.github.com>
2026-07-23 19:49:19 +02:00
Steven Palma f59eae4e27 fix(robots): add retries while recording motor ranges (#4126)
* Add retries while recording motor ranges

* fix(motors): throttle calibration reads consistently

---------

Co-authored-by: tom-doerr <tomdoerr96@gmail.com>
2026-07-23 18:41:48 +02:00
Martino Russi a993af9c51 fix(openarms): stop set_zero_position()ing on connect (#4058)
* fix(damiao): make is_calibrated a plain property, not cached

`is_calibrated` was a `@cached_property`, so it froze at its first-read
value and never reflected later changes to `self.calibration` (set by
connect/calibrate/load). This caused the OpenArm teleop to re-run
calibration even when a calibration file existed, and to skip
`set_zero_position()` after a fresh calibration.

Switch to `@property` (matching the MotorsBus base contract and the
Feetech/SO-100 buses) and drop the now-unused `functools.cached_property`
import.

Co-authored-by: Cursor <cursoragent@cursor.com>

* don't set_zero_position() on connect

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-07-23 18:34:13 +02:00
Steven Palma 392246feaf feat(diffusion): add gradient checkpointing for memory optimization (#4127)
* feat(diffusion): add gradient checkpointing for memory optimization

Add gradient_checkpointing config option to DiffusionPolicy. When
enabled, wraps the UNet encoder, mid, and decoder residual blocks
with torch.utils.checkpoint.checkpoint to trade compute for memory.

Allows training with larger batch sizes or higher-resolution inputs
on memory-constrained GPUs. Disabled by default.

Usage: --policy.gradient_checkpointing=true

Part of the 0.6.0 roadmap item 3.3 (gradient checkpointing for all
policies).

* test(diffusion): verify gradient checkpointing parity

---------

Co-authored-by: Jash Shah <jashshah.999@gmail.com>
2026-07-23 18:33:10 +02:00
Steven Palma 19dcbc19f1 fix(gamepad): Gamepad on macos often does not need fallback (#4125)
* gamepad does often work on macos

* review comments

* fix(gamepad): expose hidapi fallback in config

---------

Co-authored-by: Maxim Bonnaerens <maxim@bonnaerens.be>
2026-07-23 18:21:48 +02:00
Steven Palma 679faeaafc fix(scripts): register third-party plugins in lerobot_setup_motors (#4123)
* fix(scripts): register third-party plugins in setup-motors

* test(setup-motors): cover plugin registration

---------

Co-authored-by: Janos von Gencsy <janos.von-gencsy@tum.de>
2026-07-23 18:06:44 +02:00
YK 228cb5ddb9 Fix missing periods at end of sentences in README (#3473)
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-23 16:07:58 +02:00
Eunsung Kim ad176c6d41 Feature omx docs (#3421)
* docs(omx): add header and omx image in docs

* fix(docs):adjust image size in omx docs

---------

Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-23 15:25:46 +02:00
Duhyeon, Kim d6c605e8c5 refactor(pi05): remove unused variables in embed_suffix method (#3263)
* refactor(pi05): remove unused variables in embed_suffix method

* Refactor embed_suffix to streamline pad_masks handling

Removed unused pad_masks list and simplified its creation.

Signed-off-by: Duhyeon, Kim <49020301+dudududukim@users.noreply.github.com>

---------

Signed-off-by: Duhyeon, Kim <49020301+dudududukim@users.noreply.github.com>
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
2026-07-23 14:37:57 +02:00
Pepijn 9c82c39c7b feat(annotate): run lerobot-annotate on HF Jobs via --job.target (#4095)
* feat(annotate): run lerobot-annotate on HF Jobs via --job.target

Annotation needed a hand-edited launcher script (examples/annotations/run_hf_job.py)
to reach a GPU: users copied it, rewrote the embedded CMD string for their dataset,
and ran it with `python`. Fold that into the CLI instead, mirroring `lerobot-train`:
`lerobot-annotate --job.target=h200` submits the exact command you'd run locally.

- AnnotationJobConfig extends JobConfig with the annotation runtime's defaults
  (vllm/vllm-openai image, 2h cap) plus --job.lerobot_ref, so an unmerged branch
  can be exercised remotely without editing a script.
- lerobot.jobs.annotate builds the pod command by replaying the user's own CLI
  flags (minus --job.*/--root, with --repo_id re-emitted from the config) after a
  setup prelude that installs lerobot on top of the vLLM image. Job monitoring,
  log tailing and Ctrl-C-detaches reuse the training submitter's plumbing.
- Remote runs require --repo_id; a local-only dataset is pushed privately first.

The generated pod command is byte-for-byte the script's old CMD.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* fix(annotate): reject client-side config files on remote runs

draccus exposes `--config_path` plus a `--<field>` config-file arg for every
nested dataclass (`--vlm`, `--plan`, `--job`, ...). All name files on the
client's disk, so forwarding them to the pod silently dropped whatever settings
they carried. Reject them up front instead.

Bare `--job` also slipped past the `--job.` prefix filter, so a `--job=cfg.yaml`
holding `target: h200` would have reached the pod and had the job submit a job
of its own, recursively. It is dropped from the forwarded args as well.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* refactor(jobs): share the submit-and-follow loop between both submitters

`submit_annotate_to_hf` reused the leaf helpers (`_poll_until_done`, `_tail_logs`,
`_pod_forwarded_args`) but duplicated the orchestration around them: ~40 of the 50
lines that spawn the poll/log threads, install the Ctrl-C-detaches handler and
raise on a non-COMPLETED stage were identical in both files.

Extract that into `follow_job(job_id, *, detach, success_marker=None) -> bool`,
returning True when the job finished and False when we stopped watching without a
verdict (detach or Ctrl-C). Training keeps its model-pushed marker by passing it in;
annotation has no equivalent line (the CLI keeps working after the upload log to
write the card and tag) so its completion stays stage-based.

Kept in hf.py rather than a new module so every existing monkeypatch target in
test_hf.py still resolves.

Behaviour change: a training run whose job reaches COMPLETED without the marker
matching now prints its completion line instead of returning silently. The marker
was already documented as an optimisation with a stage-based fallback; the fallback
just never reported success.

Tests: adds annotate coverage for the non-detach path (completion and failure) —
previously only ever exercised with detach=true — plus a detach short-circuit test.
Both new annotate tests verified to fail under a mutation that stubs out follow_job.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-23 10:30:33 +02:00
96 changed files with 3484 additions and 1413 deletions
+11
View File
@@ -0,0 +1,11 @@
version: 2
updates:
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
cooldown:
default-days: 7
groups:
actions:
patterns: ["*"]
+2 -1
View File
@@ -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`).
+3 -3
View File
@@ -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
+55 -16
View File
@@ -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
+6 -1
View File
@@ -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).
+11
View File
@@ -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.
+8
View File
@@ -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.
+4
View File
@@ -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`:
-80
View File
@@ -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
View File
@@ -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]:
+6
View File
@@ -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
+10 -4
View File
@@ -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."""
+16
View File
@@ -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}
+18
View File
@@ -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:
+16 -2
View File
@@ -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
+30
View File
@@ -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:
+23 -14
View File
@@ -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()
+16 -2
View File
@@ -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,
+73 -10
View File
@@ -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(
+39 -10
View File
@@ -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
+3
View File
@@ -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
] ]
+14 -1
View File
@@ -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)
+13 -4
View File
@@ -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(
+5 -1
View File
@@ -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
+5
View File
@@ -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
+3
View File
@@ -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)
+2 -1
View File
@@ -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"]
+176
View File
@@ -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
View File
@@ -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}"
)
+1 -2
View File
@@ -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)
+9 -5
View File
@@ -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)
+1
View File
@@ -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,
+4 -12
View File
@@ -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."""
-3
View File
@@ -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,)
) )
+101 -6
View File
@@ -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": [],
}
+50 -17
View File
@@ -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]]:
+2 -4
View File
@@ -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
+2 -4
View File
@@ -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
+11 -2
View File
@@ -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")
+32 -13
View File
@@ -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(
+16 -1
View File
@@ -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)
+41 -28
View File
@@ -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):
+3 -1
View File
@@ -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()
+1
View File
@@ -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,
+6 -17
View File
@@ -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()
+12 -4
View File
@@ -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 doesnt 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
+1 -1
View File
@@ -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:
+2
View File
@@ -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"
+28
View File
@@ -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.
+23 -1
View File
@@ -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)
+44 -15
View File
@@ -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")
+27 -1
View File
@@ -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)
+7
View File
@@ -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.
+30 -1
View File
@@ -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)
+32
View File
@@ -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(
+46
View File
@@ -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)
+78
View File
@@ -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
+43
View File
@@ -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,
+1 -3
View File
@@ -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
+38
View File
@@ -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
+11
View File
@@ -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")
+245
View File
@@ -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"])]
+14
View File
@@ -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"]
+9 -3
View File
@@ -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
+9 -3
View File
@@ -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"] == [[], []]
+41 -1
View File
@@ -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."""
+19
View File
@@ -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)
+49
View File
@@ -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)
+105
View File
@@ -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
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+5 -5
View File
@@ -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():
Generated
+1029 -1011
View File
File diff suppressed because it is too large Load Diff