mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 04:36:04 +00:00
Compare commits
10 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 68d6335d5b | |||
| 258d521a89 | |||
| 95211b98f1 | |||
| 95256d766d | |||
| fd53716688 | |||
| a96540a2c4 | |||
| acd42b4d85 | |||
| bbeacfe57d | |||
| 801346e18c | |||
| ab87fd9764 |
+113
@@ -0,0 +1,113 @@
|
|||||||
|
# LeRobot Docs Audit
|
||||||
|
|
||||||
|
**Status:** Fresh baseline for a new docs redesign effort · **Date:** 2026-07-27
|
||||||
|
**Supersedes:** the June 2026 `DOCS_REDESIGN.md` proposal (lived on the unmerged, now-stale `docs/complete-docs-redesign` branch — abandoned; that branch is ~67k lines behind `main` and should not be resurrected). This document re-audits `docs/source/` from scratch against current `main`, since the docs tree grew substantially in the ~6 weeks between the two audits (many new policies, robots, and benchmarks landed).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Snapshot (current state, in numbers)
|
||||||
|
|
||||||
|
- **IA**: 15 top-level `_toctree.yml` sections, 81 `local:` entries, 0 broken toctree links.
|
||||||
|
- **Files**: 97 total files in `docs/source/` = 80 real `.mdx` pages + 16 orphaned `policy_*_README.md` stubs (up from 14 in June) + 1 `contributing.md` (a genuine symlink to root `CONTRIBUTING.md`).
|
||||||
|
- **Growth**: `.mdx` pages grew 49 → 80 in 6 months (+63%, ~5.2 pages/month), arriving in **bursts against a shared, hand-edited `_toctree.yml`** — e.g. 6 benchmark pages landed via 6 separate PRs within 48 hours (Apr 2026); 3 policy pages within 3 days (Jul 2026).
|
||||||
|
- **CLI docs coverage improved sharply since June**: 17 of 18 `lerobot-*` commands now documented (only `lerobot-info`, a diagnostics helper, is not) — the June baseline was 6/17 undocumented.
|
||||||
|
- **Decision content**: zero comparison/decision pages exist in `docs/source/` for Policies (15 options), Robots (11), Benchmarks (10), or Teleoperators — this equivalent content already exists, fully written, in root `AGENT_GUIDE.md` §6, but was never ported or linked in.
|
||||||
|
- **Teleoperators**: dedicated toctree section has only 2 pages vs. 16-17 registered teleoperator types in `src/`, but all types do get at least one doc mention somewhere (usually buried inside a robot's hardware page) — a discoverability gap, not an absence of content.
|
||||||
|
- **Orphan policy READMEs**: of the 16, 3 (Diffusion, TD-MPC, VQ-BeT) have **zero real published documentation anywhere** — the orphan stub is their only artifact. 13 duplicate an actively-maintained `.mdx` page under a different name. One (MolmoAct2) is proven to have silently rotted: a 39-line pre-refactor stub sits next to a 495-line maintained page that diverged 481 lines ago.
|
||||||
|
- **Root `README.md` itself has drifted**: its "Supported Hardware" line omits SO-101 (the flagship robot) entirely; its "SoTA Models" table links several policies to thin orphan stubs instead of the richer tutorial pages that exist for them.
|
||||||
|
- **4 concrete broken/stale copy-paste commands** verified in published pages (`act.mdx`, `il_robots.mdx` ×2 spots, `lerobot-dataset-v3.mdx`, `openarm.mdx` — the last added as recently as Jan 2026, so a fresh miss, not just aging drift).
|
||||||
|
- **Process/CI**: no CODEOWNERS, no `_redirects.yml`, no doc-consistency or toctree-completeness check anywhere; the docs CI workflow triggers only on `paths: docs/**`, so code-only PRs get zero automated docs signal (empirically responsible for the one real lag found — OMX's docs page, 41 days after its code). Despite this, 18 of 19 spot-checked new integrations (EVO1, FastWAM, LingBot-VA, MolmoAct2, RTC, VLA-JEPA, X-VLA, WallOSS, Damiao, Hope Jr, OpenArm, reBot B601, Reachy 2, Unitree G1, RoboCasa365, VLABench, IsaacLab Arena, etc.) shipped docs in the *same commit* as the feature — organic discipline is currently strong.
|
||||||
|
- **No content template** exists for policy or robot doc pages: heading sets/length vary 80-528 lines with almost no shared vocabulary; by contrast, benchmarks *do* have a written template (`adding_benchmarks.mdx`) and the resulting 8 pages are visibly consistent — direct proof a template is what produces consistency here.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Strengths
|
||||||
|
|
||||||
|
1. **Organic contributor discipline is currently much better than the June audit's numbers implied.** 18/19 spot-checked new policies/robots/benchmarks shipped docs in the same PR as the code; 17/18 CLI commands are documented; all teleoperator types get at least one mention. The underlying practice is healthier than the structural numbers (orphans, thin nav sections) suggest — the real problem is increasingly discoverability and consistency, not absence of effort.
|
||||||
|
2. **The doc-builder CI pipeline (PR previews, main build, versioned release builds) is solid, standard infrastructure already in place** — nothing bespoke to build or maintain, and it's the same tooling Transformers uses.
|
||||||
|
3. **A working, provably-effective template pattern already exists for one catalog (benchmarks)** — `adding_benchmarks.mdx`'s checklist + "at a glance" table produces 8 structurally consistent pages. This is the single best evidence in the repo for what fixes the policy/robot template gap, and it needs no new invention, just extension.
|
||||||
|
4. **Genuine single-source-of-truth patterns already exist and work**: `docs/source/contributing.md` is a real symlink to root `CONTRIBUTING.md` (verified via `ls -la`/`find -type l`). Most policy `src/README.md` files are also symlinked into `docs/source/policy_*_README.md`, eliminating drift at that layer (though the docs-side stub itself remains an unlinked orphan — see weaknesses).
|
||||||
|
5. **Zero broken toctree links**, and specific pages are genuinely high quality: `cheat-sheet.mdx` (spot-checked line-by-line against source, fully accurate), `bring_your_own_policies.mdx` (clear scope, working PR checklist), `installation.mdx`'s tables/OS tabs, and the SO-100/SO-101 hardware pages.
|
||||||
|
6. **Root `README.md` already contains a usable taxonomy for policies** (Imitation Learning / RL / VLAs / World Models / Reward Models) that the flat 15-item Policies sidebar never adopted — the raw material for a better IA already exists in-repo.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Weaknesses
|
||||||
|
|
||||||
|
### High
|
||||||
|
- **No comparison/decision-guidance page for any multi-option category** (Policies=15, Robots=11, Benchmarks=10, Teleoperators) — the exact content (decision rules, profiling snapshot) already exists in `AGENT_GUIDE.md` §6 but was never ported into `docs/source`.
|
||||||
|
- **The landing page (`index.mdx`) provides zero navigation** — 23 lines of marketing copy and a Discord link, no path to installation, cheat-sheet, or audience-specific guidance.
|
||||||
|
- **Teleoperators' dedicated nav section (2 pages) badly undersells real coverage** (16-17 implementations, all mentioned somewhere) — a new user browsing that sidebar section would wrongly conclude only phone/Isaac teleop exist. Content exists; it's a pure discoverability failure.
|
||||||
|
- **Multiple concrete broken/stale CLI examples in published, frequently copy-pasted pages**: `act.mdx` tells readers to eval with `lerobot-record` while the shown code uses `lerobot-rollout`; an identical shell-breaking snippet (`\`-terminated comment swallowing a continuation) is duplicated in `il_robots.mdx` and `lerobot-dataset-v3.mdx`; `il_robots.mdx:421` uses the dead `--control.push_to_hub` flag namespace; `openarm.mdx:208-210` (added Jan 2026) gives a `lerobot-record` example with flat pre-draccus flags that don't exist on the current config classes.
|
||||||
|
- **3 shipping policies (Diffusion, TD-MPC, VQ-BeT) have zero real published documentation** — only a thin, toctree-unreferenced stub exists for each.
|
||||||
|
- **Root `README.md` has drifted as the project scaled**: its Supported Hardware line omits SO-101 entirely; its SoTA Models table links ACT/Diffusion/VQ-BeT/Multitask-DiT/TDMPC/GR00T/SmolVLA to thin orphan stubs instead of the richer tutorial pages that exist for several of them.
|
||||||
|
- **No content template for policy/robot doc pages, unlike benchmarks** — heading sets and length vary 80-528 lines with almost no shared vocabulary (`molmoact2.mdx` has no Overview section at all; `smolvla.mdx` has no Overview/Architecture/Citation/License and reads as pure tutorial).
|
||||||
|
- **Docs CI triggers only on `paths: docs/**`** — code-only PRs get zero automated docs signal; this is the mechanism directly responsible for the one measured lag (OMX's docs page shipped 41 days after its code, in a separate PR).
|
||||||
|
- **Proof that orphan-stub drift isn't hypothetical**: `policy_molmoact2_README.md` is a stale 39-line pre-refactor snapshot sitting 481 lines behind the maintained 495-line page it once mirrored — nothing caught this silently rotting.
|
||||||
|
|
||||||
|
### Medium
|
||||||
|
- 16 orphaned `policy_*_README.md` files unreferenced by the toctree; 13 are pure duplicate cruft shadowing a maintained same-topic `.mdx` page, creating a false "two files to keep in sync" impression for contributors.
|
||||||
|
- No `_redirects.yml` anywhere — future cleanup of the orphans/renames has no 404 safety net.
|
||||||
|
- `installation.mdx` ends on a dangling forward-reference ("follow the link below to use LeRobot with your robot") — no link follows.
|
||||||
|
- `cheat-sheet.mdx`'s "Policy Types" line is stale (`act, diffusion, smolvla, pi05`) against the current 15-entry Policies section, and lists `diffusion`, which has no working toctree page at all.
|
||||||
|
- "Tutorials" (10 pages) mixes beginner, contributor, and RL-researcher content in one flat, unlabeled list.
|
||||||
|
- The EnvHub feature family is split inconsistently across two unrelated sections (`envhub.mdx`/`envhub_leisaac.mdx` under Simulation, `envhub_isaaclab_arena.mdx` under Benchmarks) despite identical framing/opening text.
|
||||||
|
- Admonition syntax is split GFM `[!NOTE]` vs. doc-builder `<Tip>` across sampled pages; since the site actually builds with `hf-doc-builder`, the majority pattern may not render as a styled callout at all (plausible, not independently render-verified).
|
||||||
|
- No CODEOWNERS file anywhere — zero designated reviewer routing for `docs/source`.
|
||||||
|
- `CONTRIBUTING.md` never links the good in-repo extension guides (`adding_benchmarks.mdx`, `bring_your_own_policies.mdx`, `integrate_hardware.mdx`) and states no docs-required policy.
|
||||||
|
- The "docs required" PR checklist convention is applied inconsistently: Policies and Benchmarks have an explicit required-docs checklist; `integrate_hardware.mdx` (Robots/Teleoperators) has none — a plausible root cause of the Teleoperators gap above.
|
||||||
|
- Growth repeatedly stacks simultaneous PRs against one shared, hand-edited `_toctree.yml` (6 benchmark PRs in 48h; 3 policy PRs in 3 days) — a merge/oversight risk even though no damage from it was found yet.
|
||||||
|
- `policy_sarm_README.md` has no source of truth left to sync against at all (SARM moved `policies/`→`rewards/`, no README carried over) — pure abandoned cruft next to the real 593-line `sarm.mdx`.
|
||||||
|
- No automated check anywhere for toctree completeness or README/mdx sync — the only CI gate is whether the doc-builder build succeeds, which doesn't catch missing coverage.
|
||||||
|
|
||||||
|
### Low
|
||||||
|
- No dedicated "Motors" section despite `motors/` being named as a distinct hardware layer in `CLAUDE.md`; `feetech.mdx`/`damiao.mdx` sit in a catch-all "Resources" section instead.
|
||||||
|
- "Sensors" is a single-page top-level nav section (`cameras.mdx` only).
|
||||||
|
- Policies (15), Benchmarks (10), and Robots (11) are flat, ungrouped lists with no internal sub-headings, despite `README.md` already having a usable taxonomy that could be reused.
|
||||||
|
- "SO-101" vs. "SO101" naming is inconsistent across ~15 files, and even within a single file (`so100.mdx` uses both).
|
||||||
|
- `hilserl.mdx` (950 lines) and `il_robots.mdx` (638 lines) each mix tutorial, reference, and troubleshooting content in one long page.
|
||||||
|
- `reachy2_camera` has zero doc mentions anywhere; the `zmq` camera backend is documented only incidentally inside `unitree_g1.mdx` rather than centrally in `cameras.mdx`.
|
||||||
|
- The PR template's "Documentation updated" line is a self-reported, unenforced checkbox.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Recommendations
|
||||||
|
|
||||||
|
### 1. Quick wins (small effort, ship this week, no maintainer proposal needed)
|
||||||
|
- Fix the 4 verified broken commands: `act.mdx` (`lerobot-record`→`lerobot-rollout`), the duplicated shell-breaking snippet in `il_robots.mdx` + `lerobot-dataset-v3.mdx`, `il_robots.mdx`'s `--control.push_to_hub`→`--dataset.push_to_hub`, and `openarm.mdx`'s invalid flat-flag record example.
|
||||||
|
- Delete the 3 fully-orphaned policy stubs (Diffusion, TD-MPC, VQ-BeT) and give each a real toctree page by promoting the existing `src/` README content.
|
||||||
|
- Delete the remaining 13 duplicate orphan stubs and the dead `policy_sarm_README.md`.
|
||||||
|
- Fix `installation.mdx`'s dangling ending and `cheat-sheet.mdx`'s stale policy-type list.
|
||||||
|
- Fix root `README.md`: add SO-101 to the Supported Hardware line; repoint the SoTA Models table's links from orphan stubs to the richer existing tutorial pages.
|
||||||
|
- Move `envhub_isaaclab_arena.mdx` into Simulation (one-line YAML change).
|
||||||
|
- Add 2-4 orientation sentences + links to the top of `index.mdx` (installation, cheat-sheet, "I have hardware" / "I don't" / "I want to contribute") without touching its marketing framing.
|
||||||
|
- Publish a single "Choosing a policy" page that ports the already-written decision rules from `AGENT_GUIDE.md` §6 — highest-leverage fix available.
|
||||||
|
- Start a `docs/source/_redirects.yml` now, before the orphan cleanup above creates the first real dead links.
|
||||||
|
- No action needed on the `contributing.md` "duplication" claim — verified it's a working symlink, not an anti-pattern.
|
||||||
|
|
||||||
|
### 2. Structural / UX changes (need a proposal + maintainer buy-in, phase as separate PRs)
|
||||||
|
- Split "Tutorials" into an explicit beginner-facing section vs. an "Advanced & Research/Contributor" section (mechanical YAML reorg, no content rewrites).
|
||||||
|
- Add a "Teleoperators" index page listing all types with one-line descriptions and links to wherever each is actually documented today.
|
||||||
|
- Re-group the flat Policies/Benchmarks/Robots lists into sub-headings reusing the taxonomy `README.md` already has.
|
||||||
|
- Treat a dedicated "Motors" section, deeper sub-grouping, and closing the remaining teleoperator/robot coverage gaps as a phased backlog of independently reviewable PRs rather than one restructure.
|
||||||
|
- Standardize SO-101/SO101 naming and (after confirming actual doc-builder rendering behavior) the admonition syntax, each as one mechanical, low-risk PR.
|
||||||
|
- Longer-term and biggest-ticket: port more of `AGENT_GUIDE.md`'s procedural content (training duration heuristics, eval targets, data-collection tips) into the published site.
|
||||||
|
|
||||||
|
### 3. Maintainability / process changes (prevent debt from reaccumulating at the current growth rate)
|
||||||
|
- Port `adding_benchmarks.mdx`'s explicit "writing a doc page" checklist/template into `bring_your_own_policies.mdx` and `integrate_hardware.mdx`.
|
||||||
|
- Add a CODEOWNERS entry for `docs/source/` (and ideally `_toctree.yml` specifically) so docs PRs get routed to a real reviewer.
|
||||||
|
- Link the extension guides from `CONTRIBUTING.md` and state a docs-required policy there.
|
||||||
|
- Add a lightweight CI/pre-commit script (not a full doc-builder run) that fails when (a) a `docs/source/*.mdx` file isn't reachable from `_toctree.yml`, or (b) a new `register_subclass` policy/robot/teleoperator/env lands with no corresponding doc file in the same diff.
|
||||||
|
- Decide the fate of the `policy_*_README.md` symlink convention going forward: it is currently self-perpetuating because `bring_your_own_policies.mdx`'s own checklist instructs new contributors to create the file that ends up orphaned. Either fold citation/paper content into the main `.mdx` tutorial, or wire the stub into the toctree as a linked citation anchor.
|
||||||
|
- Make the PR template's "Documentation updated" checkbox actionable (e.g., "(N/A if this PR only touches tests/CI/refactors)").
|
||||||
|
- Defer heavier generated-registry/support-matrix tooling (Transformers/Ultralytics-style single source of truth) until closer to 1.0.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Notes on cross-checking
|
||||||
|
|
||||||
|
Five independent audit passes fed this report; two disagreements were resolved during synthesis:
|
||||||
|
- **Toctree section count** — 15 is correct (independently parsed twice from `_toctree.yml`).
|
||||||
|
- **`contributing.md`** — it's a working symlink, not a hand-duplicated anti-pattern; one pass's `diff`-based claim didn't survive checking `find -type l`.
|
||||||
|
- **Orphan-README "sync mechanism"** — mostly fixed at the `src/`↔`docs` layer via symlinks, but that just moved the unresolved drift to the docs-side stub's absence from the toctree, and left old pre-symlink stubs (MolmoAct2) as dead leftovers.
|
||||||
|
- **Teleoperators coverage** — the nav-section framing ("~2/16") and the "mentioned somewhere" framing are both true; the real gap is discoverability, not content.
|
||||||
+1
-1
@@ -76,7 +76,7 @@ If your local computer doesn't have a powerful GPU, you can utilize Google Colab
|
|||||||
|
|
||||||
## Evaluating ACT
|
## Evaluating ACT
|
||||||
|
|
||||||
Once training is complete, you can evaluate your ACT policy using the `lerobot-record` command with your trained policy. This will run inference and record evaluation episodes:
|
Once training is complete, you can evaluate your ACT policy using the `lerobot-rollout` command with your trained policy. This will run inference and record evaluation episodes:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
lerobot-rollout \
|
lerobot-rollout \
|
||||||
|
|||||||
@@ -194,8 +194,8 @@ lerobot-record \
|
|||||||
--dataset.single_task="Navigate around obstacles" \
|
--dataset.single_task="Navigate around obstacles" \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
--dataset.encoder_threads=2 \
|
--dataset.encoder_threads=2 \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--display_data=true
|
--display_data=true
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|
||||||
Replace `your_username/dataset_name` with your Hugging Face username and a name for your dataset.
|
Replace `your_username/dataset_name` with your Hugging Face username and a name for your dataset.
|
||||||
|
|||||||
@@ -232,8 +232,8 @@ lerobot-record \
|
|||||||
--dataset.private=true \
|
--dataset.private=true \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
--dataset.encoder_threads=2 \
|
--dataset.encoder_threads=2 \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--display_data=true
|
--display_data=true
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|
||||||
### Replay
|
### Replay
|
||||||
@@ -278,6 +278,6 @@ lerobot-record \
|
|||||||
--dataset.num_episodes=10 \
|
--dataset.num_episodes=10 \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
--dataset.encoder_threads=2 \
|
--dataset.encoder_threads=2 \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--policy.path=outputs/train/hopejr_hand/checkpoints/last/pretrained_model
|
--policy.path=outputs/train/hopejr_hand/checkpoints/last/pretrained_model
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -207,8 +207,8 @@ lerobot-record \
|
|||||||
--dataset.num_episodes=5 \
|
--dataset.num_episodes=5 \
|
||||||
--dataset.single_task="Grab the black cube" \
|
--dataset.single_task="Grab the black cube" \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--dataset.encoder_threads=2
|
--dataset.encoder_threads=2
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
</hfoption>
|
</hfoption>
|
||||||
<hfoption id="API example">
|
<hfoption id="API example">
|
||||||
@@ -418,7 +418,7 @@ If you want to dive deeper into this important topic, you can check out the [blo
|
|||||||
|
|
||||||
## Visualize a dataset
|
## Visualize a dataset
|
||||||
|
|
||||||
If you uploaded your dataset to the hub with `--control.push_to_hub=true`, you can [visualize your dataset online](https://huggingface.co/spaces/lerobot/visualize_dataset) by copy pasting your repo id given by:
|
If you uploaded your dataset to the hub with `--dataset.push_to_hub=true`, you can [visualize your dataset online](https://huggingface.co/spaces/lerobot/visualize_dataset) by copy pasting your repo id given by:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
echo ${HF_USER}/so101_test
|
echo ${HF_USER}/so101_test
|
||||||
|
|||||||
@@ -44,8 +44,8 @@ lerobot-record \
|
|||||||
--dataset.num_episodes=5 \
|
--dataset.num_episodes=5 \
|
||||||
--dataset.single_task="Grab the black cube" \
|
--dataset.single_task="Grab the black cube" \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--dataset.encoder_threads=2
|
--dataset.encoder_threads=2
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|
||||||
See the [recording guide](./il_robots#record-a-dataset) for more details.
|
See the [recording guide](./il_robots#record-a-dataset) for more details.
|
||||||
|
|||||||
@@ -205,9 +205,10 @@ lerobot-record \
|
|||||||
--teleop.type=openarm_leader \
|
--teleop.type=openarm_leader \
|
||||||
--teleop.port=can1 \
|
--teleop.port=can1 \
|
||||||
--teleop.id=my_leader \
|
--teleop.id=my_leader \
|
||||||
--repo-id=my_hf_username/my_openarm_dataset \
|
--dataset.repo_id=my_hf_username/my_openarm_dataset \
|
||||||
--fps=30 \
|
--dataset.single_task="Grab the black cube" \
|
||||||
--num-episodes=10
|
--dataset.fps=30 \
|
||||||
|
--dataset.num_episodes=10
|
||||||
```
|
```
|
||||||
|
|
||||||
## Configuration Options
|
## Configuration Options
|
||||||
|
|||||||
@@ -161,8 +161,8 @@ lerobot-record \
|
|||||||
--dataset.private=true \
|
--dataset.private=true \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
--dataset.encoder_threads=2 \
|
--dataset.encoder_threads=2 \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--display_data=true
|
--display_data=true
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Specific Options
|
#### Specific Options
|
||||||
@@ -203,8 +203,8 @@ lerobot-record \
|
|||||||
--dataset.private=true \
|
--dataset.private=true \
|
||||||
--dataset.streaming_encoding=true \
|
--dataset.streaming_encoding=true \
|
||||||
--dataset.encoder_threads=2 \
|
--dataset.encoder_threads=2 \
|
||||||
# --dataset.rgb_encoder.vcodec=auto \
|
|
||||||
--display_data=true
|
--display_data=true
|
||||||
|
# Optionally, set --dataset.rgb_encoder.vcodec=auto to pick a specific video codec
|
||||||
```
|
```
|
||||||
|
|
||||||
##### `--robot.use_external_commands`
|
##### `--robot.use_external_commands`
|
||||||
|
|||||||
+13
-14
@@ -100,20 +100,19 @@ Once you are logged in, you can run inference in your setup by doing:
|
|||||||
lerobot-rollout \
|
lerobot-rollout \
|
||||||
--strategy.type=base \
|
--strategy.type=base \
|
||||||
--robot.type=so101_follower \
|
--robot.type=so101_follower \
|
||||||
--robot.port=/dev/ttyACM0 \ # <- Use your port
|
--robot.port=/dev/ttyACM0 \
|
||||||
--robot.id=my_blue_follower_arm \ # <- Use your robot id
|
--robot.id=my_blue_follower_arm \
|
||||||
--robot.cameras="{ front: {type: opencv, index_or_path: 8, width: 640, height: 480, fps: 30}}" \ # <- Use your cameras
|
--robot.cameras="{ front: {type: opencv, index_or_path: 8, width: 640, height: 480, fps: 30}}" \
|
||||||
--task="Grasp a lego block and put it in the bin." \ # <- Use the same task description you used in your dataset recording
|
--task="Grasp a lego block and put it in the bin." \
|
||||||
# <- RTC optional, use when running on low power hardware \
|
--policy.path=HF_USER/FINETUNE_MODEL_NAME
|
||||||
# --inference.type=rtc \
|
|
||||||
# --inference.rtc.execution_horizon=10 \
|
|
||||||
# --inference.rtc.max_guidance_weight=10.0 \
|
|
||||||
# <- Teleop optional if you want to teleoperate in between episodes \
|
|
||||||
# --teleop.type=so100_leader \
|
|
||||||
# --teleop.port=/dev/ttyACM0 \
|
|
||||||
# --teleop.id=my_red_leader_arm \
|
|
||||||
# --display_data=true #optional use if you want to see the camera stream \
|
|
||||||
--policy.path=HF_USER/FINETUNE_MODEL_NAME # <- Use your fine-tuned model
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Replace `--robot.port`, `--robot.id`, `--robot.cameras`, `--task`, and `--policy.path` with your own port, robot ID, camera setup, task description (matching what you used when recording your dataset), and fine-tuned model repo ID.
|
||||||
|
|
||||||
|
A few optional flags you can add to the command above:
|
||||||
|
|
||||||
|
- **RTC** (useful on low-power hardware): `--inference.type=rtc --inference.rtc.execution_horizon=10 --inference.rtc.max_guidance_weight=10.0`
|
||||||
|
- **Teleoperate in between episodes**: `--teleop.type=so100_leader --teleop.port=/dev/ttyACM0 --teleop.id=my_red_leader_arm`
|
||||||
|
- **See the camera stream**: `--display_data=true`
|
||||||
|
|
||||||
Depending on your evaluation setup, you can configure the duration and the number of episodes to record for your evaluation suite.
|
Depending on your evaluation setup, you can configure the duration and the number of episodes to record for your evaluation suite.
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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"),
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -713,6 +713,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 +733,7 @@ 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, hub_download_kwargs, is_local_source
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. Validate that all overrides were used
|
# 4. Validate that all overrides were used
|
||||||
@@ -921,6 +923,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
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 +947,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 +965,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)
|
||||||
@@ -975,7 +979,9 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
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, hub_download_kwargs, is_local_source
|
||||||
|
)
|
||||||
|
|
||||||
return steps, remaining_override_keys
|
return steps, remaining_override_keys
|
||||||
|
|
||||||
@@ -1139,6 +1145,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
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 +1164,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 +1185,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,6 +1199,12 @@ 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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -326,8 +326,17 @@ class RolloutConfig:
|
|||||||
|
|
||||||
policy_path = parser.get_path_arg("policy")
|
policy_path = parser.get_path_arg("policy")
|
||||||
if policy_path:
|
if policy_path:
|
||||||
cli_overrides = parser.get_cli_overrides("policy")
|
yaml_overrides = parser.get_yaml_overrides("policy")
|
||||||
self.policy = PreTrainedConfig.from_pretrained(policy_path, cli_overrides=cli_overrides)
|
cli_overrides = parser.get_cli_overrides("policy") or []
|
||||||
|
policy_overrides = yaml_overrides + cli_overrides
|
||||||
|
pretrained_revision = parser.parse_arg("pretrained_revision", cli_overrides)
|
||||||
|
if pretrained_revision is None:
|
||||||
|
pretrained_revision = parser.parse_arg("pretrained_revision", yaml_overrides)
|
||||||
|
self.policy = PreTrainedConfig.from_pretrained(
|
||||||
|
policy_path,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
cli_overrides=policy_overrides,
|
||||||
|
)
|
||||||
self.policy.pretrained_path = policy_path
|
self.policy.pretrained_path = policy_path
|
||||||
if self.policy is None:
|
if self.policy is None:
|
||||||
raise ValueError("--policy.path is required for rollout")
|
raise ValueError("--policy.path is required for rollout")
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from threading import Event
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.configs import FeatureType
|
from lerobot.configs import FeatureType, PreTrainedConfig
|
||||||
from lerobot.datasets import (
|
from lerobot.datasets import (
|
||||||
LeRobotDataset,
|
LeRobotDataset,
|
||||||
aggregate_pipeline_dataset_features,
|
aggregate_pipeline_dataset_features,
|
||||||
@@ -159,6 +159,35 @@ class RolloutContext:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy:
|
||||||
|
"""Load policy weights, keeping adapter and base-model revisions independent."""
|
||||||
|
pretrained_revision = policy_config.pretrained_revision
|
||||||
|
policy_class = get_policy_class(policy_config.type)
|
||||||
|
|
||||||
|
if not policy_config.use_peft:
|
||||||
|
return policy_class.from_pretrained(
|
||||||
|
policy_config.pretrained_path,
|
||||||
|
config=policy_config,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
)
|
||||||
|
|
||||||
|
from peft import PeftConfig, PeftModel
|
||||||
|
|
||||||
|
peft_path = policy_config.pretrained_path
|
||||||
|
peft_config = PeftConfig.from_pretrained(peft_path, revision=pretrained_revision)
|
||||||
|
policy = policy_class.from_pretrained(
|
||||||
|
pretrained_name_or_path=peft_config.base_model_name_or_path,
|
||||||
|
config=policy_config,
|
||||||
|
revision=peft_config.revision,
|
||||||
|
)
|
||||||
|
return PeftModel.from_pretrained(
|
||||||
|
policy,
|
||||||
|
peft_path,
|
||||||
|
config=peft_config,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_rollout_context(
|
def build_rollout_context(
|
||||||
cfg: RolloutConfig,
|
cfg: RolloutConfig,
|
||||||
shutdown_event: Event,
|
shutdown_event: Event,
|
||||||
@@ -176,7 +205,6 @@ def build_rollout_context(
|
|||||||
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
||||||
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
||||||
policy_config = cfg.policy
|
policy_config = cfg.policy
|
||||||
policy_class = get_policy_class(policy_config.type)
|
|
||||||
|
|
||||||
if hasattr(policy_config, "compile_model"):
|
if hasattr(policy_config, "compile_model"):
|
||||||
policy_config.compile_model = cfg.use_torch_compile
|
policy_config.compile_model = cfg.use_torch_compile
|
||||||
@@ -187,17 +215,7 @@ def build_rollout_context(
|
|||||||
"Please use `cpu` or `cuda` backend."
|
"Please use `cpu` or `cuda` backend."
|
||||||
)
|
)
|
||||||
|
|
||||||
if policy_config.use_peft:
|
policy = _load_pretrained_policy(policy_config)
|
||||||
from peft import PeftConfig, PeftModel
|
|
||||||
|
|
||||||
peft_path = policy_config.pretrained_path
|
|
||||||
peft_config = PeftConfig.from_pretrained(peft_path)
|
|
||||||
policy = policy_class.from_pretrained(
|
|
||||||
pretrained_name_or_path=peft_config.base_model_name_or_path, config=policy_config
|
|
||||||
)
|
|
||||||
policy = PeftModel.from_pretrained(policy, peft_path, config=peft_config)
|
|
||||||
else:
|
|
||||||
policy = policy_class.from_pretrained(policy_config.pretrained_path, config=policy_config)
|
|
||||||
|
|
||||||
if is_rtc:
|
if is_rtc:
|
||||||
policy.config.rtc_config = cfg.inference.rtc
|
policy.config.rtc_config = cfg.inference.rtc
|
||||||
@@ -392,6 +410,7 @@ def build_rollout_context(
|
|||||||
preprocessor, postprocessor = make_pre_post_processors(
|
preprocessor, postprocessor = make_pre_post_processors(
|
||||||
policy_cfg=policy_config,
|
policy_cfg=policy_config,
|
||||||
pretrained_path=cfg.policy.pretrained_path,
|
pretrained_path=cfg.policy.pretrained_path,
|
||||||
|
pretrained_revision=policy_config.pretrained_revision,
|
||||||
dataset_stats=dataset_stats,
|
dataset_stats=dataset_stats,
|
||||||
preprocessor_overrides={
|
preprocessor_overrides={
|
||||||
"device_processor": {"device": cfg.device},
|
"device_processor": {"device": cfg.device},
|
||||||
|
|||||||
@@ -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,6 +1015,12 @@ def eval_policy_all(
|
|||||||
recording_private=recording_private,
|
recording_private=recording_private,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Set the shared policy's mode before launching any workers. Restoring it
|
||||||
|
# inside individual tasks would let one task enable training mode while
|
||||||
|
# another task is still evaluating.
|
||||||
|
was_training = policy.training
|
||||||
|
policy.eval()
|
||||||
|
try:
|
||||||
if max_parallel_tasks <= 1:
|
if max_parallel_tasks <= 1:
|
||||||
prefetch_thread: threading.Thread | None = None
|
prefetch_thread: threading.Thread | None = None
|
||||||
for i, (task_group, task_id, env) in enumerate(tasks):
|
for i, (task_group, task_id, env) in enumerate(tasks):
|
||||||
@@ -1044,6 +1055,8 @@ def eval_policy_all(
|
|||||||
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()
|
||||||
|
finally:
|
||||||
|
policy.train(was_training)
|
||||||
|
|
||||||
# compute aggregated metrics helper (robust to lists/scalars)
|
# compute aggregated metrics helper (robust to lists/scalars)
|
||||||
def _agg_from_list(xs):
|
def _agg_from_list(xs):
|
||||||
|
|||||||
@@ -453,9 +453,11 @@ def record(
|
|||||||
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
||||||
)
|
)
|
||||||
|
|
||||||
robot.connect()
|
# Connect the teleoperator before the robot so the robot isn't left idle (and possibly
|
||||||
|
# tripping a firmware watchdog) during teleop init. Matches lerobot_teleoperate.py.
|
||||||
if teleop is not None:
|
if teleop is not None:
|
||||||
teleop.connect()
|
teleop.connect()
|
||||||
|
robot.connect()
|
||||||
|
|
||||||
listener, events = init_keyboard_listener()
|
listener, events = init_keyboard_listener()
|
||||||
|
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ from lerobot.robots import ( # noqa: F401
|
|||||||
earthrover_mini_plus,
|
earthrover_mini_plus,
|
||||||
hope_jr,
|
hope_jr,
|
||||||
koch_follower,
|
koch_follower,
|
||||||
|
lekiwi,
|
||||||
make_robot_from_config,
|
make_robot_from_config,
|
||||||
omx_follower,
|
omx_follower,
|
||||||
openarm_follower,
|
openarm_follower,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user