mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
Compare commits
29 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c112ba6957 | |||
| 741005d719 | |||
| 2e8345a5cc | |||
| ef88d4e52b | |||
| 64b23178d5 | |||
| f66e5128ec | |||
| 1e3a158e13 | |||
| dc0cee9c75 | |||
| 7e241bd630 | |||
| e867359d09 | |||
| f1efa588b8 | |||
| 3e37269dc6 | |||
| 8b2678318c | |||
| 6b56bf299b | |||
| b1ee35e637 | |||
| 6312be2d3f | |||
| d36f429a30 | |||
| bad0260a46 | |||
| adccdea1cf | |||
| 8135a8a8d1 | |||
| 2aba372b4e | |||
| c841a0c258 | |||
| a3ddba2454 | |||
| 99443a936d | |||
| 7a3298ea26 | |||
| 2d8f5f314e | |||
| 732a12108e | |||
| 29fcf057dc | |||
| 81db623b44 |
@@ -53,7 +53,7 @@ permissions:
|
||||
contents: read
|
||||
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
|
||||
# Cancel in-flight runs for the same branch/PR.
|
||||
|
||||
@@ -27,7 +27,7 @@ on:
|
||||
|
||||
# Sets up the environment variables
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
DOCKER_IMAGE_NAME_CPU: huggingface/lerobot-cpu:latest
|
||||
DOCKER_IMAGE_NAME_GPU: huggingface/lerobot-gpu:latest
|
||||
|
||||
@@ -24,19 +24,24 @@ on:
|
||||
required: false
|
||||
type: string
|
||||
|
||||
# Triggers the workflow on push events to main for the docs folder
|
||||
# Triggers on pushes to main that touch the docs or the sources the API reference is generated from.
|
||||
# `src/**` is included because the API reference is built from docstrings via `[[autodoc]]`: without it,
|
||||
# published API pages would go stale as soon as a docstring changed.
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**"
|
||||
- "src/**"
|
||||
|
||||
# Triggers the workflow on pull request events targeting main for the docs folder
|
||||
# Same for pull requests, so a docstring change gets a preview build and a broken `[[autodoc]]` path
|
||||
# fails the PR rather than main.
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**"
|
||||
- "src/**"
|
||||
|
||||
release:
|
||||
types: [published]
|
||||
@@ -59,12 +64,21 @@ jobs:
|
||||
with:
|
||||
commit_sha: ${{ github.sha }}
|
||||
package: lerobot
|
||||
# doc-builder ships a mock-deps registry entry for lerobot, so the reusable workflow takes its
|
||||
# "light install" path: `pip install ./lerobot --no-deps` plus a handful of real dependencies.
|
||||
# That is not enough to import lerobot — draccus runs `register_subclass` at import time and
|
||||
# `processor/converters.py` calls `functools.singledispatch.register(torch.Tensor)`, neither of
|
||||
# which works against a mock. Install the package for real before the build.
|
||||
pre_command: uv pip install "./lerobot[dataset]"
|
||||
# `--version main` is load-bearing: without `--not_python_module`, doc-builder falls back to
|
||||
# `lerobot.__version__` and only maps that to the default branch when it contains "dev". Our main
|
||||
# branch carries a release version (0.6.2), so omitting this would publish the main docs to
|
||||
# /lerobot/v0.6.2/ instead of /lerobot/main/ and disable notebook building.
|
||||
additional_args: >-
|
||||
--not_python_module
|
||||
${{
|
||||
(github.event_name == 'release' && format('--version {0}', github.event.release.tag_name)) ||
|
||||
(inputs.version != '' && format('--version {0}', inputs.version)) ||
|
||||
''
|
||||
'--version main'
|
||||
}}
|
||||
secrets:
|
||||
token: ${{ secrets.HUGGINGFACE_PUSH }}
|
||||
@@ -83,4 +97,6 @@ jobs:
|
||||
commit_sha: ${{ github.event.pull_request.head.sha }}
|
||||
pr_number: ${{ github.event.number }}
|
||||
package: lerobot
|
||||
additional_args: --not_python_module
|
||||
# See the comment on build_main_docs. The PR workflow passes its own `--version pr_<n>`, so no
|
||||
# additional_args are needed here.
|
||||
pre_command: uv pip install "./lerobot[dataset]"
|
||||
|
||||
@@ -48,7 +48,7 @@ permissions:
|
||||
|
||||
# Sets up the environment variables
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
|
||||
# Ensures that only the latest commit for a PR or branch is built, canceling older runs.
|
||||
|
||||
@@ -37,7 +37,7 @@ permissions:
|
||||
|
||||
# Sets up the environment variables
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
DOCKER_IMAGE_NAME: huggingface/lerobot-gpu
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ on:
|
||||
|
||||
# Sets up the environment variables
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
DOCKER_IMAGE_NAME: huggingface/lerobot-gpu:latest-deps
|
||||
|
||||
|
||||
@@ -56,3 +56,41 @@ jobs:
|
||||
uses: pre-commit/action@2c7b3805fd2a0fd8c1884dcaebf91fc102a13ecd # v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --show-diff-on-failure --color=always
|
||||
|
||||
# This job runs the examples in our docstrings and validates the doctest allowlist.
|
||||
# See docs/source/writing_docstrings.mdx for the standard these enforce.
|
||||
doc-checks:
|
||||
name: Run Documentation Checks (Doctests)
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
# Examples that need a physical robot, a serial port or a Hub download are skipped by content.
|
||||
# Everything else has to actually run. See src/lerobot/utils/doctest_utils.py.
|
||||
SKIP_HARDWARE_DOCTEST: "1"
|
||||
SKIP_CUDA_DOCTEST: "1"
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Setup uv and Python
|
||||
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2
|
||||
with:
|
||||
enable-cache: true
|
||||
version: "0.11.30"
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv sync --locked --extra test --extra dataset
|
||||
|
||||
- name: Check the doctest list is sorted and its paths exist
|
||||
run: make check-doctest-list
|
||||
|
||||
- name: Check documented arguments match their signatures
|
||||
run: make check-docstrings
|
||||
|
||||
- name: Check docstring coverage has not regressed
|
||||
run: uv run --with interrogate interrogate --config=pyproject.toml
|
||||
|
||||
- name: Run doctests
|
||||
run: make doctest
|
||||
|
||||
@@ -21,7 +21,7 @@ on:
|
||||
|
||||
# Sets up the environment variables
|
||||
env:
|
||||
UV_VERSION: "0.8.0"
|
||||
UV_VERSION: "0.11.30"
|
||||
PYTHON_VERSION: "3.12"
|
||||
|
||||
jobs:
|
||||
|
||||
@@ -19,8 +19,8 @@ on:
|
||||
workflow_dispatch:
|
||||
|
||||
# Runs at 02:00
|
||||
# schedule:
|
||||
# - cron: "0 2 * * *"
|
||||
schedule:
|
||||
- cron: "0 2 * * *"
|
||||
|
||||
env:
|
||||
CLOSE_ISSUE_MESSAGE: >
|
||||
@@ -31,7 +31,7 @@ env:
|
||||
Feel free to reopen if is still relevant, or to ping a collaborator if you have any questions.
|
||||
WARN_ISSUE_MESSAGE: >
|
||||
This issue has been automatically marked as stale because it has not had
|
||||
recent activity (1 year). It will be closed if no further activity occurs.
|
||||
recent activity (1 year). It will be closed if no further activity occurs within 30 days.
|
||||
Any change, comment or update to this issue will reset this count.
|
||||
Thank you for your contributions.
|
||||
WARN_PR_MESSAGE: >
|
||||
@@ -61,8 +61,8 @@ jobs:
|
||||
exempt-pr-labels: never-stale
|
||||
days-before-issue-stale: 365
|
||||
days-before-issue-close: 30
|
||||
days-before-pr-stale: 365
|
||||
days-before-pr-close: 30
|
||||
days-before-pr-stale: -1
|
||||
days-before-pr-close: -1
|
||||
delete-branch: true
|
||||
close-issue-message: ${{ env.CLOSE_ISSUE_MESSAGE }}
|
||||
close-pr-message: ${{ env.CLOSE_PR_MESSAGE }}
|
||||
|
||||
+11
-2
@@ -67,7 +67,11 @@ repos:
|
||||
args: [--prose-wrap=preserve]
|
||||
# Jinja2 model-card templates use a .md extension but contain {% ... %} /
|
||||
# {{ ... }} tags that prettier's Markdown formatter mangles (e.g. table loops).
|
||||
exclude: ^src/lerobot/templates/.*\.md$
|
||||
#
|
||||
# docs/source/api/ holds the generated API reference. Its `[[autodoc]]` blocks restrict output
|
||||
# to an indented `- member` list, which prettier reads as a lazy paragraph continuation and
|
||||
# joins onto one line — silently turning a member list into part of the directive.
|
||||
exclude: ^(src/lerobot/templates/.*\.md|docs/source/api/.*\.mdx)$
|
||||
|
||||
##### Security #####
|
||||
- repo: https://github.com/gitleaks/gitleaks
|
||||
@@ -104,8 +108,13 @@ repos:
|
||||
# args: ["--docstring-style", "google", "-v", "2"]
|
||||
# exclude: ^tests/.*$
|
||||
|
||||
# interrogate runs in CI (quality.yml, doc-checks job) rather than here. Its 1.7.0 release still imports
|
||||
# the deprecated `py` package, which resolves against whatever `py` happens to be importable in
|
||||
# pre-commit's isolated env — on a machine with miniconda on the path that is a stray `py.py` and the
|
||||
# hook dies before it reads any config. The gate is the same either way; the CI step is just reliable.
|
||||
# - repo: https://github.com/econchick/interrogate
|
||||
# rev: 1.7.0
|
||||
# hooks:
|
||||
# - id: interrogate
|
||||
# args: ["-vv", "--config=pyproject.toml"]
|
||||
# args: ["--config=pyproject.toml"]
|
||||
# pass_filenames: false
|
||||
|
||||
@@ -50,6 +50,10 @@ To run checks manually on all files:
|
||||
pre-commit run --all-files
|
||||
```
|
||||
|
||||
### Docstrings
|
||||
|
||||
The API reference is generated from the docstrings in `src/lerobot/`. If you add or change anything public, follow the [docstring standard](https://huggingface.co/docs/lerobot/writing_docstrings) — the format is parsed by the renderer and checked in CI.
|
||||
|
||||
### Running Tests
|
||||
|
||||
We use `pytest`. First, ensure you have test artifacts by installing **git-lfs**:
|
||||
|
||||
@@ -184,3 +184,29 @@ test-smolvla-ete-eval:
|
||||
# backend, so it does not require a real model checkpoint or GPU.
|
||||
annotation-e2e:
|
||||
uv run python -m tests.annotations.run_e2e_smoke
|
||||
|
||||
# Docstring & doctest checks. See docs/source/writing_docstrings.mdx for the standard these enforce.
|
||||
|
||||
# Run the examples in the docstrings listed in utils/documentation_tests.txt. Hardware and GPU examples are
|
||||
# skipped by content (see src/lerobot/utils/doctest_utils.py); CI sets both flags.
|
||||
doctest:
|
||||
@files=$$(grep -v '^\s*#' utils/documentation_tests.txt | grep -v '^\s*$$'); \
|
||||
if [ -z "$$files" ]; then \
|
||||
echo "utils/documentation_tests.txt lists no files; nothing to run."; \
|
||||
else \
|
||||
SKIP_HARDWARE_DOCTEST=1 uv run pytest --doctest-modules --no-header -q $$files; \
|
||||
fi
|
||||
|
||||
check-doctest-list:
|
||||
uv run python utils/check_doctest_list.py
|
||||
|
||||
fix-doctest-list:
|
||||
uv run python utils/check_doctest_list.py --fix_and_overwrite
|
||||
|
||||
check-docstrings:
|
||||
uv run python utils/check_docstrings.py
|
||||
uv run python utils/check_config_docstrings.py
|
||||
|
||||
fix-docstrings:
|
||||
uv run python utils/check_docstrings.py --fix_and_overwrite
|
||||
uv run python utils/check_doctest_list.py --fix_and_overwrite
|
||||
|
||||
@@ -128,6 +128,23 @@ lerobot-eval \
|
||||
|
||||
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).
|
||||
|
||||
### Third-Party Hardware
|
||||
|
||||
Beyond the natively supported hardware, the community maintains a growing ecosystem of plugins for other robots, teleoperators, cameras, and sensors - UFACTORY xArm, Universal Robots UR5e, Franka, AgileX Piper, Trossen WidowX, ARX5, I2RT YAM, GELLO, SpaceMouse, Meta Quest, ROS 2 bridges, tactile and depth cameras, and more.
|
||||
|
||||
Plugins are auto-discovered by package name: LeRobot imports any installed package prefixed with `lerobot_robot_`, `lerobot_teleoperator_`, or `lerobot_camera_`. Install one and use the `type` it registers straight from the CLI:
|
||||
|
||||
```bash
|
||||
pip install lerobot_robot_<name> lerobot_teleoperator_<name>
|
||||
|
||||
lerobot-record \
|
||||
--robot.type=<robot_name> \
|
||||
--teleop.type=<teleoperator_name> \
|
||||
--dataset.repo_id=${HF_USER}/my-dataset
|
||||
```
|
||||
|
||||
Browse the full list in the [Third-Party Robots & Teleoperators](https://huggingface.co/docs/lerobot/main/third_party_robots) and [Third-Party Cameras & Sensors](https://huggingface.co/docs/lerobot/main/third_party_sensors) documentation.
|
||||
|
||||
## Resources
|
||||
|
||||
- **[Documentation](https://huggingface.co/docs/lerobot/index):** The complete guide to tutorials & API.
|
||||
|
||||
+60
@@ -0,0 +1,60 @@
|
||||
# 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.
|
||||
|
||||
"""Root conftest: makes doctest collection use LeRobot's parser.
|
||||
|
||||
This only affects `--doctest-modules` runs (see `make doctest`). The test suite itself is configured by
|
||||
`tests/conftest.py`.
|
||||
"""
|
||||
|
||||
import doctest
|
||||
|
||||
import _pytest.doctest
|
||||
|
||||
from lerobot.utils.doctest_utils import LeRobotDoctestModule, LeRobotDocTestParser
|
||||
|
||||
# Lets an example opt out of output comparison with `# doctest: +IGNORE_RESULT`, for calls whose output is
|
||||
# a progress bar or otherwise not reproducible.
|
||||
IGNORE_RESULT = doctest.register_optionflag("IGNORE_RESULT")
|
||||
|
||||
OutputChecker = doctest.OutputChecker
|
||||
|
||||
|
||||
class CustomOutputChecker(OutputChecker):
|
||||
"""An output checker that honours the `IGNORE_RESULT` flag."""
|
||||
|
||||
def check_output(self, want, got, optionflags):
|
||||
"""Return `True` when `IGNORE_RESULT` is set, otherwise defer to stdlib.
|
||||
|
||||
Args:
|
||||
want (`str`):
|
||||
The expected output.
|
||||
got (`str`):
|
||||
The actual output.
|
||||
optionflags (`int`):
|
||||
Bitmask of active doctest option flags.
|
||||
|
||||
Returns:
|
||||
`bool`: Whether the output is considered a match.
|
||||
"""
|
||||
if IGNORE_RESULT & optionflags:
|
||||
return True
|
||||
return OutputChecker.check_output(self, want, got, optionflags)
|
||||
|
||||
|
||||
# Reassigning these module attributes is how doctest behaviour is customised; mypy sees it as assigning to
|
||||
# a type, which is exactly what is intended here.
|
||||
doctest.OutputChecker = CustomOutputChecker # type: ignore[misc]
|
||||
_pytest.doctest.DoctestModule = LeRobotDoctestModule
|
||||
doctest.DocTestParser = LeRobotDocTestParser # type: ignore[misc]
|
||||
@@ -165,6 +165,8 @@
|
||||
title: OpenArm
|
||||
- local: rebot_b601
|
||||
title: reBot B601-DM
|
||||
- local: third_party_robots
|
||||
title: Third-Party Robots & Teleoperators
|
||||
title: "Robots"
|
||||
- sections:
|
||||
- local: phone_teleop
|
||||
@@ -175,6 +177,8 @@
|
||||
- sections:
|
||||
- local: cameras
|
||||
title: Cameras
|
||||
- local: third_party_sensors
|
||||
title: Third-Party Cameras & Sensors
|
||||
title: "Sensors"
|
||||
- sections:
|
||||
- local: notebooks
|
||||
@@ -187,6 +191,28 @@
|
||||
- sections:
|
||||
- local: contributing
|
||||
title: Contribute to LeRobot
|
||||
- local: writing_docstrings
|
||||
title: Writing docstrings
|
||||
- local: backwardcomp
|
||||
title: Backward compatibility
|
||||
title: "About"
|
||||
- sections:
|
||||
- local: api/robots
|
||||
title: Robots
|
||||
- local: api/teleoperators
|
||||
title: Teleoperators
|
||||
- local: api/cameras
|
||||
title: Cameras
|
||||
- local: api/motors
|
||||
title: Motors
|
||||
- local: api/datasets
|
||||
title: Datasets
|
||||
- local: api/policies
|
||||
title: Policies
|
||||
- local: api/processor
|
||||
title: Processors
|
||||
- local: api/envs
|
||||
title: Environments
|
||||
- local: api/configs
|
||||
title: Configuration
|
||||
title: "API Reference"
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
# Cameras
|
||||
|
||||
Cameras supply the image observations a policy sees. Every backend — OpenCV, Intel RealSense, Reachy 2 —
|
||||
implements the [`Camera`] interface, so swapping hardware does not change the code that reads frames.
|
||||
|
||||
See the [Cameras guide](../cameras) for choosing and configuring a camera, and
|
||||
[Third-Party Cameras & Sensors](../third_party_sensors) for devices outside the core set.
|
||||
|
||||
## Camera
|
||||
|
||||
[[autodoc]] lerobot.cameras.Camera
|
||||
- connect
|
||||
- disconnect
|
||||
- read
|
||||
- async_read
|
||||
- find_cameras
|
||||
|
||||
## CameraConfig
|
||||
|
||||
[[autodoc]] lerobot.cameras.CameraConfig
|
||||
|
||||
## make_cameras_from_configs
|
||||
|
||||
[[autodoc]] lerobot.cameras.make_cameras_from_configs
|
||||
@@ -0,0 +1,27 @@
|
||||
# Configuration
|
||||
|
||||
LeRobot configuration is plain dataclasses parsed by [draccus](https://github.com/dlwh/draccus), so every
|
||||
field is settable from the CLI. [`TrainPipelineConfig`] is the top-level object for `lerobot-train`.
|
||||
|
||||
Polymorphic configs (policies, robots, environments) use `draccus.ChoiceRegistry`: a subclass registers
|
||||
itself with `@register_subclass("name")` and is then selectable by that name on the command line.
|
||||
|
||||
## TrainPipelineConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.train.TrainPipelineConfig
|
||||
|
||||
## PreTrainedConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.PreTrainedConfig
|
||||
|
||||
## DatasetConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.DatasetConfig
|
||||
|
||||
## EvalConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.EvalConfig
|
||||
|
||||
## WandBConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.WandBConfig
|
||||
@@ -0,0 +1,23 @@
|
||||
# Datasets
|
||||
|
||||
[`LeRobotDataset`] is the format every LeRobot script reads and writes. It is episode-aware, decodes video
|
||||
observations on the fly, and round-trips to the Hugging Face Hub.
|
||||
|
||||
See [Using LeRobotDataset](../lerobot-dataset-v3) for the format and the common operations,
|
||||
[Porting Large Datasets](../porting_datasets_v3) for migration, and [Tools](../tools) for the CLI.
|
||||
|
||||
## LeRobotDataset
|
||||
|
||||
[[autodoc]] lerobot.datasets.LeRobotDataset
|
||||
|
||||
## LeRobotDatasetMetadata
|
||||
|
||||
[[autodoc]] lerobot.datasets.LeRobotDatasetMetadata
|
||||
|
||||
## MultiLeRobotDataset
|
||||
|
||||
[[autodoc]] lerobot.datasets.MultiLeRobotDataset
|
||||
|
||||
## StreamingLeRobotDataset
|
||||
|
||||
[[autodoc]] lerobot.datasets.StreamingLeRobotDataset
|
||||
@@ -0,0 +1,19 @@
|
||||
# Environments
|
||||
|
||||
Simulation environments are configured through [`EnvConfig`] and built by [`make_env`]. Each subclass
|
||||
declares its `gym_kwargs` and how to construct the vectorised environments.
|
||||
|
||||
See [Environments from the Hub](../envhub) for using published environments and
|
||||
[Adding a New Benchmark](../adding_benchmarks) for contributing one.
|
||||
|
||||
## EnvConfig
|
||||
|
||||
[[autodoc]] lerobot.envs.EnvConfig
|
||||
|
||||
## make_env
|
||||
|
||||
[[autodoc]] lerobot.envs.make_env
|
||||
|
||||
## make_env_config
|
||||
|
||||
[[autodoc]] lerobot.envs.make_env_config
|
||||
@@ -0,0 +1,23 @@
|
||||
# Motors
|
||||
|
||||
`MotorsBus` is the low-level interface to a chain of servos on a serial bus. Robots use it to read positions
|
||||
and write goal positions; you rarely touch it directly unless you are adding hardware.
|
||||
|
||||
See [Bring Your Own Hardware](../integrate_hardware) for adding a new bus, and
|
||||
[Updating Feetech Firmware](../feetech) and [Damiao Motors and CAN Bus](../damiao) for device-specific notes.
|
||||
|
||||
## MotorsBus
|
||||
|
||||
[[autodoc]] lerobot.motors.motors_bus.MotorsBus
|
||||
|
||||
## Motor
|
||||
|
||||
[[autodoc]] lerobot.motors.Motor
|
||||
|
||||
## MotorCalibration
|
||||
|
||||
[[autodoc]] lerobot.motors.MotorCalibration
|
||||
|
||||
## MotorNormMode
|
||||
|
||||
[[autodoc]] lerobot.motors.MotorNormMode
|
||||
@@ -0,0 +1,177 @@
|
||||
# Policies
|
||||
|
||||
Every policy inherits [`PreTrainedPolicy`], which combines a `torch.nn.Module` with the Hub mixin, so any
|
||||
policy can be pushed to and loaded from the Hugging Face Hub with the same two calls.
|
||||
|
||||
Each policy has its own guide with training recipes and results — [ACT](../act), [SmolVLA](../smolvla),
|
||||
[π₀](../pi0), [π₀.₅](../pi05) and the rest are listed under Policies. To add one, see
|
||||
[Adding a Policy](../bring_your_own_policies).
|
||||
|
||||
## PreTrainedPolicy
|
||||
|
||||
The abstract base class every policy subclasses. `forward` computes the training loss, `select_action`
|
||||
returns one action at a time for control loops, and `predict_action_chunk` returns a full action chunk.
|
||||
|
||||
[[autodoc]] lerobot.policies.pretrained.PreTrainedPolicy
|
||||
- forward
|
||||
- predict_action_chunk
|
||||
- select_action
|
||||
- get_optim_params
|
||||
- reset
|
||||
- from_pretrained
|
||||
- supports_rtc
|
||||
- push_model_to_hub
|
||||
- wrap_with_peft
|
||||
|
||||
## PreTrainedConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.PreTrainedConfig
|
||||
|
||||
## make_policy
|
||||
|
||||
[[autodoc]] lerobot.policies.factory.make_policy
|
||||
|
||||
## get_policy_class
|
||||
|
||||
[[autodoc]] lerobot.policies.factory.get_policy_class
|
||||
|
||||
## make_policy_config
|
||||
|
||||
[[autodoc]] lerobot.policies.factory.make_policy_config
|
||||
|
||||
## make_pre_post_processors
|
||||
|
||||
[[autodoc]] lerobot.policies.factory.make_pre_post_processors
|
||||
|
||||
## ACT
|
||||
|
||||
[[autodoc]] lerobot.policies.act.modeling_act.ACTPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.act.configuration_act.ACTConfig
|
||||
|
||||
## SmolVLA
|
||||
|
||||
[[autodoc]] lerobot.policies.smolvla.modeling_smolvla.SmolVLAPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.smolvla.configuration_smolvla.SmolVLAConfig
|
||||
|
||||
## π₀ (PI0)
|
||||
|
||||
[[autodoc]] lerobot.policies.pi0.modeling_pi0.PI0Policy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.pi0.configuration_pi0.PI0Config
|
||||
|
||||
## π₀-FAST (PI0Fast)
|
||||
|
||||
[[autodoc]] lerobot.policies.pi0_fast.modeling_pi0_fast.PI0FastPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.pi0_fast.configuration_pi0_fast.PI0FastConfig
|
||||
|
||||
## π₀.₅ (PI05)
|
||||
|
||||
[[autodoc]] lerobot.policies.pi05.modeling_pi05.PI05Policy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.pi05.configuration_pi05.PI05Config
|
||||
|
||||
## MolmoAct2
|
||||
|
||||
[[autodoc]] lerobot.policies.molmoact2.modeling_molmoact2.MolmoAct2Policy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.molmoact2.configuration_molmoact2.MolmoAct2Config
|
||||
|
||||
## VLA-JEPA
|
||||
|
||||
[[autodoc]] lerobot.policies.vla_jepa.modeling_vla_jepa.VLAJEPAPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.vla_jepa.configuration_vla_jepa.VLAJEPAConfig
|
||||
|
||||
## EO-1
|
||||
|
||||
[[autodoc]] lerobot.policies.eo1.modeling_eo1.EO1Policy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.eo1.configuration_eo1.EO1Config
|
||||
|
||||
## LingBot-VA
|
||||
|
||||
[[autodoc]] lerobot.policies.lingbot_va.modeling_lingbot_va.LingBotVAPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.lingbot_va.configuration_lingbot_va.LingBotVAConfig
|
||||
|
||||
## FastWAM
|
||||
|
||||
[[autodoc]] lerobot.policies.fastwam.modeling_fastwam.FastWAMPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.fastwam.configuration_fastwam.FastWAMConfig
|
||||
|
||||
## EVO1
|
||||
|
||||
[[autodoc]] lerobot.policies.evo1.modeling_evo1.Evo1Policy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.evo1.configuration_evo1.Evo1Config
|
||||
|
||||
## NVIDIA GR00T
|
||||
|
||||
[[autodoc]] lerobot.policies.groot.modeling_groot.GrootPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.groot.configuration_groot.GrootConfig
|
||||
|
||||
## X-VLA
|
||||
|
||||
[[autodoc]] lerobot.policies.xvla.modeling_xvla.XVLAPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.xvla.configuration_xvla.XVLAConfig
|
||||
|
||||
## Multitask DiT Policy
|
||||
|
||||
[[autodoc]] lerobot.policies.multi_task_dit.modeling_multi_task_dit.MultiTaskDiTPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.multi_task_dit.configuration_multi_task_dit.MultiTaskDiTConfig
|
||||
|
||||
## WALL-OSS
|
||||
|
||||
[[autodoc]] lerobot.policies.wall_x.modeling_wall_x.WallXPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.wall_x.configuration_wall_x.WallXConfig
|
||||
|
||||
## Diffusion Policy
|
||||
|
||||
[[autodoc]] lerobot.policies.diffusion.modeling_diffusion.DiffusionPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.diffusion.configuration_diffusion.DiffusionConfig
|
||||
|
||||
## Gaussian Actor
|
||||
|
||||
[[autodoc]] lerobot.policies.gaussian_actor.modeling_gaussian_actor.GaussianActorPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.gaussian_actor.configuration_gaussian_actor.GaussianActorConfig
|
||||
|
||||
## TD-MPC
|
||||
|
||||
[[autodoc]] lerobot.policies.tdmpc.modeling_tdmpc.TDMPCPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.tdmpc.configuration_tdmpc.TDMPCConfig
|
||||
|
||||
## VQ-BeT
|
||||
|
||||
[[autodoc]] lerobot.policies.vqbet.modeling_vqbet.VQBeTPolicy
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.policies.vqbet.configuration_vqbet.VQBeTConfig
|
||||
@@ -0,0 +1,20 @@
|
||||
# Processors
|
||||
|
||||
Processors are the data transformation layer between a robot, a dataset and a policy. A pipeline is a chain
|
||||
of [`ProcessorStep`]s; each step declares how it transforms both the data and the feature contract.
|
||||
|
||||
See [Introduction to Robot Processors](../introduction_processors) for the concepts,
|
||||
[Implement your own processor](../implement_your_own_processor) to write a step, and
|
||||
[Debug your processor pipeline](../debug_processor_pipeline) when a pipeline misbehaves.
|
||||
|
||||
## ProcessorStep
|
||||
|
||||
[[autodoc]] lerobot.processor.pipeline.ProcessorStep
|
||||
|
||||
## DataProcessorPipeline
|
||||
|
||||
[[autodoc]] lerobot.processor.pipeline.DataProcessorPipeline
|
||||
|
||||
## PolicyProcessorPipeline
|
||||
|
||||
[[autodoc]] lerobot.processor.pipeline.PolicyProcessorPipeline
|
||||
@@ -0,0 +1,147 @@
|
||||
# Robots
|
||||
|
||||
Every robot in LeRobot implements the [`Robot`] interface: connect, read an observation, send an action,
|
||||
disconnect. Writing a policy or a recording script against that interface means it works with any supported
|
||||
arm without change.
|
||||
|
||||
This page is the generated reference. For wiring, calibration and first-run instructions, start with the
|
||||
hardware guides — [SO-101](../so101), [LeKiwi](../lekiwi), [Hope Jr](../hope_jr), [Reachy 2](../reachy2),
|
||||
[OpenArm](../openarm) — or [Imitation Learning for Robots](../il_robots) for the end-to-end workflow. To add
|
||||
a robot of your own, see [Bring Your Own Hardware](../integrate_hardware).
|
||||
|
||||
## Robot
|
||||
|
||||
The abstract base class. Subclasses implement every method below; the contract described here is what a
|
||||
policy or recording loop can rely on.
|
||||
|
||||
[[autodoc]] lerobot.robots.Robot
|
||||
- connect
|
||||
- disconnect
|
||||
- configure
|
||||
- calibrate
|
||||
- get_observation
|
||||
- send_action
|
||||
- observation_features
|
||||
- action_features
|
||||
- is_connected
|
||||
- is_calibrated
|
||||
|
||||
## RobotConfig
|
||||
|
||||
[[autodoc]] lerobot.robots.RobotConfig
|
||||
|
||||
## make_robot_from_config
|
||||
|
||||
[[autodoc]] lerobot.robots.make_robot_from_config
|
||||
|
||||
## SO-100 and SO-101 followers
|
||||
|
||||
`SO100Follower` and `SO101Follower` are aliases of the same `SOFollower` class; the two arms differ in their
|
||||
configuration, not their control code. `SO100FollowerConfig` and `SO101FollowerConfig` are likewise aliases
|
||||
of `SOFollowerRobotConfig`.
|
||||
|
||||
[[autodoc]] lerobot.robots.so_follower.SOFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.so_follower.SOFollowerRobotConfig
|
||||
|
||||
## BiSOFollower
|
||||
|
||||
Two SO followers driven as one bimanual robot.
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_so_follower.BiSOFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_so_follower.BiSOFollowerConfig
|
||||
|
||||
## KochFollower
|
||||
|
||||
[[autodoc]] lerobot.robots.koch_follower.KochFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.koch_follower.KochFollowerConfig
|
||||
|
||||
## LeKiwi
|
||||
|
||||
`LeKiwi` runs on the robot itself. `LeKiwiClient` is the host-side proxy that talks to it over the network
|
||||
and presents the same [`Robot`] interface.
|
||||
|
||||
[[autodoc]] lerobot.robots.lekiwi.LeKiwi
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.lekiwi.LeKiwiConfig
|
||||
|
||||
[[autodoc]] lerobot.robots.lekiwi.LeKiwiClient
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.lekiwi.LeKiwiClientConfig
|
||||
|
||||
## OpenArmFollower
|
||||
|
||||
[[autodoc]] lerobot.robots.openarm_follower.OpenArmFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.openarm_follower.OpenArmFollowerConfig
|
||||
|
||||
## BiOpenArmFollower
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_openarm_follower.BiOpenArmFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_openarm_follower.BiOpenArmFollowerConfig
|
||||
|
||||
## OmxFollower
|
||||
|
||||
[[autodoc]] lerobot.robots.omx_follower.OmxFollower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.omx_follower.OmxFollowerConfig
|
||||
|
||||
## Reachy2Robot
|
||||
|
||||
[[autodoc]] lerobot.robots.reachy2.Reachy2Robot
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.reachy2.Reachy2RobotConfig
|
||||
|
||||
## UnitreeG1
|
||||
|
||||
[[autodoc]] lerobot.robots.unitree_g1.UnitreeG1
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.unitree_g1.UnitreeG1Config
|
||||
|
||||
## Hope Jr
|
||||
|
||||
The Hope Jr humanoid is exposed as two independent robots, an arm and a hand.
|
||||
|
||||
[[autodoc]] lerobot.robots.hope_jr.HopeJrArm
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.hope_jr.HopeJrArmConfig
|
||||
|
||||
[[autodoc]] lerobot.robots.hope_jr.HopeJrHand
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.hope_jr.HopeJrHandConfig
|
||||
|
||||
## RebotB601Follower
|
||||
|
||||
[[autodoc]] lerobot.robots.rebot_b601_follower.RebotB601Follower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.rebot_b601_follower.RebotB601FollowerRobotConfig
|
||||
|
||||
## BiRebotB601Follower
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_rebot_b601_follower.BiRebotB601Follower
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.bi_rebot_b601_follower.BiRebotB601FollowerConfig
|
||||
|
||||
## EarthRoverMiniPlus
|
||||
|
||||
[[autodoc]] lerobot.robots.earthrover_mini_plus.EarthRoverMiniPlus
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.robots.earthrover_mini_plus.EarthRoverMiniPlusConfig
|
||||
@@ -0,0 +1,30 @@
|
||||
# Teleoperators
|
||||
|
||||
A teleoperator produces actions for a robot to follow — a leader arm, a gamepad, a keyboard, a phone. All of
|
||||
them implement the [`Teleoperator`] interface, so a recording script written against it works with any input
|
||||
device.
|
||||
|
||||
See [Phone teleoperation](../phone_teleop) and [Isaac Teleop](../isaac_teleop) for setup guides, and
|
||||
[Imitation Learning for Robots](../il_robots) for the recording workflow.
|
||||
|
||||
## Teleoperator
|
||||
|
||||
[[autodoc]] lerobot.teleoperators.Teleoperator
|
||||
- connect
|
||||
- disconnect
|
||||
- configure
|
||||
- calibrate
|
||||
- get_action
|
||||
- send_feedback
|
||||
- action_features
|
||||
- feedback_features
|
||||
- is_connected
|
||||
- is_calibrated
|
||||
|
||||
## TeleoperatorConfig
|
||||
|
||||
[[autodoc]] lerobot.teleoperators.TeleoperatorConfig
|
||||
|
||||
## make_teleoperator_from_config
|
||||
|
||||
[[autodoc]] lerobot.teleoperators.make_teleoperator_from_config
|
||||
@@ -161,6 +161,16 @@ The methods called by the train/eval loops:
|
||||
|
||||
Batches are flat dictionaries keyed by the constants in [`lerobot.utils.constants`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/utils/constants.py): `OBS_STATE` (`observation.state.<motor>`), `OBS_IMAGES` (`observation.images.<camera>`), `OBS_LANGUAGE`, `ACTION`, etc. Reuse the constants — don't invent new prefixes.
|
||||
|
||||
If your model is large enough to warrant [sharded multi-GPU training](./multi_gpu_training#sharded-training-fsdp), also declare its FSDP wrap units — the repeated block classes sharding operates on:
|
||||
|
||||
```python
|
||||
class MyPolicy(PreTrainedPolicy):
|
||||
...
|
||||
_fsdp_wrap_modules = ["MyTransformerBlock"]
|
||||
```
|
||||
|
||||
With this one declaration, `--parallelism.dp_shard=N` works out of the box for your policy (users can still override it with `--accelerator.fsdp.wrap_modules`). Without any wrap source, sharded runs fail at startup by design.
|
||||
|
||||
### Processor functions
|
||||
|
||||
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).
|
||||
@@ -300,7 +310,7 @@ The file names are load-bearing: the factory does lazy imports by name, and the
|
||||
Two places need to know about your policy. All by name.
|
||||
|
||||
1. **`policies/__init__.py`** — re-export `MyPolicyConfig` and add it to `__all__`. This import is what registers your policy: `@PreTrainedConfig.register_subclass("my_policy")` runs, and from then on the factory resolves everything by convention. **Don't** re-export the modeling class; it loads lazily through the factory (so `import lerobot` stays fast).
|
||||
2. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what `push_model_to_hub` renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page.
|
||||
2. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what the end-of-training publisher renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page.
|
||||
|
||||
Mirror an existing policy that's structurally similar to yours; the diff is small.
|
||||
|
||||
@@ -344,7 +354,7 @@ A new policy is much easier to review — and far more useful — when it ships
|
||||
|
||||
**Pick at least one in-tree benchmark.** LeRobot ships sim benchmarks with per-benchmark Docker images (LIBERO, LIBERO-plus, Meta-World, RoboTwin 2.0, RoboCasa365, RoboCerebra, RoboMME, VLABench and more). Pick the one that matches your policy's modality — VLAs usually go to LIBERO or VLABench; image-only BC to LIBERO or Meta-World. The full list lives under [Benchmarks](./libero) in the docs sidebar.
|
||||
|
||||
**Push the checkpoint & processors** to the Hub under `lerobot/<policy>_<benchmark>` (or your namespace if you don't have write access; a maintainer can mirror it). Use `PreTrainedPolicy.push_model_to_hub` so the repo gets `config.json`, `model.safetensors`, and a model card.
|
||||
**Push the checkpoint & processors** to the Hub under `lerobot/<policy>_<benchmark>` (or your namespace if you don't have write access; a maintainer can mirror it). The easiest way is training with `--policy.repo_id=<namespace>/<repo>` and `--policy.push_to_hub=true`: `lerobot-train` publishes the model, both processors, and a model card at the end of the run. To publish an existing checkpoint after the fact, upload its `pretrained_model/` directory (e.g. `huggingface-cli upload`), or use `lerobot-convert-dcp --push_to_hub=...` for sharded-format checkpoints.
|
||||
|
||||
**Report results in your policy's MDX**, with the exact `lerobot-eval` command and hardware so anyone can re-run:
|
||||
|
||||
|
||||
+218
-11
@@ -23,18 +23,18 @@ The broader EVO1 project may include additional training scripts and dataset too
|
||||
2. Install EVO1 dependencies:
|
||||
|
||||
```bash
|
||||
pip install -e ".[evo1]"
|
||||
pip install -e ".[training,evo1]"
|
||||
```
|
||||
|
||||
For LIBERO evaluation, install the LIBERO extra as well:
|
||||
For LIBERO training and evaluation, install the LIBERO extra as well:
|
||||
|
||||
```bash
|
||||
pip install -e ".[evo1,libero]"
|
||||
pip install -e ".[training,evo1,libero]"
|
||||
```
|
||||
|
||||
3. Install a `flash-attn` wheel only if it is compatible with your Python, PyTorch, CUDA, and GPU stack. EVO1 falls back to standard attention when `flash_attn` is not available.
|
||||
|
||||
EVO1 uses the native Hugging Face `transformers` InternVL implementation, so `policy.vlm_model_name` must point to a natively converted checkpoint such as `OpenGVLab/InternVL3-1B-hf` (note the `-hf` suffix). The first run may download the configured VLM checkpoint unless `policy.vlm_model_name` points to a local model directory.
|
||||
EVO1 uses the native Hugging Face `transformers` InternVL implementation, so `policy.vlm_model_name` must point to a natively converted checkpoint such as `OpenGVLab/InternVL3-1B-hf` (note the `-hf` suffix). The first run downloads the configured VLM checkpoint and later runs reuse it from the Hugging Face cache.
|
||||
|
||||
## Data Requirements
|
||||
|
||||
@@ -92,7 +92,7 @@ lerobot-train \
|
||||
|
||||
### Stage 2
|
||||
|
||||
Stage 2 finetunes the VLM branches and action head. A common workflow starts from a Stage 1 checkpoint:
|
||||
Stage 2 loads the Stage 1 policy, but starts a fresh optimizer and scheduler:
|
||||
|
||||
```bash
|
||||
lerobot-train \
|
||||
@@ -152,16 +152,154 @@ lerobot-rollout \
|
||||
|
||||
### LIBERO Evaluation
|
||||
|
||||
> [!NOTE]
|
||||
> Benchmark results for a `lerobot`-hosted LIBERO checkpoint trained with this implementation
|
||||
> will be added once training completes.
|
||||
#### Reference result
|
||||
|
||||
The official EVO1 LIBERO rollout protocol uses the raw LIBERO camera feature names
|
||||
> [!NOTE]
|
||||
> The released Stage-2 checkpoint passed clean-download and rollout verification:
|
||||
> [`zuoxingdong/evo1_libero`](https://huggingface.co/zuoxingdong/evo1_libero), revision
|
||||
> [`515921f4a2c1d3f3ad523721eafa26fdf2af315b`](https://huggingface.co/zuoxingdong/evo1_libero/commit/515921f4a2c1d3f3ad523721eafa26fdf2af315b).
|
||||
> The clean-download evaluation used LeRobot revision
|
||||
> [`e40b58a8dfa9e7b86918c374791599d070518d11`](https://github.com/huggingface/lerobot/commit/e40b58a8dfa9e7b86918c374791599d070518d11).
|
||||
|
||||
The single-run Stage-2 checkpoint at step 70,000 produced:
|
||||
|
||||
| Suite | Successful episodes | Episodes | Success rate |
|
||||
| -------------- | ------------------: | --------: | -----------: |
|
||||
| LIBERO Spatial | 485 | 500 | 97.0% |
|
||||
| LIBERO Object | 496 | 500 | 99.2% |
|
||||
| LIBERO Goal | 483 | 500 | 96.6% |
|
||||
| LIBERO-10 | 469 | 500 | 93.8% |
|
||||
| **Overall** | **1,933** | **2,000** | **96.65%** |
|
||||
|
||||
These results use one trained checkpoint and evaluation seed `1000`; they are not a multi-seed
|
||||
mean or confidence estimate.
|
||||
|
||||
#### Reference training recipe
|
||||
|
||||
The released checkpoint records the complete resolved Stage-2 configuration in
|
||||
[`train_config.json`](https://huggingface.co/zuoxingdong/evo1_libero/blob/515921f4a2c1d3f3ad523721eafa26fdf2af315b/train_config.json).
|
||||
The measured run used two H100 GPUs with two DDP processes and batch 64 per process, giving global batch 128. Both stages used the same topology. The base VLM came from revision
|
||||
`014c0583a0d4bedf29fbe2dbff4f865eb998e171` of `OpenGVLab/InternVL3-1B-hf`.
|
||||
The released artifact does not record the exact LeRobot training commit or its original dependency lock,
|
||||
so the commands below reproduce the recorded configuration and topology from a current checkout rather
|
||||
than reconstructing the software environment bit for bit.
|
||||
|
||||
From a LeRobot source checkout, install the locked dependencies and download that exact VLM revision:
|
||||
|
||||
```bash
|
||||
uv sync --locked --extra training --extra evo1 --extra libero
|
||||
VLM_DIR=$(uv run hf download OpenGVLab/InternVL3-1B-hf \
|
||||
--revision=014c0583a0d4bedf29fbe2dbff4f865eb998e171)
|
||||
```
|
||||
|
||||
Stage 1 freezes the VLM and trains the action head for 5,000 steps:
|
||||
|
||||
```bash
|
||||
uv run accelerate launch --num_processes=2 -m lerobot.scripts.lerobot_train \
|
||||
--dataset.repo_id=lerobot/libero \
|
||||
--dataset.revision=a1aaacb7f6cd6ee5fb43120f673cebb0cfea7dd4 \
|
||||
--dataset.video_backend=torchcodec \
|
||||
--dataset.return_uint8=true \
|
||||
--dataset.image_transforms.enable=true \
|
||||
--dataset.use_imagenet_stats=true \
|
||||
--dataset.eval_split=0.0 \
|
||||
--policy.type=evo1 \
|
||||
--policy.training_stage=stage1 \
|
||||
--policy.apply_training_stage_defaults=true \
|
||||
--policy.vlm_model_name="${VLM_DIR}" \
|
||||
--policy.vlm_num_layers=14 \
|
||||
--policy.vlm_dtype=bfloat16 \
|
||||
--policy.device=cuda \
|
||||
--policy.use_amp=true \
|
||||
--policy.use_flash_attn=true \
|
||||
--policy.enable_gradient_checkpointing=true \
|
||||
--policy.gradient_checkpointing_use_reentrant=false \
|
||||
--policy.image_resolution='[448,448]' \
|
||||
--policy.chunk_size=50 \
|
||||
--policy.n_action_steps=50 \
|
||||
--policy.max_state_dim=24 \
|
||||
--policy.max_action_dim=24 \
|
||||
--policy.dropout=0.2 \
|
||||
--policy.optimizer_lr=1e-5 \
|
||||
--policy.optimizer_weight_decay=1e-3 \
|
||||
--policy.optimizer_grad_clip_norm=1.0 \
|
||||
--policy.scheduler_warmup_steps=1000 \
|
||||
--policy.push_to_hub=false \
|
||||
--use_policy_training_preset=true \
|
||||
--batch_size=64 \
|
||||
--steps=5000 \
|
||||
--save_checkpoint=true \
|
||||
--save_checkpoint_to_hub=false \
|
||||
--save_freq=2500 \
|
||||
--log_freq=10 \
|
||||
--env_eval_freq=0 \
|
||||
--num_workers=4 \
|
||||
--prefetch_factor=2 \
|
||||
--persistent_workers=true \
|
||||
--seed=1000 \
|
||||
--wandb.enable=false \
|
||||
--output_dir=./outputs/evo1-libero-stage1-g128-5k
|
||||
```
|
||||
|
||||
Stage 2 loads the Stage-1 policy but starts a fresh optimizer and scheduler. It trains for 80,000 steps;
|
||||
the reported checkpoint is the save at step 70,000:
|
||||
|
||||
```bash
|
||||
uv run accelerate launch --num_processes=2 -m lerobot.scripts.lerobot_train \
|
||||
--dataset.repo_id=lerobot/libero \
|
||||
--dataset.revision=a1aaacb7f6cd6ee5fb43120f673cebb0cfea7dd4 \
|
||||
--dataset.video_backend=torchcodec \
|
||||
--dataset.return_uint8=true \
|
||||
--dataset.image_transforms.enable=true \
|
||||
--dataset.use_imagenet_stats=true \
|
||||
--dataset.eval_split=0.0 \
|
||||
--policy.path=./outputs/evo1-libero-stage1-g128-5k/checkpoints/005000/pretrained_model \
|
||||
--policy.training_stage=stage2 \
|
||||
--policy.apply_training_stage_defaults=true \
|
||||
--policy.vlm_model_name="${VLM_DIR}" \
|
||||
--policy.vlm_num_layers=14 \
|
||||
--policy.vlm_dtype=float32 \
|
||||
--policy.device=cuda \
|
||||
--policy.use_amp=true \
|
||||
--policy.use_flash_attn=true \
|
||||
--policy.enable_gradient_checkpointing=true \
|
||||
--policy.gradient_checkpointing_use_reentrant=false \
|
||||
--policy.image_resolution='[448,448]' \
|
||||
--policy.chunk_size=50 \
|
||||
--policy.n_action_steps=50 \
|
||||
--policy.max_state_dim=24 \
|
||||
--policy.max_action_dim=24 \
|
||||
--policy.dropout=0.2 \
|
||||
--policy.optimizer_lr=1e-5 \
|
||||
--policy.optimizer_weight_decay=1e-3 \
|
||||
--policy.optimizer_grad_clip_norm=1.0 \
|
||||
--policy.scheduler_warmup_steps=1000 \
|
||||
--policy.push_to_hub=false \
|
||||
--use_policy_training_preset=true \
|
||||
--batch_size=64 \
|
||||
--steps=80000 \
|
||||
--resume=false \
|
||||
--save_checkpoint=true \
|
||||
--save_checkpoint_to_hub=false \
|
||||
--save_freq=10000 \
|
||||
--log_freq=10 \
|
||||
--env_eval_freq=0 \
|
||||
--num_workers=4 \
|
||||
--prefetch_factor=2 \
|
||||
--persistent_workers=true \
|
||||
--seed=1000 \
|
||||
--wandb.enable=false \
|
||||
--output_dir=./outputs/evo1-libero-stage2-g128-80k
|
||||
```
|
||||
|
||||
#### Author-format evaluation profile
|
||||
|
||||
The author-format EVO1 LIBERO profile uses the raw LIBERO camera feature names
|
||||
(`observation.images.agentview_image` and `observation.images.robot0_eye_in_hand_image`), replans every
|
||||
14 actions, and binarizes the gripper command before stepping the simulator. The EVO1 policy postprocessor
|
||||
can crop the padded 24D action back to the 7D LIBERO action space and apply that gripper binarization. To
|
||||
evaluate a LIBERO checkpoint under the same one-episode-per-task setting, keep the raw camera names instead
|
||||
of the default `image`/`image2` mapping and set the LIBERO action postprocessing flags:
|
||||
evaluate an author-format checkpoint under the same one-episode-per-task setting, keep the raw camera names
|
||||
instead of the default `image`/`image2` mapping and set the LIBERO action postprocessing flags:
|
||||
|
||||
```bash
|
||||
lerobot-eval \
|
||||
@@ -181,6 +319,75 @@ lerobot-eval \
|
||||
--eval.n_episodes=1
|
||||
```
|
||||
|
||||
#### Native `lerobot/libero` v3 profile
|
||||
|
||||
Revision `a1aaacb7f6cd6ee5fb43120f673cebb0cfea7dd4` stores camera features as `image` and
|
||||
`image2`. This example evaluates all ten LIBERO Object tasks, launching each task in a fresh process:
|
||||
|
||||
```bash
|
||||
export MUJOCO_GL=egl
|
||||
export PYOPENGL_PLATFORM=egl
|
||||
|
||||
suite=libero_object
|
||||
horizon=280
|
||||
for task_id in {0..9}; do
|
||||
lerobot-eval \
|
||||
--policy.path=zuoxingdong/evo1_libero \
|
||||
--policy.pretrained_revision=515921f4a2c1d3f3ad523721eafa26fdf2af315b \
|
||||
--policy.vlm_model_name=OpenGVLab/InternVL3-1B-hf \
|
||||
--policy.device=cuda \
|
||||
--policy.use_amp=true \
|
||||
--policy.vlm_dtype=bfloat16 \
|
||||
--policy.use_flash_attn=false \
|
||||
--policy.enable_gradient_checkpointing=false \
|
||||
--policy.vlm_num_layers=14 \
|
||||
--policy.image_resolution='[448,448]' \
|
||||
--policy.max_text_length=1024 \
|
||||
--policy.chunk_size=50 \
|
||||
--policy.n_action_steps=14 \
|
||||
--policy.max_state_dim=24 \
|
||||
--policy.max_action_dim=24 \
|
||||
--policy.num_inference_timesteps=32 \
|
||||
--policy.postprocess_action_dim=7 \
|
||||
--policy.binarize_gripper=true \
|
||||
--policy.gripper_threshold=0.0 \
|
||||
--policy.gripper_below_threshold_value=-1.0 \
|
||||
--policy.gripper_above_threshold_value=1.0 \
|
||||
--env.type=libero \
|
||||
--env.task="${suite}" \
|
||||
--env.task_ids="[${task_id}]" \
|
||||
--env.camera_name=agentview_image,robot0_eye_in_hand_image \
|
||||
--env.camera_name_mapping="{agentview_image: image, robot0_eye_in_hand_image: image2}" \
|
||||
--env.control_mode=relative \
|
||||
--env.obs_type=pixels_agent_pos \
|
||||
--env.observation_width=448 \
|
||||
--env.observation_height=448 \
|
||||
--env.init_states=true \
|
||||
--env.episode_length="${horizon}" \
|
||||
--env.render_mode=rgb_array \
|
||||
--env.max_parallel_tasks=1 \
|
||||
--eval.n_episodes=50 \
|
||||
--eval.batch_size=1 \
|
||||
--eval.use_async_envs=false \
|
||||
--eval.recording=false \
|
||||
--seed=1000 \
|
||||
--output_dir="./outputs/evo1-libero-stage2-70k-eval/${suite}/task-${task_id}" \
|
||||
--job_name="evo1-libero-stage2-70k-${suite}-task-${task_id}"
|
||||
done
|
||||
```
|
||||
|
||||
Run all ten task IDs for each suite with these horizons:
|
||||
|
||||
| `env.task` | `env.episode_length` |
|
||||
| ---------------- | -------------------: |
|
||||
| `libero_spatial` | `280` |
|
||||
| `libero_object` | `280` |
|
||||
| `libero_goal` | `300` |
|
||||
| `libero_10` | `520` |
|
||||
|
||||
Set `suite` and `horizon` for each row. This gives 500 episodes per suite and 2,000 episodes overall, while
|
||||
the loop's fresh process per task matches the measured RNG-reset topology.
|
||||
|
||||
## References
|
||||
|
||||
- [EVO1 repository](https://github.com/MINT-SJTU/Evo-1)
|
||||
|
||||
@@ -30,7 +30,7 @@ The goal: lower the barrier to entry for robotics, so that everyone can contribu
|
||||
</div>
|
||||
|
||||
<div align="center">
|
||||
<img src="../../media/readme/robots_control_video.webp" width="640px" alt="Reachy 2 Demo">
|
||||
<img src="https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/lerobot/robots_control_video.webp" width="640px" alt="Reachy 2 Demo">
|
||||
</div>
|
||||
|
||||
## How It Works
|
||||
|
||||
@@ -108,6 +108,7 @@ own binding plus a matching image block, e.g.
|
||||
|
||||
```yaml
|
||||
ask_vqa_top:
|
||||
route: vqa
|
||||
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)"
|
||||
@@ -127,7 +128,9 @@ ask_vqa_top:
|
||||
}
|
||||
```
|
||||
|
||||
Add one such sub-recipe per camera the dataset records.
|
||||
Add one such sub-recipe per camera the dataset records. The explicit
|
||||
`route: vqa` marker makes a matching sparse VQA annotation take precedence
|
||||
over normal weighted blend selection; component names are purely descriptive.
|
||||
|
||||
## Layer 3 — training format
|
||||
|
||||
@@ -141,7 +144,20 @@ 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.
|
||||
|
||||
## 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.
|
||||
|
||||
`recipes/subtask_joint.yaml` demonstrates joint sequence training rather than a
|
||||
weighted blend. For the same sample, its assistant subtask is supervised with
|
||||
text cross-entropy on the `low_level` stream while action prediction remains
|
||||
active, matching the joint setup from the π0.5 paper. Enable
|
||||
`--policy.joint_subtask_conditioning=true` to use that subtask conditioning at inference.
|
||||
|
||||
## Graceful absence
|
||||
|
||||
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
||||
If an event-scoped branch is selected on a frame without the required event row, rendering returns `None`, allowing a loader to retry another sample.
|
||||
If both language columns are missing, `None`, or empty, `RenderMessagesStep` uses
|
||||
the task string as low-level supervision when available and otherwise leaves the
|
||||
sample unchanged. For an annotated sample, if no recipe branch applies and no
|
||||
task fallback exists, rendering returns `None`, allowing a loader to retry another sample.
|
||||
|
||||
@@ -142,6 +142,22 @@ repo_id = "yaak-ai/L2D-v3"
|
||||
dataset = StreamingLeRobotDataset(repo_id) # streams directly from the Hub
|
||||
```
|
||||
|
||||
Datasets stored in an [HF Storage Bucket](https://huggingface.co/docs/hub/storage-buckets) (`hf://buckets/`) can be streamed the same way by passing `repo_type="bucket"`:
|
||||
|
||||
```python
|
||||
dataset = StreamingLeRobotDataset("my-org/my-bucket", repo_type="bucket")
|
||||
```
|
||||
|
||||
Both options are available in `lerobot-train` through `--dataset.streaming=true`, and `--dataset.repo_type=bucket` to stream from a bucket instead of a Hub dataset repo:
|
||||
|
||||
```bash
|
||||
lerobot-train \
|
||||
--dataset.repo_id=my-org/my-bucket \
|
||||
--dataset.repo_type=bucket \
|
||||
--dataset.streaming=true \
|
||||
...
|
||||
```
|
||||
|
||||
<div style="display:flex; justify-content:center; gap:12px; flex-wrap:wrap;">
|
||||
<figure style="margin:0; text-align:center;">
|
||||
<img
|
||||
|
||||
@@ -92,6 +92,20 @@ LIBERO supports two control modes — `relative` (default) and `absolute`. Diffe
|
||||
--env.control_mode=relative # or "absolute"
|
||||
```
|
||||
|
||||
### Reset performance
|
||||
|
||||
By default, LeRobot preserves LIBERO's hard-reset behavior. With fixed initial
|
||||
states enabled, you can opt into soft resets to skip rebuilding the simulator
|
||||
model and renderer on every episode:
|
||||
|
||||
```bash
|
||||
--env.init_states=true --env.hard_reset=false
|
||||
```
|
||||
|
||||
Soft resets are faster but are not bit-identical to hard resets after the
|
||||
environment's settling steps, so camera observations and policy results may
|
||||
differ slightly. Use hard resets when reproducing benchmark results.
|
||||
|
||||
### Policy inputs and outputs
|
||||
|
||||
**Observations:**
|
||||
|
||||
@@ -134,6 +134,20 @@ LIBERO-plus supports two control modes — `relative` (default) and `absolute`.
|
||||
--env.control_mode=relative # or "absolute"
|
||||
```
|
||||
|
||||
### Reset performance
|
||||
|
||||
By default, LeRobot preserves LIBERO's hard-reset behavior. With fixed initial
|
||||
states enabled, you can opt into soft resets to skip rebuilding the simulator
|
||||
model and renderer on every episode:
|
||||
|
||||
```bash
|
||||
--env.init_states=true --env.hard_reset=false
|
||||
```
|
||||
|
||||
Soft resets are faster but are not bit-identical to hard resets after the
|
||||
environment's settling steps, so camera observations and policy results may
|
||||
differ slightly. Use hard resets when reproducing benchmark results.
|
||||
|
||||
### Policy inputs and outputs
|
||||
|
||||
**Observations:**
|
||||
|
||||
+114
-118
@@ -1,28 +1,29 @@
|
||||
# Multi-GPU Training
|
||||
|
||||
This guide shows you how to train policies on multiple GPUs using [Hugging Face Accelerate](https://huggingface.co/docs/accelerate).
|
||||
LeRobot trains on multiple GPUs through [Hugging Face Accelerate](https://huggingface.co/docs/accelerate). Three data-parallel layouts are supported:
|
||||
|
||||
| Layout | What it does | Config |
|
||||
| -------- | ------------------------------------------------------------- | ------------------------------------------------------- |
|
||||
| **DDP** | Replicates the full model on every GPU | default on any multi-GPU launch |
|
||||
| **FSDP** | Shards parameters, gradients, and optimizer state across GPUs | `--parallelism.dp_shard=N` |
|
||||
| **HSDP** | Shards within groups of GPUs, replicates across groups | `--parallelism.dp_replicate=R --parallelism.dp_shard=S` |
|
||||
|
||||
## Installation
|
||||
|
||||
`accelerate` is included in the `training` extra. Install it with:
|
||||
`accelerate` is included in the `training` extra:
|
||||
|
||||
```bash
|
||||
pip install 'lerobot[training]'
|
||||
```
|
||||
|
||||
## Training with Multiple GPUs
|
||||
## Launching
|
||||
|
||||
You can launch training in two ways:
|
||||
Distributed training can be launched through both `torchrun` and `accelerate launch`. Accelerate is used as a plain launcher: it does not manage the training configuration, and every distributed training setting lives in LeRobot's own config system.
|
||||
|
||||
### Option 1: Without config (specify parameters directly)
|
||||
|
||||
You can specify all parameters directly in the command without running `accelerate config`:
|
||||
With `torchrun`:
|
||||
|
||||
```bash
|
||||
accelerate launch \
|
||||
--multi_gpu \
|
||||
--num_processes=2 \
|
||||
$(which lerobot-train) \
|
||||
torchrun --nproc-per-node=2 $(which lerobot-train) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=act \
|
||||
--policy.repo_id=${HF_USER}/my_trained_policy \
|
||||
@@ -31,32 +32,10 @@ accelerate launch \
|
||||
--wandb.enable=true
|
||||
```
|
||||
|
||||
**Key accelerate parameters:**
|
||||
|
||||
- `--multi_gpu`: Enable multi-GPU training
|
||||
- `--num_processes=2`: Number of GPUs to use
|
||||
- `--mixed_precision=fp16`: Use fp16 mixed precision (or `bf16` if supported)
|
||||
|
||||
### Option 2: Using accelerate config
|
||||
|
||||
If you prefer to save your configuration, you can optionally configure accelerate for your hardware setup by running:
|
||||
With `accelerate launch` (as a plain launcher):
|
||||
|
||||
```bash
|
||||
accelerate config
|
||||
```
|
||||
|
||||
This interactive setup will ask you questions about your training environment (number of GPUs, mixed precision settings, etc.) and saves the configuration for future use. For a simple multi-GPU setup on a single machine, you can use these recommended settings:
|
||||
|
||||
- Compute environment: This machine
|
||||
- Number of machines: 1
|
||||
- Number of processes: (number of GPUs you want to use)
|
||||
- GPU ids to use: (leave empty to use all)
|
||||
- Mixed precision: fp16 or bf16 (recommended for faster training)
|
||||
|
||||
Then launch training with:
|
||||
|
||||
```bash
|
||||
accelerate launch $(which lerobot-train) \
|
||||
accelerate launch --num_processes=2 $(which lerobot-train) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=act \
|
||||
--policy.repo_id=${HF_USER}/my_trained_policy \
|
||||
@@ -65,116 +44,133 @@ accelerate launch $(which lerobot-train) \
|
||||
--wandb.enable=true
|
||||
```
|
||||
|
||||
## How It Works
|
||||
With no `--parallelism.*` flags, a multi-process launch runs plain DDP. Multi-node runs use the standard `torchrun --nnodes/--node-rank/--rdzv-endpoint` flags (or `accelerate launch --num_machines/--machine_rank/--main_process_ip`).
|
||||
|
||||
When you launch training with accelerate:
|
||||
> [!WARNING]
|
||||
> Accelerate's YAML config files (`accelerate launch --config_file some.yaml`, `accelerate config`) are not supported. They configure the engine through environment variables, bypassing LeRobot's configuration system, so `train_config.json` would no longer describe the settings a run actually used. `lerobot-train` therefore refuses to start when [accelerate environment variables](https://huggingface.co/docs/accelerate/usage_guides/fsdp) are set. Put the settings in `--parallelism.*` / `--accelerator.*` flags instead, or set `LEROBOT_ALLOW_ACCELERATE_ENV=1` to acknowledge the override and proceed anyway.
|
||||
|
||||
1. **Automatic detection**: LeRobot automatically detects if it's running under accelerate
|
||||
2. **Data distribution**: Your batch is automatically split across GPUs
|
||||
3. **Gradient synchronization**: Gradients are synchronized across GPUs during backpropagation
|
||||
4. **Single process logging**: Only the main process logs to wandb and saves checkpoints
|
||||
## Batch semantics, learning rate, and steps
|
||||
|
||||
## Learning Rate and Training Steps Scaling
|
||||
Each of the `dp_replicate × dp_shard` data-parallel workers loads its own `--batch_size` micro-batch every step, so one training step consumes `batch_size × dp_world_size` samples, and `× gradient_accumulation_steps` of those go into each optimizer update:
|
||||
|
||||
**Important:** LeRobot does **NOT** automatically scale learning rates or training steps based on the number of GPUs. This gives you full control over your training hyperparameters.
|
||||
|
||||
### Why No Automatic Scaling?
|
||||
|
||||
Many distributed training frameworks automatically scale the learning rate by the number of GPUs (e.g., `lr = base_lr × num_gpus`).
|
||||
However, LeRobot keeps the learning rate exactly as you specify it.
|
||||
|
||||
### When and How to Scale
|
||||
|
||||
If you want to scale your hyperparameters when using multiple GPUs, you should do it manually:
|
||||
|
||||
**Learning Rate Scaling:**
|
||||
|
||||
```bash
|
||||
# Example: 2 GPUs with linear LR scaling
|
||||
# Base LR: 1e-4, with 2 GPUs -> 2e-4
|
||||
accelerate launch --num_processes=2 $(which lerobot-train) \
|
||||
--optimizer.lr=2e-4 \
|
||||
--dataset.repo_id=lerobot/pusht \
|
||||
--policy.type=act
|
||||
```
|
||||
effective_batch_size = batch_size × dp_world_size × gradient_accumulation_steps
|
||||
```
|
||||
|
||||
**Training Steps Scaling:**
|
||||
The training banner prints this factorization at startup. `--steps` counts loop steps (micro-batches per worker), not optimizer updates.
|
||||
|
||||
Since the effective batch size `bs` increases with multiple GPUs (batch_size × num_gpus), you may want to reduce the number of training steps proportionally:
|
||||
Gradient accumulation is a first-class flag:
|
||||
|
||||
```bash
|
||||
# Example: 2 GPUs with effective batch size 2x larger
|
||||
# Original: batch_size=8, steps=100000
|
||||
# With 2 GPUs: batch_size=8 (16 in total), steps=50000
|
||||
accelerate launch --num_processes=2 $(which lerobot-train) \
|
||||
--batch_size=8 \
|
||||
--steps=50000 \
|
||||
--dataset.repo_id=lerobot/pusht \
|
||||
--policy.type=act
|
||||
torchrun --nproc-per-node=2 $(which lerobot-train) \
|
||||
--batch_size=8 --accelerator.gradient_accumulation.steps=4 ...
|
||||
```
|
||||
|
||||
## Training Large Models with FSDP
|
||||
**LeRobot does not auto-scale the learning rate or the number of steps** when the effective batch size grows. If you scale out and want equivalent training, please adjust manually, e.g. with 2 GPUs: double `--optimizer.lr` (linear scaling), or halve `--steps`.
|
||||
|
||||
DDP replicates the full model on every GPU, so a model that doesn't fit on one GPU won't fit under
|
||||
DDP either. For large models, use **FSDP** (Fully Sharded Data Parallel), which shards parameters,
|
||||
gradients, and optimizer state across GPUs. See the [accelerate FSDP guide](https://huggingface.co/docs/accelerate/usage_guides/fsdp) for background.
|
||||
## Sharded training (FSDP)
|
||||
|
||||
An example on how to launch LeRobot training with FSDP across 4 GPUs (1 machine):
|
||||
If a model is too large to train with DDP, shard it with FSDP2:
|
||||
|
||||
```bash
|
||||
accelerate launch --config_file fsdp.yaml --num_processes=4 $(which lerobot-train) \
|
||||
torchrun --nproc-per-node=4 $(which lerobot-train) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=<your_policy> \
|
||||
--parallelism.dp_shard=4 \
|
||||
--accelerator.mixed_precision=bf16 \
|
||||
--output_dir=outputs/train/my_policy_fsdp
|
||||
```
|
||||
|
||||
A minimal `fsdp.yaml` (FSDP1; shards params/grads/optimizer — ZeRO-3-equivalent):
|
||||
`--parallelism.dp_shard=-1` shards over however many processes the launcher started.
|
||||
|
||||
```yaml
|
||||
compute_environment: LOCAL_MACHINE
|
||||
distributed_type: FSDP
|
||||
mixed_precision: bf16
|
||||
num_machines: 1
|
||||
num_processes: 4
|
||||
fsdp_config:
|
||||
fsdp_version: 1
|
||||
fsdp_sharding_strategy: FULL_SHARD # params + grads + optimizer (ZeRO-3)
|
||||
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
|
||||
fsdp_transformer_layer_cls_to_wrap: <YourTransformerBlock> # repeated block class to shard
|
||||
fsdp_use_orig_params: true # required: optimizer is built pre-prepare
|
||||
fsdp_state_dict_type: FULL_STATE_DICT
|
||||
### Wrap units
|
||||
|
||||
FSDP shards the model in units (typically the repeated transformer block) and gathers one unit at a time during forward/backward. Policies declare their wrap units via `_fsdp_wrap_modules` on the policy class. For example, ACT declares `["ACTEncoderLayer", "ACTDecoderLayer"]` and FastWAM declares `["MoTLayer"]`. For a policy without a `_fsdp_wrap_modules` declaration, pass one of the flags below. You can specify the module class name explicitly, or use a size-based policy instead:
|
||||
|
||||
```bash
|
||||
--accelerator.fsdp.wrap_modules='["MyTransformerBlock"]' # explicit class names
|
||||
--accelerator.fsdp.min_num_params=1000000 # or: wrap every submodule above 1M params
|
||||
```
|
||||
|
||||
Set `fsdp_transformer_layer_cls_to_wrap` to your model's repeated transformer-block class so each
|
||||
block is sharded as its own unit. `fsdp_use_orig_params: true` is required because LeRobot builds the
|
||||
optimizer before `accelerator.prepare()`.
|
||||
If a policy doesn't declare `_fsdp_wrap_modules` and no `--accelerator.fsdp.wrap_modules` or `--accelerator.fsdp.min_num_params` is passed, the run fails at startup rather than silently wrapping only the root module (which would forfeit all sharding memory savings).
|
||||
|
||||
### FSDP checkpoints
|
||||
Other sharding settings:
|
||||
|
||||
LeRobot gathers the full state dict across all ranks and the main process writes it as a single
|
||||
`model.safetensors`, loadable as usual with `Policy.from_pretrained(...)`. Two things to look out for:
|
||||
- `--accelerator.fsdp.reshard_after_forward`: whether to keep each unit's parameters resident after forward.
|
||||
- `--accelerator.fsdp.cpu_offload`: keeps parameters, gradients and optimizer states on CPU.
|
||||
- `--accelerator.fsdp.ignored_modules`: a regex of module paths to keep unsharded.
|
||||
|
||||
- **Checkpoints store fp32 weights.** Under mixed precision (`bf16`/`fp16`) FSDP keeps an fp32 master
|
||||
copy, and the checkpoint saves it (~2× the bf16 size on disk) so training can resume consistently
|
||||
with the fp32 optimizer state; `from_pretrained` casts back to the policy dtype on load. FSDP-specific
|
||||
caveat: an fp32 checkpoint is materialized in full precision on the target device _before_ casting,
|
||||
so loading it for inference on a tight GPU can OOM even when the bf16 model would fit — load on CPU
|
||||
first, or cast `model.safetensors` to the deployment dtype offline.
|
||||
- The sharded optimizer state is gathered into a full (world-size-independent) state dict and saved
|
||||
alongside the model in the same `optimizer_state.safetensors` / `optimizer_param_groups.json`
|
||||
format as single-GPU training, so **resume-from-checkpoint is supported** with `--resume=true`.
|
||||
Resume reshards both the model and the optimizer state to the _current_ FSDP topology, so you can
|
||||
resume an FSDP checkpoint on a different number of GPUs. Note that the data sampler is only
|
||||
sample-exact when the world size and batch size match the original run (a warning is logged
|
||||
otherwise); the optimizer/model state itself is unaffected.
|
||||
### HSDP
|
||||
|
||||
Hybrid Sharded Data Parallel: parameters, gradients and optimizer states are sharded across `dp_shard` ranks, and that sharding is replicated `dp_replicate` times. Parameter all-gathers and gradient reduce-scatters stay inside a shard group; only the all-reduce that synchronizes the replicas crosses between groups. The two degrees must multiply to the world size:
|
||||
|
||||
```bash
|
||||
# 16 GPUs = 2 nodes × 8: shard within each node, replicate across nodes
|
||||
torchrun --nnodes=2 --nproc-per-node=8 ... $(which lerobot-train) \
|
||||
--parallelism.dp_replicate=2 --parallelism.dp_shard=8 ...
|
||||
```
|
||||
|
||||
## Checkpoints
|
||||
|
||||
Every checkpoint contains a `pretrained_model/` directory and a `training_state/` directory:
|
||||
|
||||
```text
|
||||
005000/ # the training step at that checkpoint
|
||||
├── pretrained_model/
|
||||
│ ├── config.json # policy config
|
||||
│ ├── train_config.json # the full training config
|
||||
│ ├── model.safetensors # full weights (checkpoint_format ∈ {safetensors, safetensors_dcp}, or any non-sharded run)
|
||||
│ ├── pytorch_model_fsdp_0/ # DCP weight shards (checkpoint_format ∈ {dcp, safetensors_dcp})
|
||||
│ ├── policy_preprocessor.json # preprocessor config (when the run has a preprocessor)
|
||||
│ ├── policy_preprocessor_step_*.safetensors # state of the stateful preprocessor steps
|
||||
│ ├── policy_postprocessor.json # postprocessor config (when the run has a postprocessor)
|
||||
│ └── policy_postprocessor_step_*.safetensors # state of the stateful postprocessor steps
|
||||
└── training_state/
|
||||
├── training_step.json # step counter, topology, and batch semantics
|
||||
├── rng_state.safetensors # rng states
|
||||
├── scheduler_state.json # scheduler state (when the run has a scheduler)
|
||||
├── optimizer_state.safetensors # full optimizer state (non-sharded runs)
|
||||
├── optimizer_param_groups.json # optimizer param groups (non-sharded runs)
|
||||
└── optimizer_0/ # DCP optimizer shards (sharded runs)
|
||||
```
|
||||
|
||||
During single-GPU or DDP training, the pipeline serializes each state dict into a single file: `model.safetensors` for the model and `optimizer_state.safetensors` for the optimizer.
|
||||
|
||||
During sharded training, the optimizer state is saved as DCP shards under `training_state/optimizer_0/`, and the layout of the model under `pretrained_model/` can be configured through `--checkpoint_format`:
|
||||
|
||||
| `--checkpoint_format` | Weights artifact | Use when |
|
||||
| ------------------------- | -------------------------------------------- | --------------------------------------------------------------------- |
|
||||
| `safetensors` _(default)_ | single `model.safetensors` only | you want every checkpoint immediately loadable with `from_pretrained` |
|
||||
| `dcp` | `pytorch_model_fsdp_0/` shard directory only | gathering the full weights makes saves and resumes too slow |
|
||||
| `safetensors_dcp` | both | you want fast resume _and_ immediately loadable checkpoints |
|
||||
|
||||
Two things to know about gathered (`safetensors`) checkpoints from sharded runs:
|
||||
|
||||
- **They store fp32 weights.** Under mixed precision training, FSDP keeps an fp32 master copy, and the checkpoint saves the master copy to make sure training resumes consistently.
|
||||
- The gather is collective (all ranks participate) but only the main process writes.
|
||||
|
||||
### Converting DCP checkpoints
|
||||
|
||||
`lerobot-convert-dcp` merges a DCP shard directory into a regular `model.safetensors`, offline and without GPUs:
|
||||
|
||||
```bash
|
||||
lerobot-convert-dcp --checkpoint_dir=outputs/train/run/checkpoints/005000
|
||||
lerobot-convert-dcp --checkpoint_dir=... --delete_dcp=true --push_to_hub=${HF_USER}/my_policy
|
||||
```
|
||||
|
||||
`--push_to_hub` publishes the converted directory as a model repo.
|
||||
|
||||
### Resuming
|
||||
|
||||
Resume with `--resume=true --config_path=.../checkpoints/last/pretrained_model/train_config.json`. Resuming from a DCP checkpoint supports resharding the model and optimizer state to the _current_ topology, which means you can resume with a different `dp_replicate/dp_shard` split. The data sampler can always resume at the right epoch and offset, but is only _sample-exact_ when the world size and batch size match the original run (a warning is logged otherwise).
|
||||
|
||||
> [!NOTE]
|
||||
> FSDP checkpoints written by LeRobot 0.6.x and earlier used a different on-disk layout (a gathered full optimizer state) and **cannot be resumed**.
|
||||
|
||||
## Notes
|
||||
|
||||
- The `--policy.use_amp` flag in `lerobot-train` is only used when **not** running with accelerate. When using accelerate, mixed precision is controlled by accelerate's configuration.
|
||||
- Training logs, checkpoints, and hub uploads are only done by the main process to avoid conflicts. Non-main processes have console logging disabled to prevent duplicate output.
|
||||
- The effective batch size is `batch_size × num_gpus`. If you use 4 GPUs with `--batch_size=8`, your effective batch size is 32.
|
||||
- Learning rate scheduling is handled correctly across multiple processes—LeRobot sets `step_scheduler_with_optimizer=False` to prevent accelerate from adjusting scheduler steps based on the number of processes.
|
||||
- When saving or pushing models, LeRobot automatically unwraps the model from accelerate's distributed wrapper to ensure compatibility.
|
||||
- WandB integration automatically initializes only on the main process, preventing multiple runs from being created.
|
||||
- Checkpoint saves and end-of-training publishes are collective (every rank enters them). Gathered weights, sidecar files and Hub uploads are written by the main process alone.
|
||||
- Metrics are reduced across ranks before logging: losses are averaged, and `samples/s` reports cluster-wide throughput.
|
||||
- Learning-rate scheduling is stepped once per training step regardless of the number of processes (`step_scheduler_with_optimizer=False` is baked in).
|
||||
|
||||
For more advanced configurations and troubleshooting, see the [Accelerate documentation](https://huggingface.co/docs/accelerate). If you want to learn more about how to train on a large number of GPUs, checkout this awesome guide: [Ultrascale Playbook](https://huggingface.co/spaces/nanotron/ultrascale-playbook).
|
||||
For background on the underlying machinery, see the [Accelerate FSDP guide](https://huggingface.co/docs/accelerate/usage_guides/fsdp). To go deeper on large-scale training, check out the [Ultrascale Playbook](https://huggingface.co/spaces/nanotron/ultrascale-playbook).
|
||||
|
||||
@@ -0,0 +1,339 @@
|
||||
# Third-Party Robots & Teleoperators
|
||||
|
||||
The LeRobot ecosystem extends far beyond its officially supported hardware. Thanks to LeRobot's plugin architecture, the community has built integrations for a wide range of robot arms and teleoperation devices — from industrial manipulators to affordable hobbyist platforms, VR headsets, haptic devices, and full arm-plus-teleoperator kits. This page showcases community-maintained integrations you can use for teleoperation, data collection, and policy deployment.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> These projects are developed and maintained by third parties. Please refer to each repository for installation instructions, hardware requirements, and support.
|
||||
|
||||
Drop-in plugins are auto-discovered by package name: LeRobot imports any installed package prefixed with `lerobot_robot_` or `lerobot_teleoperator_`. Once installed, reference the `type` the plugin registers (see its README — it may differ from the package name) directly from any LeRobot command:
|
||||
|
||||
```bash
|
||||
pip install lerobot_robot_<name> lerobot_teleoperator_<name>
|
||||
|
||||
lerobot-record \
|
||||
--robot.type=<robot_name> \
|
||||
--teleop.type=<teleoperator_name> \
|
||||
--dataset.repo_id=${HF_USER}/my-dataset \
|
||||
--dataset.num_episodes=5
|
||||
```
|
||||
|
||||
> [!TIP]
|
||||
> ⚠️ marks projects that are forks/extensions of LeRobot. They may require custom setup rather than working with an unmodified install. All other entries are drop-in plugins.
|
||||
|
||||
## Industrial & Collaborative Arms
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/SpesRobotics/lerobot-robot-xarm">lerobot-robot-xarm</a></td>
|
||||
<td>Plugin for the xArm collaborative arm series from <a href="https://www.ufactory.cc/">UFACTORY</a>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/lebai-robotics/lerobot_lebai">lerobot_lebai</a></td>
|
||||
<td>Plugin for the six-axis collaborative arms from <a href="https://lebai.ltd/en/">Lebai</a>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/wengmister/LeFranX">LeFranX</a> ⚠️</td>
|
||||
<td>LeRobot extension for the <a href="https://franka.de/">Franka</a> research arm, paired with the <a href="https://www.robotera.com/">RobotEra XHand</a> hand for VR teleoperation.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
#### Universal Robots UR5e
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/yechen056/UR5e-LeRobot">UR5e-LeRobot</a> ⚠️</td>
|
||||
<td>LeRobot extension for the <a href="https://www.universal-robots.com/">Universal Robots UR5e</a>, with single-arm and bimanual support.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/scy-v/lerobot_ur5e_auto">lerobot_ur5e_auto</a> ⚠️</td>
|
||||
<td>LeRobot extension for a mobile <a href="https://www.universal-robots.com/">Universal Robots UR5e</a>, adding automated recording at scale with minimal supervision.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/F-Fer/lerobot_ur5e_gello">lerobot_robot_ur5e</a></td>
|
||||
<td>Plugin for the <a href="https://www.universal-robots.com/">Universal Robots UR5e</a> with a <a href="https://robotiq.com/">Robotiq</a> gripper, over RTDE control.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Research & Learning Arms
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/TrossenRobotics/lerobot_trossen">lerobot_trossen</a></td>
|
||||
<td>Plugin for the WidowX and ALOHA-style arms from <a href="https://www.trossenrobotics.com/">Trossen Robotics</a>.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
#### AgileX Piper
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/AgRoboticsResearch/lerobot_robot_piper">lerobot_robot_piper (AgRobotics Research)</a></td>
|
||||
<td>Plugin for the <a href="https://global.agilex.ai/">AgileX Piper</a> arm.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/WeGo-Robotics/lerobot_robot_piper">lerobot_robot_piper (WeGo Robotics)</a></td>
|
||||
<td>Plugin for the <a href="https://global.agilex.ai/">AgileX Piper</a> arm, with multi-arm teleoperation, safety limits, and GUI tools.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Affordable & Hobbyist Arms
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/servodevelop/fashionstar-lerobot-robot-cello">fashionstar-lerobot-robot-cello</a></td>
|
||||
<td>Plugin for the StarAI Cello 6+1 degrees of freedom robot arm from <a href="https://fashionstar.com.hk/">FashionStar</a>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/servodevelop/fashionstar-lerobot-robot-viola">fashionstar-lerobot-robot-viola</a></td>
|
||||
<td>Plugin for the compact StarAI Viola 6+1 degrees of freedom robot arm from <a href="https://fashionstar.com.hk/">FashionStar</a>.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Service, Mobile & Utility Robots
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/ugo-plus/lerobot-robot-ugo-pro">lerobot-robot-ugo-pro</a></td>
|
||||
<td>Plugin for the ugo Pro dual-arm service robot from <a href="https://ugo.plus/products/ugo-pro/">ugo</a>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/zuoxingdong/lerobot_robot_lekiwi_pincopen">lerobot_robot_lekiwi_pincopen</a></td>
|
||||
<td>Plugin for a LeKiwi mobile manipulator with a <a href="https://github.com/pollen-robotics/PincOpen">PincOpen</a> gripper.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/KillingJacky/lerobot-robot-dummy">lerobot-robot-dummy</a></td>
|
||||
<td>Plugin simulating a robot for recording without hardware. Useful for debugging !</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Teleoperators
|
||||
|
||||
### VR & Motion Controllers
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/SpesRobotics/lerobot-teleoperator-teleop">lerobot-teleoperator-teleop</a></td>
|
||||
<td>Plugin turning a phone or VR headset into a teleoperator via <a href="https://immersiveweb.dev">WebXR</a>, wrapping the open-source <a href="https://github.com/SpesRobotics/teleop"><code>teleop</code></a> library.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/Jas000n/lerobot-teleoperator-spacemouse">lerobot-teleoperator-spacemouse</a></td>
|
||||
<td>Plugin for the <a href="https://3dconnexion.com/">3Dconnexion SpaceMouse</a>, with inverse kinematics for SO-ARMS robots.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/Dream-Machines-Robotics/vr-teleop-kit">vr-teleop-kit</a></td>
|
||||
<td>Plugin teleoperating arms from a <a href="https://www.meta.com/quest/">Meta Quest</a> (WebXR), relying on URDF descriptions for inverse kinematics.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/xensedyl/lerobot-teleoperator-pico4">lerobot-teleoperator-pico4</a></td>
|
||||
<td>Plugin for the <a href="https://www.picoxr.com/">PICO 4</a> VR headset, with a companion controller-free <a href="https://github.com/xensedyl/lerobot-teleoperator-pico4-hand">hand-tracking variant</a>.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
### Leader Arms
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/F-Fer/lerobot_ur5e_gello">lerobot_teleoperator_gello</a></td>
|
||||
<td>Plugin for the 7 degrees of freedom <a href="https://wuphilipp.github.io/gello_site/">GELLO</a> teleoperator.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/uynitsuj/lerobot_teleoperator_yamactiveleader">lerobot_teleoperator_yamactiveleader</a></td>
|
||||
<td>Plugin for the active YAM teleoperator from <a href="https://i2rt.com/">I2RT</a>, a bilateral force-feedback arm.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/charlie8612/lerobot_teleoperator_omy">lerobot_teleoperator_omy</a></td>
|
||||
<td>Plugin for the OMY-L100 6 degrees of freedom teleoperator from <a href="https://www.robotis.com/">ROBOTIS</a>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://pypi.org/project/lerobot-teleoperator-pipermate/">lerobot-teleoperator-pipermate</a></td>
|
||||
<td>Plugin for the PiperMate teleoperator (<a href="https://fashionstar.com.hk/">FashionStar</a> UART servos), driving the <a href="https://global.agilex.ai/">AgileX Piper</a> arm.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
### Haptic Devices
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/chohh7391/lerobot_teleoperator_inverse3">lerobot_teleoperator_inverse3</a></td>
|
||||
<td>Plugin for the <a href="https://www.haply.co/">Haply Inverse3</a> haptic device, adding force-feedback teleoperation.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/hzhz112/lerobot_teleoperator_omega7">lerobot_teleoperator_omega7</a></td>
|
||||
<td>Plugin for the <a href="https://www.forcedimension.com/">Force Dimension omega.7</a> haptic device, adding force-feedback teleoperation.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
### Networked & Remote
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://pypi.org/project/lerobot-teleoperator-livekit/">lerobot-teleoperator-livekit</a></td>
|
||||
<td>Plugin receiving teleoperation commands over a <a href="https://livekit.io/">LiveKit</a> Portal (WebRTC) for remote control.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Full Kits (Robot + Teleoperator)
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/villekuosmanen/lerobot-arx5">lerobot-arx5</a></td>
|
||||
<td>Plugin for the <a href="https://www.arx-x.com/">ARX5</a> arm: <a href="https://pypi.org/project/lerobot-robot-arx5/"><code>lerobot-arx5</code></a> robot arm with its <a href="https://pypi.org/project/lerobot-teleoperator-arx5/"><code>lerobot-teleoperator-arx5</code></a> teleoperator arm.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/robertorobotics/Nextis-AIRA-3D">Nextis-AIRA-3D</a></td>
|
||||
<td>Plugin for the 7 degrees of freedom arm from <a href="https://www.nextis.tech">Nextis</a>: robot arm <code>aira_follower</code> and teleoperator arm <code>aira_leader</code>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/pravsels/lerobot_yam">lerobot_yam</a></td>
|
||||
<td>Plugin suite for the YAM arm from <a href="https://i2rt.com/">I2RT</a>: robot arm <code>yam_follower</code> and teleoperator arm <code>yam_leader</code>.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/robot-learning-co/trlc-dk1">trlc-dk1</a></td>
|
||||
<td>Plugin for the development kit from <a href="https://www.robot-learning.co/">The Robot Learning Company</a>: single and bimanual arms follower/teleoperator types.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/hexfellow/hex_lerobot_drivers">hex_lerobot_drivers</a></td>
|
||||
<td>Plugin suite for <a href="https://hexfellow.com/">HEXFELLOW</a> devices: robots, teleoperators, and cameras (see <a href="./third_party_sensors">Cameras & Sensors</a>).</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/Hiwonder-official/lerobot-robot-nexarm-follower">lerobot-robot-nexarm-follower</a></td>
|
||||
<td>Plugin for the NexArm from <a href="https://www.hiwonder.com/">Hiwonder</a>: the <a href="https://github.com/Hiwonder-official/lerobot-robot-nexarm-follower">robot arm</a> and its matching <a href="https://github.com/Hiwonder-official/lerobot-teleoperator-nexarm-leader">teleoperator arm</a>.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## ROS 2 Bridges
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/ngres/leros2">leros2</a></td>
|
||||
<td>Plugin bridging ROS 2 topics and actions to LeRobot robots and teleoperators.</td>
|
||||
</tr>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/ROBOTIS-GIT/lerobot_robot_ros2_zenoh">lerobot_robot_ros2_zenoh</a></td>
|
||||
<td>Plugin bridging ROS 2 robots to LeRobot over <a href="https://zenoh.io">Zenoh</a> pub/sub transport.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Contributing
|
||||
|
||||
Built your own LeRobot hardware integration? The plugin system makes it straightforward to add new robots and teleoperators — check out the [Bring Your Own Hardware](./integrate_hardware) guide to get started, and share your project with the community!
|
||||
@@ -0,0 +1,99 @@
|
||||
# Third-Party Cameras & Sensors
|
||||
|
||||
The LeRobot ecosystem extends far beyond its natively supported cameras (OpenCV, Intel RealSense, ZMQ, Reachy 2). Thanks to LeRobot's plugin architecture, the community has built drop-in camera and sensor integrations — from depth cameras to vision-based tactile sensors. This page showcases community-maintained camera and sensor integrations you can use for teleoperation, data collection, and policy deployment.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> These projects are developed and maintained by third parties. Please refer to each repository for installation instructions, hardware requirements, and support.
|
||||
|
||||
Drop-in plugins are auto-discovered by package name: LeRobot imports any installed package prefixed with `lerobot_camera_`. Once installed, reference the camera `type` the plugin registers (see its README — it may differ from the package name) directly from any LeRobot command:
|
||||
|
||||
```bash
|
||||
pip install lerobot_camera_<name>
|
||||
|
||||
lerobot-record \
|
||||
--robot.type=so101_follower \
|
||||
--robot.port=/dev/ttyACM0 \
|
||||
--robot.cameras="{ front: {type: <name>, width: 640, height: 480, fps: 30} }" \
|
||||
--dataset.repo_id=${HF_USER}/my-dataset \
|
||||
--dataset.num_episodes=5
|
||||
```
|
||||
|
||||
## Tactile Sensors
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/xensedyl/lerobot-camera-xense">lerobot-camera-xense</a></td>
|
||||
<td>Plugin for <a href="https://www.xenserobotics.com/">Xense</a> vision-based tactile sensors, exposing rectified/difference images, depth, and 2D markers.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Depth Cameras
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/hexfellow/hex_lerobot_drivers/tree/main/lerobot_camera_berxel">lerobot_camera_berxel</a></td>
|
||||
<td>Plugin for the <a href="https://www.berxel.com/">Berxel</a> depth camera, part of the broader <a href="https://hexfellow.com/">HEXFELLOW</a> driver suite.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Networked Cameras
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/F-Fer/lerobot_ur5e_gello">lerobot_camera_zmq</a></td>
|
||||
<td>Plugin streaming <a href="https://www.stereolabs.com/">Stereolabs ZED</a> and USB camera frames from a Raspberry Pi over the network.</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Virtual Cameras
|
||||
|
||||
<!-- prettier-ignore -->
|
||||
<table width="100%" style="display:table; width:100%; table-layout:fixed;">
|
||||
<colgroup>
|
||||
<col width="30%" />
|
||||
<col width="70%" />
|
||||
</colgroup>
|
||||
<thead style="border:0">
|
||||
<tr style="border:0"><th>Project</th><th>Description</th></tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr style="border:0">
|
||||
<td><a href="https://github.com/hexfellow/hex_lerobot_drivers/tree/main/lerobot_camera_dummy">lerobot_camera_dummy</a></td>
|
||||
<td>Plugin simulating a camera for recording without hardware. Useful for debugging !</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## Contributing
|
||||
|
||||
Built your own LeRobot camera or sensor integration? Package it as an installable `lerobot_camera_<name>` plugin and it will be auto-discovered by the LeRobot CLI — see the [Bring Your Own Hardware](./integrate_hardware) guide and the [Cameras](./cameras) reference to get started, then share your project with the community!
|
||||
@@ -40,3 +40,15 @@ lerobot-eval \
|
||||
```
|
||||
|
||||
However, in most cases, presence of an accelerator is detected automatically and `policy.device` parameter can be omitted from CLI commands.
|
||||
|
||||
## Mixed precision
|
||||
|
||||
Training precision is owned by `--accelerator.mixed_precision`, which accepts `no` (default) and `bf16`:
|
||||
|
||||
```bash
|
||||
lerobot-train \
|
||||
--policy.type=act \
|
||||
--accelerator.mixed_precision=bf16 ...
|
||||
```
|
||||
|
||||
`bf16` requires an accelerator that supports it.
|
||||
|
||||
+12
-44
@@ -30,19 +30,10 @@ Only Qwen + the action head are used. The world model is not needed at inference
|
||||
|
||||
Available presets via `action_model_type`:
|
||||
|
||||
| Preset | Heads | Head dim |
|
||||
| ------- | ----- | -------- |
|
||||
| `DiT-B` | 12 | 64 |
|
||||
| `DiT-L` | 32 | 48 |
|
||||
|
||||
The preset only sets the attention geometry, and each entry can be overridden by
|
||||
`action_num_heads` / `action_attention_head_dim`. Two widths follow from it:
|
||||
|
||||
- the DiT's **internal** width is `heads x head_dim` (768 for `DiT-B`), derived rather than configured;
|
||||
- the DiT's **output** width, and the width of the action-decoder and state-encoder MLPs, is
|
||||
`action_hidden_size` (default 1024).
|
||||
|
||||
So `DiT-B` runs a 768-wide transformer that projects to 1024. The two are independent.
|
||||
| Preset | Hidden dim | Heads | Head dim |
|
||||
| ------- | ---------- | ----- | -------- |
|
||||
| `DiT-B` | 768 | 12 | 64 |
|
||||
| `DiT-L` | 1536 | 32 | 48 |
|
||||
|
||||
### World model details
|
||||
|
||||
@@ -83,27 +74,10 @@ Key parameters in `VLAJEPAConfig`:
|
||||
| `num_inference_timesteps` | 4 | Euler integration steps for action denoising |
|
||||
| `freeze_qwen` | `False` | Freeze the Qwen3-VL backbone and only train the action head |
|
||||
| `reinit_modules` | `None` | Key prefixes allowed to be randomly re-initialised on load (for cross-embodiment transfer, see [Fine-tuning on a different embodiment](#fine-tuning-on-a-different-embodiment)) |
|
||||
| `resize_images_to` | `None` | `(height, width)` every camera frame is resized to before the Qwen3-VL vision tower. `None` keeps the native resolution, and Qwen3-VL's patch count grows with it, so a 720x1280 camera can exhaust GPU memory. The published checkpoints use `[224, 224]` |
|
||||
| `gripper_dim` | 6 | Index of the gripper dimension in the action vector. Ignored when `gripper_joint_names` matches a dataset action name |
|
||||
| `gripper_joint_names` | `["gripper"]` | Action-dimension names identifying the gripper; the matched index wins over `gripper_dim` |
|
||||
| `gripper_threshold` | 0.5 | Threshold used by `pre_snap_gripper_action` and `binarize_gripper_action`. Note `binarize` runs *after* unnormalization, so this is compared against the gripper's physical value |
|
||||
| `pre_snap_gripper_action` | `False` | Snap the gripper dim to {0, 1} before unnormalization. LIBERO-specific, see below |
|
||||
| `binarize_gripper_action` | `False` | Binarize the gripper dim to {-1, 1} after unnormalization. LIBERO-specific, see below |
|
||||
| `clip_normalized_actions` | `True` | Clip normalized actions to [-1, 1] before unnormalizing. Only applied when `ACTION` uses `MIN_MAX`; ignored (with a warning) under `MEAN_STD`, where it would truncate at 1 sigma |
|
||||
| `world_model_num_views` | `None` | Camera views the world-model predictor is built for. Baked into checkpoint shapes. `None` falls back to `jepa_tubelet_size`, which is what the published checkpoints encode |
|
||||
|
||||
<Tip warning={true}>
|
||||
|
||||
`pre_snap_gripper_action` and `binarize_gripper_action` are a port of the starVLA LIBERO eval
|
||||
loop and are only correct for LIBERO's action convention. `pre_snap` writes {0, 1} into
|
||||
*normalized* space, the unnormalizer maps those to the midpoint and the max, and `binarize` then
|
||||
compares that **physical** value against `gripper_threshold` (0.5). For a gripper measured in
|
||||
degrees, mm or [0, 100], both values land above the threshold and the commanded gripper becomes a
|
||||
constant. They default to `False` for that reason; enable them only for LIBERO-style setups, and
|
||||
set `gripper_threshold` in the gripper's own units if you do. The processor factory warns when the
|
||||
dataset stats show the range cannot work.
|
||||
|
||||
</Tip>
|
||||
| `gripper_dim` | 6 | Index of the gripper dimension in the action vector (e.g. 6 for a 7-DoF arm with gripper as the last joint) |
|
||||
| `gripper_threshold` | 0.5 | Threshold used by `pre_snap_gripper_action` and `binarize_gripper_action` to binarize the gripper dimension |
|
||||
| `pre_snap_gripper_action` | `True` | Snap the gripper dim to {0, 1} before unnormalization. Set to `False` for robots without a binary gripper |
|
||||
| `binarize_gripper_action` | `True` | Binarize the gripper dim to {-1, 1} after unnormalization. Set to `False` for robots without a binary gripper |
|
||||
|
||||
---
|
||||
|
||||
@@ -213,20 +187,14 @@ lerobot-eval \
|
||||
|
||||
## Fine-tuning on datasets with a different number of cameras
|
||||
|
||||
The pretrained world model predictor was trained with `embed_dim = world_model_num_views × 1024`, i.e. two camera views.
|
||||
|
||||
<Tip>
|
||||
|
||||
This view count used to be read from `jepa_tubelet_size`, which also names the JEPA encoder's *temporal* tubelet size. `world_model_num_views` is the field for it now; leaving it at `None` falls back to `jepa_tubelet_size` so the published checkpoints keep loading unchanged.
|
||||
|
||||
</Tip>
|
||||
The pretrained world model predictor was trained with `embed_dim = jepa_tubelet_size × 1024` (default `jepa_tubelet_size=2`).
|
||||
|
||||
**Default behaviour — view padding / trimming (no action required)**
|
||||
|
||||
When fine-tuning from `VLA-JEPA-Pretrain` the model automatically adjusts the number of views fed to the world model to match `world_model_num_views`:
|
||||
When fine-tuning from `VLA-JEPA-Pretrain` the model automatically adjusts the number of views fed to the world model to match `jepa_tubelet_size`:
|
||||
|
||||
- **Single-view datasets (e.g. BridgeV2):** the single-view latent is duplicated to produce a two-view world-model input, preserving the JEPA self-supervised signal without any weight mismatch.
|
||||
- **>2-view datasets (e.g. DROID with 3 views):** all views are passed to the Qwen backbone (for richer context), but only the first `world_model_num_views` views (one wrist + one third-person, following the configured view order) are used for the world model.
|
||||
- **>2-view datasets (e.g. DROID with 3 views):** all views are passed to the Qwen backbone (for richer context), but only the first `jepa_tubelet_size` views (one wrist + one third-person, following the configured view order) are used for the world model.
|
||||
|
||||
**Option 1 — Disable the world model**
|
||||
|
||||
@@ -242,7 +210,7 @@ lerobot-train \
|
||||
|
||||
**Option 2 — Reinitialize the predictor input projection**
|
||||
|
||||
If you want to change `world_model_num_views` to a value other than 2, load the checkpoint with `strict=False` and reinitialize `model.video_predictor.predictor_embed` for the new `embed_dim`. All other predictor block weights (attention, MLP, norm, output projection) are camera-count-agnostic and can be reused from the pretrained checkpoint.
|
||||
If you want to change `jepa_tubelet_size` to a value other than 2, load the checkpoint with `strict=False` and reinitialize `model.video_predictor.predictor_embed` for the new `embed_dim`. All other predictor block weights (attention, MLP, norm, output projection) are camera-count-agnostic and can be reused from the pretrained checkpoint.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -0,0 +1,287 @@
|
||||
# Writing docstrings
|
||||
|
||||
LeRobot's API reference is generated directly from the docstrings in `src/lerobot/`. A docstring is not a
|
||||
comment — it is the published documentation for that object, and the format below is what the renderer and
|
||||
the CI checks parse.
|
||||
|
||||
This page is the contract. If you are adding or editing anything public in `src/lerobot/`, follow it.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **An undocumented public method is an invisible one.** `[[autodoc]]` silently skips members that have no
|
||||
> docstring — no warning, no error, it simply does not appear on the rendered page. Coverage and
|
||||
> API-reference completeness are the same problem.
|
||||
|
||||
## The format in one example
|
||||
|
||||
Google section headers, Hugging Face type formatting. Both, not one or the other.
|
||||
|
||||
````python
|
||||
def send_action(self, action: RobotAction, rate_hz: float = 30.0) -> RobotAction:
|
||||
"""Command the robot to move to a target joint configuration.
|
||||
|
||||
Values are clipped by the configured maximum relative target before reaching the motors, so the
|
||||
returned action may differ from the requested one.
|
||||
|
||||
Args:
|
||||
action (`dict[str, float]`):
|
||||
Target values keyed by motor name, e.g. `{"shoulder_pan.pos": 0.0}`. Keys must match the
|
||||
robot's action features.
|
||||
rate_hz (`float`, *optional*, defaults to `30.0`):
|
||||
Control loop frequency.
|
||||
|
||||
Returns:
|
||||
`dict[str, float]`: The action actually written to the motors after safety clipping.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot has not been connected.
|
||||
|
||||
Example:
|
||||
```python
|
||||
>>> from lerobot.robots.so_follower import SO101Follower, SO101FollowerConfig
|
||||
>>> robot = SO101Follower(SO101FollowerConfig(port="/dev/ttyACM0")) # doctest: +SKIP
|
||||
>>> robot.connect() # doctest: +SKIP
|
||||
>>> robot.send_action({"shoulder_pan.pos": 0.0}) # doctest: +SKIP
|
||||
```
|
||||
"""
|
||||
````
|
||||
|
||||
Cross-references are omitted from the examples on this page — see [Cross-references](#cross-references) for
|
||||
their syntax and why they cannot be shown inside a code block.
|
||||
|
||||
## Rules
|
||||
|
||||
### Sections
|
||||
|
||||
`Args:` · `Returns:` · `Raises:` · `Yields:` · `Example:` · `Note:`
|
||||
|
||||
In that order. No other section headers. A one-line summary comes first, then an optional free-form
|
||||
description, then the sections.
|
||||
|
||||
### The `Args:` line is machine-parsed
|
||||
|
||||
```
|
||||
name (`type`, *optional*, defaults to `X`):
|
||||
Description, indented on its own line.
|
||||
```
|
||||
|
||||
The `*optional*, defaults to` clause is **checked against the real signature default** by
|
||||
`make check-docstrings`. It is not decorative — if you write a default that has drifted from the code, CI
|
||||
fails. Omit the clause entirely for required parameters:
|
||||
|
||||
```python
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to, e.g. `/dev/ttyACM0`.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps the magnitude of the relative positional target vector. `None` disables clipping.
|
||||
use_degrees (`bool`, *optional*, defaults to `True`):
|
||||
Keep `True` for backward compatibility with existing policies and datasets.
|
||||
```
|
||||
|
||||
Types go in backticks. Use `*optional*` with no `defaults to` when the default is `None` or is otherwise not
|
||||
worth restating.
|
||||
|
||||
### `Returns:` is type-first
|
||||
|
||||
One indented line, type first, then a colon, then the description:
|
||||
|
||||
```python
|
||||
Returns:
|
||||
`dict[str, float]`: The action actually written to the motors after safety clipping.
|
||||
```
|
||||
|
||||
`Yields:` takes the same shape.
|
||||
|
||||
### `**Attributes**:`, never `Attributes:`
|
||||
|
||||
doc-builder parses a bare `Attributes:` as a **synonym for `Parameters:`**, so your attributes get rendered
|
||||
as constructor arguments. This is silent and wrong. Whenever the attributes differ from the constructor
|
||||
parameters, use the bold form with a `--` separator:
|
||||
|
||||
```python
|
||||
class Robot(abc.ABC):
|
||||
"""The base abstract class for all LeRobot-compatible robots.
|
||||
|
||||
**Attributes**:
|
||||
- **config_class** (`type[RobotConfig]`) -- The expected configuration class for this robot.
|
||||
- **name** (`str`) -- The unique robot name used to identify this robot type.
|
||||
"""
|
||||
```
|
||||
|
||||
Note `--`, not `:`.
|
||||
|
||||
### Cross-references
|
||||
|
||||
Use doc-builder's bracket syntax: a square-bracketed backtick-quoted path. **Sphinx roles (`:pymeth:`,
|
||||
`:pyattr:`) are not supported** and render as literal text on the page.
|
||||
|
||||
| Want | Write |
|
||||
| ---------------------------- | ----------------------------------- |
|
||||
| Class in the main package | [`Robot`] |
|
||||
| Method, show the full path | [`Robot.connect`] |
|
||||
| Method, show the bare name | [`~Robot.connect`] |
|
||||
| Nested path | [`~robots.Robot.connect`] |
|
||||
| Object in another HF library | [`~accelerate.Accelerator`] |
|
||||
|
||||
The `~` strips the path from the **link text only**; the link still resolves to the full path.
|
||||
|
||||
> [!NOTE]
|
||||
> doc-builder resolves this syntax everywhere in a page — including inside fenced code blocks. That is why
|
||||
> the docstring examples on this page use plain prose instead of cross-references: a code block containing
|
||||
> one would render the resolved link rather than the syntax you need to type. In your own docstrings, use
|
||||
> cross-references freely; this restriction only affects documentation _about_ the syntax.
|
||||
|
||||
### Callouts
|
||||
|
||||
Use GitHub-style blockquotes:
|
||||
|
||||
```markdown
|
||||
> [!TIP]
|
||||
> Call this once at startup — it takes about two seconds.
|
||||
|
||||
> [!WARNING]
|
||||
> Torque is disabled on disconnect. The arm will drop if it is holding a load.
|
||||
```
|
||||
|
||||
The `<Tip>` component is legacy per doc-builder; don't add new ones.
|
||||
|
||||
### Examples must be fenced
|
||||
|
||||
An example lives inside a fenced ` ```python ` block containing `>>> `. The fence is what makes it render
|
||||
as a code block, and it is what the doctest preprocessor's regex looks for:
|
||||
|
||||
````python
|
||||
Example:
|
||||
```python
|
||||
>>> from lerobot.robots.so_follower import SO101FollowerConfig
|
||||
>>> cfg = SO101FollowerConfig(port="/dev/ttyACM0")
|
||||
>>> cfg.use_degrees
|
||||
True
|
||||
```
|
||||
````
|
||||
|
||||
> [!WARNING]
|
||||
> An unfenced `>>>` is still collected — doctest finds prompts anywhere in a docstring. What you lose is the
|
||||
> rendering, so it shows up as a wall of prose on the page. Every example needs the fence.
|
||||
|
||||
Every example either executes in CI or carries `# doctest: +SKIP`. Anything that touches hardware, a GPU, or
|
||||
downloads from the Hub gets `+SKIP`:
|
||||
|
||||
````python
|
||||
Example:
|
||||
```python
|
||||
>>> robot.connect() # doctest: +SKIP
|
||||
>>> policy = ACTPolicy.from_pretrained("lerobot/act_aloha_sim_transfer_cube_human") # doctest: +SKIP
|
||||
```
|
||||
````
|
||||
|
||||
Add files containing runnable examples to `utils/documentation_tests.txt`.
|
||||
|
||||
Put examples on the three to five genuine entry points of a module. Examples on trivial accessors are noise.
|
||||
|
||||
## Three patterns you will hit constantly
|
||||
|
||||
### Config dataclasses
|
||||
|
||||
Configuration fields are historically documented with `#` comments above each field. **doc-builder cannot
|
||||
see inline comments** — such a class renders with every field listed and not a single description. Move them
|
||||
into an `Args:` block on the class docstring:
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class SOFollowerConfig:
|
||||
"""Configuration for SO-family follower arms.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to, e.g. `/dev/ttyACM0`.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps the magnitude of the relative positional target vector. A scalar applies to all motors;
|
||||
a dict maps motor name to a per-motor cap. `None` disables clipping.
|
||||
use_degrees (`bool`, *optional*, defaults to `True`):
|
||||
Keep `True` for backward compatibility with existing policies and datasets.
|
||||
"""
|
||||
|
||||
port: str
|
||||
max_relative_target: float | dict[str, float] | None = None
|
||||
use_degrees: bool = True
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **doc-builder does not inherit docstrings from base classes.** LeRobot's registered config classes are
|
||||
> often thin multiple-inheritance shims:
|
||||
>
|
||||
> ```python
|
||||
> @RobotConfig.register_subclass("so101_follower")
|
||||
> @dataclass
|
||||
> class SOFollowerRobotConfig(RobotConfig, SOFollowerConfig):
|
||||
> pass
|
||||
> ```
|
||||
>
|
||||
> That class renders **every** field — including the ones it inherits — with no descriptions at all, no
|
||||
> matter how well the bases are documented. The `Args:` block must live on the concrete class that
|
||||
> `[[autodoc]]` names, and it must cover inherited fields too.
|
||||
|
||||
### Base class, then concrete subclass
|
||||
|
||||
The abstract base carries the canonical contract. Subclasses document only what deviates — port semantics,
|
||||
calibration quirks, motor layout, supported feature keys. Do not copy the base contract into every subclass.
|
||||
|
||||
`Robot`, `Teleoperator`, `Camera`, `MotorsBus`, `ProcessorStep`, and `PreTrainedPolicy` all follow this
|
||||
shape.
|
||||
|
||||
### Module-level aliases
|
||||
|
||||
Several public names are aliases rather than distinct classes:
|
||||
|
||||
```python
|
||||
SO100FollowerConfig = SOFollowerRobotConfig
|
||||
SO101FollowerConfig = SOFollowerRobotConfig
|
||||
```
|
||||
|
||||
`[[autodoc]]` resolves the alias and renders the **canonical** class name, so a `## SO101FollowerConfig`
|
||||
heading will show `class lerobot.robots.so_follower.SOFollowerRobotConfig` in the body. Document the
|
||||
canonical class once, and mention the aliases in the page's prose rather than giving each alias its own
|
||||
autodoc block.
|
||||
|
||||
## What not to document
|
||||
|
||||
- **Private members.** Anything starting with `_` is not part of the public API.
|
||||
- **The type annotation restated as prose.** `port (`str`): A string.` adds nothing. Say what it is for.
|
||||
- **Vendored upstream code.** `src/lerobot/policies/molmoact2/molmoact2_hf_model/` is vendored from
|
||||
`transformers` and already carries upstream-style docstrings. Leave it alone — restyling it only creates
|
||||
conflicts on the next sync. It is excluded from the API reference and from the docstring checks.
|
||||
|
||||
## How this is enforced
|
||||
|
||||
| Check | What it catches |
|
||||
| ------------------------- | ---------------------------------------------------------------------------------------------------------- |
|
||||
| `make check-docstrings` | An `Args:` entry that doesn't match the signature; a documented default that has drifted from the real one |
|
||||
| `make doctest` | Examples that no longer run |
|
||||
| `make check-doctest-list` | Stale or unsorted entries in `utils/documentation_tests.txt` |
|
||||
| `ruff` (`D` rules) | Google-convention style violations |
|
||||
| `interrogate` | Docstring coverage falling below the current threshold |
|
||||
| doc-builder | A `[[autodoc]]` path that points at something that doesn't exist — this breaks the docs build |
|
||||
|
||||
Run them together before opening a PR:
|
||||
|
||||
```bash
|
||||
make check-docstrings && make doctest && pre-commit run --all-files
|
||||
```
|
||||
|
||||
Then render the page and actually look at it:
|
||||
|
||||
```bash
|
||||
doc-builder build lerobot docs/source/ --build_dir /tmp/doc-build
|
||||
```
|
||||
|
||||
## Checklist
|
||||
|
||||
- [ ] Every public member you touched has a docstring.
|
||||
- [ ] Every `Args:` entry matches the signature, including the `*optional*, defaults to` clause.
|
||||
- [ ] `Returns:` is type-first on one indented line.
|
||||
- [ ] No bare `Attributes:` — use `**Attributes**:` with `--` separators.
|
||||
- [ ] No Sphinx roles — cross-references use [`~module.Class.method`].
|
||||
- [ ] Examples are inside a fenced ` ```python ` block, and either run in CI or carry `# doctest: +SKIP`.
|
||||
- [ ] Config dataclass fields are in an `Args:` block on the concrete class, not `#` comments.
|
||||
- [ ] The rendered page has been eyeballed.
|
||||
+121
-20
@@ -25,7 +25,7 @@ discord = "https://discord.gg/s3KuuzsPFb"
|
||||
|
||||
[project]
|
||||
name = "lerobot"
|
||||
version = "0.6.1"
|
||||
version = "0.6.2"
|
||||
description = "🤗 LeRobot: State-of-the-art Machine Learning for Real-World Robotics in Pytorch"
|
||||
dynamic = ["readme"]
|
||||
license = { text = "Apache-2.0" }
|
||||
@@ -87,7 +87,7 @@ dependencies = [
|
||||
|
||||
# Build tools (required by opencv-python-headless on some platforms)
|
||||
"cmake>=3.29.0.1,<4.2.0",
|
||||
"setuptools>=71.0.0,<81.0.0",
|
||||
"setuptools>=71.0.0,<82.0.0", # torch 2.11 requires setuptools<82; a higher cap makes the resolver downgrade torch
|
||||
]
|
||||
|
||||
# Optional dependencies
|
||||
@@ -261,7 +261,7 @@ annotations = [
|
||||
# Development
|
||||
dev = ["pre-commit>=3.7.0,<5.0.0", "debugpy>=1.8.1,<1.9.0", "lerobot[grpcio-dep]", "grpcio-tools>=1.73.1,<2.0.0", "mypy>=1.19.1", "ruff>=0.14.1", "lerobot[notebook]"]
|
||||
notebook = ["jupyter>=1.0.0,<2.0.0", "ipykernel>=6.0.0,<7.0.0"]
|
||||
test = ["pytest>=8.1.0,<9.0.0", "pytest-timeout>=2.4.0,<3.0.0", "pytest-cov>=5.0.0,<8.0.0", "mock-serial>=0.0.1,<0.1.0 ; sys_platform != 'win32'"]
|
||||
test = ["pytest>=8.1.0,<10.0.0", "pytest-timeout>=2.4.0,<3.0.0", "pytest-cov>=5.0.0,<8.0.0", "mock-serial>=0.0.1,<0.1.0 ; sys_platform != 'win32'"]
|
||||
video_benchmark = ["scikit-image>=0.23.2,<0.26.0", "pandas>=2.2.2,<2.4.0"]
|
||||
|
||||
# Simulation
|
||||
@@ -346,6 +346,7 @@ lerobot-record="lerobot.scripts.lerobot_record:main"
|
||||
lerobot-replay="lerobot.scripts.lerobot_replay:main"
|
||||
lerobot-setup-motors="lerobot.scripts.lerobot_setup_motors:main"
|
||||
lerobot-teleoperate="lerobot.scripts.lerobot_teleoperate:main"
|
||||
lerobot-convert-dcp="lerobot.scripts.lerobot_convert_dcp:main"
|
||||
lerobot-eval="lerobot.scripts.lerobot_eval:main"
|
||||
lerobot-train="lerobot.scripts.lerobot_train:main"
|
||||
lerobot-train-tokenizer="lerobot.scripts.lerobot_train_tokenizer:main"
|
||||
@@ -400,19 +401,101 @@ exclude = ["tests/artifacts/**/*.safetensors", "*_pb2.py", "*_pb2_grpc.py"]
|
||||
# N: pep8-naming
|
||||
# TODO: Uncomment rules when ready to use
|
||||
select = [
|
||||
"E", "W", "F", "I", "B", "C4", "T20", "N", "UP", "SIM" #, "A", "S", "D", "RUF"
|
||||
"E", "W", "F", "I", "B", "C4", "T20", "N", "UP", "SIM", "D" #, "A", "S", "RUF"
|
||||
]
|
||||
ignore = [
|
||||
"E501", # Line too long
|
||||
"T201", # Print statement found
|
||||
"T203", # Pprint statement found
|
||||
"B008", # Perform function call in argument defaults
|
||||
# D100/D104: module- and package-level docstrings. The API reference is generated from class and
|
||||
# function docstrings; a banner at the top of every file and every __init__.py would not appear on any
|
||||
# rendered page. Coverage of the things that do get rendered is enforced by interrogate instead.
|
||||
"D100",
|
||||
"D104",
|
||||
]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"__init__.py" = ["F401", "F403", "E402"]
|
||||
"__init__.py" = ["F401", "F403", "E402", "D104"]
|
||||
# E402: conditional-import guards (TYPE_CHECKING / is_package_available) must precede the imports they protect
|
||||
"src/lerobot/scripts/convert_dataset_v21_to_v30.py" = ["E402"]
|
||||
|
||||
# D (pydocstyle) is enabled globally, but only holds for code that has been converted to the docstring
|
||||
# standard in docs/source/writing_docstrings.mdx. Every module below is still on the old style; each entry
|
||||
# is deleted as that module is converted, and this block can be removed once it is empty.
|
||||
#
|
||||
# Not part of the API reference and not planned for conversion: tests, examples, benchmarks, templates,
|
||||
# CI helper scripts and the packaging shim.
|
||||
"tests/**" = ["D"]
|
||||
"examples/**" = ["D"]
|
||||
"benchmarks/**" = ["D"]
|
||||
"scripts/**" = ["D"]
|
||||
"setup.py" = ["D"]
|
||||
"src/lerobot/templates/**" = ["D"]
|
||||
# Vendored from transformers; keeps its upstream docstring style so syncs stay clean.
|
||||
"src/lerobot/policies/molmoact2/molmoact2_hf_model/**" = ["D"]
|
||||
# Awaiting conversion, one PR per module.
|
||||
"src/lerobot/annotations/**" = ["D"]
|
||||
"src/lerobot/async_inference/**" = ["D"]
|
||||
"src/lerobot/cameras/**" = ["D"]
|
||||
"src/lerobot/common/**" = ["D"]
|
||||
"src/lerobot/configs/**" = ["D"]
|
||||
"src/lerobot/data_processing/**" = ["D"]
|
||||
"src/lerobot/datasets/**" = ["D"]
|
||||
"src/lerobot/distributed/**" = ["D"]
|
||||
"src/lerobot/envs/**" = ["D"]
|
||||
"src/lerobot/jobs/**" = ["D"]
|
||||
"src/lerobot/model/**" = ["D"]
|
||||
"src/lerobot/motors/**" = ["D"]
|
||||
"src/lerobot/optim/**" = ["D"]
|
||||
"src/lerobot/policies/__init__.py" = ["D"]
|
||||
"src/lerobot/policies/pi_gemma.py" = ["D"]
|
||||
"src/lerobot/policies/common/**" = ["D"]
|
||||
# Wave 3 of the docstring initiative documents each policy family's config class in full, plus only
|
||||
# the public forward/select_action surface of modeling_*.py's main <Family>Policy class and the
|
||||
# processor_*.py's make_<family>_pre_post_processors factory. modeling_*.py and processor_*.py also
|
||||
# contain internal building blocks (nn.Module helpers, ProcessorStep internals) that remain out of
|
||||
# scope, so those two file patterns stay D-ignored wholesale rather than enumerated per symbol; the
|
||||
# narrower Policy/processor-factory scope is instead enforced via the AST coverage check and
|
||||
# utils/check_docstrings.py's leaf-module entries. configuration_*.py is fully documented and stays
|
||||
# checked here.
|
||||
"src/lerobot/policies/*/modeling_*.py" = ["D"]
|
||||
"src/lerobot/policies/*/processor_*.py" = ["D"]
|
||||
"src/lerobot/policies/evo1/evo1_model.py" = ["D"]
|
||||
"src/lerobot/policies/evo1/flow_matching.py" = ["D"]
|
||||
"src/lerobot/policies/evo1/internvl3_embedder.py" = ["D"]
|
||||
"src/lerobot/policies/fastwam/wan/**" = ["D"]
|
||||
"src/lerobot/policies/groot/action_head/**" = ["D"]
|
||||
"src/lerobot/policies/groot/groot_n1_7.py" = ["D"]
|
||||
"src/lerobot/policies/groot/utils.py" = ["D"]
|
||||
"src/lerobot/policies/lingbot_va/utils.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/action_interpolator.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/action_queue.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/debug_tracker.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/debug_visualizer.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/latency_tracker.py" = ["D"]
|
||||
"src/lerobot/policies/rtc/relative.py" = ["D"]
|
||||
"src/lerobot/policies/smolvla/smolvlm_with_expert.py" = ["D"]
|
||||
"src/lerobot/policies/vla_jepa/action_head.py" = ["D"]
|
||||
"src/lerobot/policies/vla_jepa/qwen_interface.py" = ["D"]
|
||||
"src/lerobot/policies/vla_jepa/world_model.py" = ["D"]
|
||||
"src/lerobot/policies/vqbet/vqbet_utils.py" = ["D"]
|
||||
"src/lerobot/policies/wall_x/constant.py" = ["D"]
|
||||
"src/lerobot/policies/wall_x/qwen_model/**" = ["D"]
|
||||
"src/lerobot/policies/wall_x/utils.py" = ["D"]
|
||||
"src/lerobot/policies/xvla/action_hub.py" = ["D"]
|
||||
"src/lerobot/policies/xvla/soft_transformer.py" = ["D"]
|
||||
"src/lerobot/policies/xvla/utils.py" = ["D"]
|
||||
"src/lerobot/processor/**" = ["D"]
|
||||
"src/lerobot/rewards/**" = ["D"]
|
||||
"src/lerobot/rl/**" = ["D"]
|
||||
"src/lerobot/rollout/**" = ["D"]
|
||||
"src/lerobot/scripts/**" = ["D"]
|
||||
"src/lerobot/teleoperators/**" = ["D"]
|
||||
"src/lerobot/transforms/**" = ["D"]
|
||||
"src/lerobot/transport/**" = ["D"]
|
||||
"src/lerobot/utils/**" = ["D"]
|
||||
"src/lerobot/lerobot_types.py" = ["D"]
|
||||
[tool.ruff.lint.isort]
|
||||
combine-as-imports = true
|
||||
known-first-party = ["lerobot"]
|
||||
@@ -456,25 +539,34 @@ default.extend-ignore-identifiers-re = [
|
||||
"seperated_timestep",
|
||||
]
|
||||
|
||||
# TODO: Uncomment when ready to use
|
||||
# [tool.interrogate]
|
||||
# ignore-init-module = true
|
||||
# ignore-init-method = true
|
||||
# ignore-nested-functions = false
|
||||
# ignore-magic = false
|
||||
# ignore-semiprivate = false
|
||||
# ignore-private = false
|
||||
# ignore-property-decorators = false
|
||||
# ignore-module = false
|
||||
# ignore-setters = false
|
||||
# fail-under = 80
|
||||
# output-format = "term-missing"
|
||||
# color = true
|
||||
# paths = ["src/lerobot"]
|
||||
# Docstring coverage gate. `fail-under` is a RATCHET, not a target: it is set just below the currently
|
||||
# measured coverage so it passes today, and is raised in the same PR that documents a module. Never set it
|
||||
# to a value that fails on main. The destination is 100; see docs/source/writing_docstrings.mdx.
|
||||
[tool.interrogate]
|
||||
ignore-init-module = true
|
||||
ignore-init-method = true
|
||||
ignore-nested-functions = false
|
||||
ignore-magic = false
|
||||
ignore-semiprivate = false
|
||||
ignore-private = false
|
||||
ignore-property-decorators = false
|
||||
ignore-module = false
|
||||
ignore-setters = false
|
||||
fail-under = 58
|
||||
output-format = "term-missing"
|
||||
color = true
|
||||
paths = ["src/lerobot"]
|
||||
exclude = ["src/lerobot/policies/molmoact2/molmoact2_hf_model"]
|
||||
|
||||
# TODO: Enable mypy gradually module by module across multiple PRs
|
||||
# Uncomment [tool.mypy] first, then uncomment individual module overrides as they get proper type annotations
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
markers = [
|
||||
"multigpu: distributed tests needing 2-4 GPUs (CI: docker_publish.yml lane)",
|
||||
"multigpu_heavy: 8-GPU sweeps and soak tests; never run in CI",
|
||||
]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.12"
|
||||
ignore_missing_imports = true
|
||||
@@ -521,6 +613,15 @@ disallow_untyped_defs = true
|
||||
disallow_incomplete_defs = true
|
||||
check_untyped_defs = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = "lerobot.distributed.*"
|
||||
ignore_errors = false
|
||||
|
||||
# extra strictness for the distributed engine
|
||||
disallow_untyped_defs = true
|
||||
disallow_incomplete_defs = true
|
||||
check_untyped_defs = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = "lerobot.optim.*"
|
||||
ignore_errors = false
|
||||
|
||||
@@ -14,8 +14,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""
|
||||
LeRobot -- PyTorch library for real-world robotics.
|
||||
"""LeRobot -- PyTorch library for real-world robotics.
|
||||
|
||||
Provides datasets, pretrained policies, and tools for training, evaluation,
|
||||
data collection, and robot control. Integrates with Hugging Face Hub for
|
||||
|
||||
@@ -13,7 +13,7 @@
|
||||
# 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.
|
||||
"""To enable `lerobot.__version__`"""
|
||||
"""To enable `lerobot.__version__`."""
|
||||
|
||||
from importlib.metadata import PackageNotFoundError, version
|
||||
|
||||
|
||||
@@ -33,10 +33,10 @@ class Camera(abc.ABC):
|
||||
- Connection/disconnection
|
||||
- Frame capture (sync/async/latest)
|
||||
|
||||
Attributes:
|
||||
fps (int | None): Configured frames per second
|
||||
width (int | None): Frame width in pixels
|
||||
height (int | None): Frame height in pixels
|
||||
**Attributes**:
|
||||
- **fps** (`int | None`) -- Configured frames per second.
|
||||
- **width** (`int | None`) -- Frame width in pixels.
|
||||
- **height** (`int | None`) -- Frame height in pixels.
|
||||
"""
|
||||
|
||||
def __init__(self, config: CameraConfig):
|
||||
|
||||
@@ -40,17 +40,20 @@ class OpenCVCameraConfig(CameraConfig):
|
||||
OpenCVCameraConfig(0, 30, 1280, 720, fourcc="YUYV") # With YUYV format
|
||||
```
|
||||
|
||||
Attributes:
|
||||
index_or_path: Either an integer representing the camera device index,
|
||||
or a Path object pointing to a video file.
|
||||
fps: Requested frames per second for the color stream.
|
||||
width: Requested frame width in pixels for the color stream.
|
||||
height: Requested frame height in pixels for the color stream.
|
||||
color_mode: Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation.
|
||||
warmup_s: Time reading frames before returning from connect (in seconds)
|
||||
fourcc: FOURCC code for video format (e.g., "MJPG", "YUYV", "I420"). Defaults to None (auto-detect).
|
||||
backend: OpenCV backend identifier (https://docs.opencv.org/3.4/d4/d15/group__videoio__flags__base.html). Defaults to ANY.
|
||||
**Attributes**:
|
||||
- **index_or_path** (`int | Path`) -- Either an integer representing the camera device index, or a
|
||||
Path object pointing to a video file.
|
||||
- **fps** -- Requested frames per second for the color stream.
|
||||
- **width** -- Requested frame width in pixels for the color stream.
|
||||
- **height** -- Requested frame height in pixels for the color stream.
|
||||
- **color_mode** (`ColorMode`) -- Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
- **rotation** (`Cv2Rotation`) -- Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no
|
||||
rotation.
|
||||
- **warmup_s** (`int`) -- Time reading frames before returning from connect (in seconds)
|
||||
- **fourcc** (`str | None`) -- FOURCC code for video format (e.g., "MJPG", "YUYV", "I420"). Defaults
|
||||
to None (auto-detect).
|
||||
- **backend** (`Cv2Backends`) -- OpenCV backend identifier
|
||||
(https://docs.opencv.org/3.4/d4/d15/group__videoio__flags__base.html). Defaults to ANY.
|
||||
|
||||
Note:
|
||||
- Only 3-channel color output (RGB/BGR) is currently supported.
|
||||
|
||||
@@ -43,16 +43,16 @@ class Reachy2CameraConfig(CameraConfig):
|
||||
) # Left teleop camera, 640x480 @ 30FPS
|
||||
```
|
||||
|
||||
Attributes:
|
||||
name: Name of the camera device. Can be "teleop" or "depth".
|
||||
image_type: Type of image stream. For "teleop" camera, can be "left" or "right".
|
||||
For "depth" camera, can be "rgb" or "depth". (depth is not supported yet)
|
||||
fps: Requested frames per second for the color stream. Not configurable for Reachy 2 cameras.
|
||||
width: Requested frame width in pixels for the color stream.
|
||||
height: Requested frame height in pixels for the color stream.
|
||||
color_mode: Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
ip_address: IP address of the robot. Defaults to "localhost".
|
||||
port: Port number for the camera server. Defaults to 50065.
|
||||
**Attributes**:
|
||||
- **name** (`str`) -- Name of the camera device. Can be "teleop" or "depth".
|
||||
- **image_type** (`str`) -- Type of image stream. For "teleop" camera, can be "left" or "right". For
|
||||
"depth" camera, can be "rgb" or "depth". (depth is not supported yet)
|
||||
- **fps** -- Requested frames per second for the color stream. Not configurable for Reachy 2 cameras.
|
||||
- **width** -- Requested frame width in pixels for the color stream.
|
||||
- **height** -- Requested frame height in pixels for the color stream.
|
||||
- **color_mode** (`ColorMode`) -- Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
- **ip_address** (`str | None`) -- IP address of the robot. Defaults to "localhost".
|
||||
- **port** (`int`) -- Port number for the camera server. Defaults to 50065.
|
||||
|
||||
Note:
|
||||
- Only 3-channel color output (RGB/BGR) is currently supported.
|
||||
|
||||
@@ -109,6 +109,11 @@ class RealSenseCamera(Camera):
|
||||
```
|
||||
"""
|
||||
|
||||
# Maximum number of warmup attempts made by connect(). A failed attempt is first
|
||||
# retried with a plain pipeline stop/start, which is usually enough to recover the
|
||||
# stream; a USB hardware reset is performed before the final attempt as a last resort.
|
||||
_MAX_CONNECT_ATTEMPTS = 3
|
||||
|
||||
def __init__(self, config: RealSenseCameraConfig):
|
||||
"""
|
||||
Initializes the RealSenseCamera instance.
|
||||
@@ -173,6 +178,76 @@ class RealSenseCamera(Camera):
|
||||
"""Checks if the camera pipeline is started and streams are active."""
|
||||
return self.rs_pipeline is not None and self.rs_profile is not None
|
||||
|
||||
def _hardware_reset(self, wait_s: float = 5.0) -> None:
|
||||
"""Issue a USB hardware reset to recover an unresponsive device (common on D405)."""
|
||||
context = rs.context()
|
||||
for device in context.query_devices():
|
||||
if device.get_info(rs.camera_info.serial_number) == self.serial_number:
|
||||
logger.info(f"{self} performing hardware reset.")
|
||||
device.hardware_reset()
|
||||
time.sleep(wait_s)
|
||||
return
|
||||
logger.warning(f"{self} device not found for hardware reset, skipping.")
|
||||
|
||||
def _open_pipeline(self) -> None:
|
||||
"""Initializes the RealSense pipeline, starts it, and starts the background read thread.
|
||||
|
||||
Raises:
|
||||
ValueError: If the configuration is invalid, a requested sensor option is unsupported,
|
||||
or a requested sensor value is invalid.
|
||||
ConnectionError: If the camera is found but fails to start the pipeline or no RealSense devices are detected at all.
|
||||
RuntimeError: If the pipeline starts but fails to apply requested settings.
|
||||
"""
|
||||
rs_pipeline = rs.pipeline()
|
||||
rs_config = rs.config()
|
||||
self._configure_rs_pipeline_config(rs_config)
|
||||
|
||||
try:
|
||||
rs_profile = rs_pipeline.start(rs_config)
|
||||
except RuntimeError as e:
|
||||
raise ConnectionError(
|
||||
f"Failed to open {self}.Run `lerobot-find-cameras realsense` to find available cameras."
|
||||
) from e
|
||||
|
||||
self.rs_pipeline = rs_pipeline
|
||||
self.rs_profile = rs_profile
|
||||
|
||||
try:
|
||||
self._configure_capture_settings()
|
||||
self._configure_sensor_options()
|
||||
self._start_read_thread()
|
||||
except BaseException:
|
||||
self._release_after_failed_setup()
|
||||
raise
|
||||
|
||||
def _run_warmup(self) -> None:
|
||||
"""Blocks until at least one valid frame has been captured by the background thread.
|
||||
|
||||
Raises:
|
||||
ConnectionError: If no frame arrives before ``warmup_s`` elapses.
|
||||
"""
|
||||
# NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise.
|
||||
self.warmup_s = max(self.warmup_s, 1)
|
||||
|
||||
warmup_read = self.async_read if self.use_rgb else self.async_read_depth
|
||||
start_time = time.time()
|
||||
while time.time() - start_time < self.warmup_s:
|
||||
warmup_read(timeout_ms=self.warmup_s * 1000)
|
||||
time.sleep(0.1)
|
||||
with self.frame_lock:
|
||||
if (self.use_rgb and self.latest_color_frame is None) or (
|
||||
self.use_depth and self.latest_depth_frame is None
|
||||
):
|
||||
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
||||
|
||||
def _release_after_failed_setup(self) -> None:
|
||||
"""Releases the device handle and restores auto-detected settings after a failed attempt."""
|
||||
try:
|
||||
self._cleanup_resources()
|
||||
except Exception:
|
||||
logger.exception(f"Failed to fully clean up {self} after connect() failed.")
|
||||
self._reset_connection_settings()
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, warmup: bool = True) -> None:
|
||||
"""
|
||||
@@ -181,58 +256,53 @@ class RealSenseCamera(Camera):
|
||||
Initializes the RealSense pipeline, configures the required streams (color
|
||||
and optionally depth), starts the pipeline, and validates the actual stream settings.
|
||||
|
||||
If the pipeline starts but no frames arrive during warmup, retries up to
|
||||
``_MAX_CONNECT_ATTEMPTS`` times, performing a USB hardware reset before the
|
||||
final attempt.
|
||||
|
||||
Args:
|
||||
warmup (bool): If True, waits at connect() time until at least one valid frame
|
||||
has been captured by the background thread. Defaults to True.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the camera is already connected.
|
||||
ValueError: If the configuration is invalid, a requested sensor option is unsupported,
|
||||
or a requested sensor value is invalid.
|
||||
ValueError: If the configuration is invalid (e.g., missing serial/name, name not unique).
|
||||
ConnectionError: If the camera is found but fails to start the pipeline or no RealSense devices are detected at all.
|
||||
RuntimeError: If the pipeline starts but fails to apply requested settings.
|
||||
"""
|
||||
|
||||
self.rs_pipeline = rs.pipeline()
|
||||
rs_config = rs.config()
|
||||
self._configure_rs_pipeline_config(rs_config)
|
||||
if not warmup:
|
||||
self._open_pipeline()
|
||||
logger.info(f"{self} connected.")
|
||||
return
|
||||
|
||||
try:
|
||||
self.rs_profile = self.rs_pipeline.start(rs_config)
|
||||
except RuntimeError as e:
|
||||
self.rs_profile = None
|
||||
self.rs_pipeline = None
|
||||
raise ConnectionError(
|
||||
f"Failed to open {self}.Run `lerobot-find-cameras realsense` to find available cameras."
|
||||
) from e
|
||||
last_error: Exception | None = None
|
||||
|
||||
try:
|
||||
self._configure_capture_settings()
|
||||
self._configure_sensor_options()
|
||||
self._start_read_thread()
|
||||
for attempt in range(1, self._MAX_CONNECT_ATTEMPTS + 1):
|
||||
if attempt == self._MAX_CONNECT_ATTEMPTS:
|
||||
self._hardware_reset()
|
||||
|
||||
# NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise.
|
||||
self.warmup_s = max(self.warmup_s, 1)
|
||||
self._open_pipeline()
|
||||
|
||||
warmup_read = self.async_read if self.use_rgb else self.async_read_depth
|
||||
start_time = time.time()
|
||||
while time.time() - start_time < self.warmup_s:
|
||||
warmup_read(timeout_ms=self.warmup_s * 1000)
|
||||
time.sleep(0.1)
|
||||
with self.frame_lock:
|
||||
if (self.use_rgb and self.latest_color_frame is None) or (
|
||||
self.use_depth and self.latest_depth_frame is None
|
||||
):
|
||||
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
||||
except BaseException:
|
||||
connected = False
|
||||
try:
|
||||
self._cleanup_resources()
|
||||
except Exception:
|
||||
logger.exception(f"Failed to fully clean up {self} after connect() failed.")
|
||||
self._reset_connection_settings()
|
||||
raise
|
||||
self._run_warmup()
|
||||
connected = True
|
||||
except (TimeoutError, ConnectionError) as e:
|
||||
last_error = e
|
||||
finally:
|
||||
if not connected:
|
||||
self._release_after_failed_setup()
|
||||
|
||||
logger.info(f"{self} connected.")
|
||||
if connected:
|
||||
logger.info(f"{self} connected.")
|
||||
return
|
||||
|
||||
logger.warning(f"{self} warmup failed (attempt {attempt}/{self._MAX_CONNECT_ATTEMPTS}).")
|
||||
|
||||
raise ConnectionError(
|
||||
f"{self} failed to capture frames after {self._MAX_CONNECT_ATTEMPTS} attempts."
|
||||
) from last_error
|
||||
|
||||
@staticmethod
|
||||
def find_cameras() -> list[dict[str, Any]]:
|
||||
@@ -629,6 +699,9 @@ class RealSenseCamera(Camera):
|
||||
capture_time = time.perf_counter()
|
||||
|
||||
with self.frame_lock:
|
||||
# Under the lock, so a late frame cannot resurrect the buffer _stop_read_thread() cleared.
|
||||
if stop_event.is_set():
|
||||
break
|
||||
if self.use_rgb:
|
||||
self.latest_color_frame = processed_color_frame
|
||||
if self.use_depth:
|
||||
@@ -839,4 +912,5 @@ class RealSenseCamera(Camera):
|
||||
)
|
||||
|
||||
self._cleanup_resources()
|
||||
|
||||
logger.info(f"{self} disconnected.")
|
||||
|
||||
@@ -36,27 +36,28 @@ class RealSenseCameraConfig(CameraConfig):
|
||||
RealSenseCameraConfig("0123456789", 30, 640, 480, rotation=Cv2Rotation.ROTATE_90) # With 90° rotation
|
||||
```
|
||||
|
||||
Attributes:
|
||||
fps: Requested frames per second for the color stream.
|
||||
width: Requested frame width in pixels for the color stream.
|
||||
height: Requested frame height in pixels for the color stream.
|
||||
serial_number_or_name: Unique serial number or human-readable name to identify the camera.
|
||||
color_mode: Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
use_rgb: Whether to enable the color stream. Defaults to True.
|
||||
use_depth: Whether to enable depth stream. Defaults to False.
|
||||
rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation.
|
||||
warmup_s: Time reading frames before returning from connect (in seconds)
|
||||
exposure: Manual exposure value for the color sensor. When set, auto-exposure is
|
||||
disabled and this fixed value is used. Valid ranges are camera-model specific
|
||||
and reported if the value is rejected. Defaults to None (leave unchanged).
|
||||
gain: Manual gain value for the color sensor. When set, auto-exposure is disabled
|
||||
and this fixed gain is used, which also freezes exposure at its current value
|
||||
when no exposure is configured. Valid ranges are camera-model specific and
|
||||
reported if the value is rejected. Defaults to None (leave unchanged).
|
||||
white_balance: Manual white balance value for the color sensor. When set, auto
|
||||
white balance is disabled and this fixed value is used. Valid ranges are
|
||||
camera-model specific and reported if the value is rejected. Defaults to None
|
||||
(leave unchanged).
|
||||
**Attributes**:
|
||||
- **fps** -- Requested frames per second for the color stream.
|
||||
- **width** -- Requested frame width in pixels for the color stream.
|
||||
- **height** -- Requested frame height in pixels for the color stream.
|
||||
- **serial_number_or_name** (`str`) -- Unique serial number or human-readable name to identify the
|
||||
camera.
|
||||
- **color_mode** (`ColorMode`) -- Color mode for image output (RGB or BGR). Defaults to RGB.
|
||||
- **use_rgb** (`bool`) -- Whether to enable the color stream. Defaults to True.
|
||||
- **use_depth** (`bool`) -- Whether to enable depth stream. Defaults to False.
|
||||
- **rotation** (`Cv2Rotation`) -- Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no
|
||||
rotation.
|
||||
- **warmup_s** (`int`) -- Time reading frames before returning from connect (in seconds)
|
||||
- **exposure** (`int | None`) -- Manual exposure value for the color sensor. When set, auto-exposure
|
||||
is disabled and this fixed value is used. Valid ranges are camera-model specific and reported if the
|
||||
value is rejected. Defaults to None (leave unchanged).
|
||||
- **gain** (`int | None`) -- Manual gain value for the color sensor. When set, auto-exposure is
|
||||
disabled and this fixed gain is used, which also freezes exposure at its current value when no
|
||||
exposure is configured. Valid ranges are camera-model specific and reported if the value is
|
||||
rejected. Defaults to None (leave unchanged).
|
||||
- **white_balance** (`int | None`) -- Manual white balance value for the color sensor. When set, auto
|
||||
white balance is disabled and this fixed value is used. Valid ranges are camera-model specific and
|
||||
reported if the value is rejected. Defaults to None (leave unchanged).
|
||||
|
||||
Note:
|
||||
- Either name or serial_number must be specified.
|
||||
|
||||
@@ -102,6 +102,7 @@ class ImageServer:
|
||||
fps=self.fps,
|
||||
width=shape[1],
|
||||
height=shape[0],
|
||||
fourcc=cfg.get("fourcc", "MJPG"),
|
||||
color_mode=ColorMode.RGB,
|
||||
)
|
||||
camera = OpenCVCamera(cam_config)
|
||||
|
||||
+604
-174
@@ -13,16 +13,41 @@
|
||||
# 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 pathlib import Path
|
||||
"""Training-output persistence: checkpoints, two-phase resume, and hub publishing.
|
||||
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
Rank discipline: every function here that can
|
||||
contain a collective is documented as such and must run on ALL ranks; rank-0-only file writes
|
||||
sit under one grouped ``is_main_process()`` gate per contiguous region, placed below all
|
||||
collectives. The leaf save/load helpers carry no rank gates of their own — the exception is
|
||||
``PreTrainedPolicy._save_pretrained``, whose gate is internal because its collective gather and
|
||||
its writes live in the same method.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from importlib.resources import files
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch.distributed as dist
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, snapshot_download
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
|
||||
from lerobot.__version__ import __version__
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.configs.rewards import RewardModelConfig
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.distributed.checkpoint import (
|
||||
is_sharded_module,
|
||||
load_sharded_model,
|
||||
load_sharded_optimizer,
|
||||
save_sharded_model,
|
||||
save_sharded_optimizer,
|
||||
)
|
||||
from lerobot.distributed.utils import is_main_process
|
||||
from lerobot.optim import (
|
||||
load_optimizer_state,
|
||||
load_optimizer_state_dict,
|
||||
load_scheduler_state,
|
||||
save_optimizer_state,
|
||||
save_scheduler_state,
|
||||
@@ -40,14 +65,39 @@ from lerobot.utils.hub import find_latest_hub_checkpoint
|
||||
from lerobot.utils.io_utils import load_json, write_json
|
||||
from lerobot.utils.random_utils import load_rng_state, save_rng_state
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from accelerate import Accelerator
|
||||
|
||||
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
||||
from lerobot.rewards.pretrained import PreTrainedRewardModel
|
||||
|
||||
|
||||
def get_step_identifier(step: int, total_steps: int) -> str:
|
||||
"""Format a step number as the zero-padded identifier used for checkpoint directory names.
|
||||
|
||||
Args:
|
||||
step (int): The training step to format.
|
||||
total_steps (int): The total number of training steps; sets the padding width
|
||||
(minimum 6 digits).
|
||||
|
||||
Returns:
|
||||
str: The zero-padded step identifier, e.g. `"005000"`.
|
||||
"""
|
||||
num_digits = max(6, len(str(total_steps)))
|
||||
return f"{step:0{num_digits}d}"
|
||||
|
||||
|
||||
def get_step_checkpoint_dir(output_dir: Path, total_steps: int, step: int) -> Path:
|
||||
"""Returns the checkpoint sub-directory corresponding to the step number."""
|
||||
"""Returns the checkpoint sub-directory corresponding to the step number.
|
||||
|
||||
Args:
|
||||
output_dir (Path): The training run's output directory.
|
||||
total_steps (int): The total number of training steps; sets the identifier padding.
|
||||
step (int): The training step of the checkpoint.
|
||||
|
||||
Returns:
|
||||
Path: The checkpoint step directory, `output_dir/checkpoints/<step-identifier>`.
|
||||
"""
|
||||
step_identifier = get_step_identifier(step, total_steps)
|
||||
return output_dir / CHECKPOINTS_DIR / step_identifier
|
||||
|
||||
@@ -63,37 +113,15 @@ def should_save_checkpoint(step: int, save_freq: int, total_steps: int) -> bool:
|
||||
return (save_freq > 0 and step % save_freq == 0) or step == total_steps
|
||||
|
||||
|
||||
def save_training_step(
|
||||
step: int, save_dir: Path, num_processes: int | None = None, batch_size: int | None = None
|
||||
) -> None:
|
||||
state: dict = {"step": step}
|
||||
# num_processes and batch_size are recorded so a resumed run can detect a changed world size or
|
||||
# batch size: the sampler's resume offset is computed from the (num_processes, batch_size) that
|
||||
# produced `step`, since both scale how many sampler positions a step consumes (see
|
||||
# compute_sampler_state).
|
||||
if num_processes is not None:
|
||||
state["num_processes"] = num_processes
|
||||
if batch_size is not None:
|
||||
state["batch_size"] = batch_size
|
||||
write_json(state, save_dir / TRAINING_STEP)
|
||||
def update_last_checkpoint(checkpoint_dir: Path) -> None:
|
||||
"""Point the `last` symlink in the checkpoints directory at the given checkpoint.
|
||||
|
||||
Any existing `last` symlink is replaced. The link target is relative to the checkpoints
|
||||
directory, so the tree stays valid when the run directory is moved.
|
||||
|
||||
def load_training_step(save_dir: Path) -> int:
|
||||
training_step = load_json(save_dir / TRAINING_STEP)
|
||||
return training_step["step"]
|
||||
|
||||
|
||||
def load_training_num_processes(checkpoint_dir: Path) -> int | None:
|
||||
"""World size recorded at checkpoint time, or None for checkpoints written before it was stored."""
|
||||
return load_json(checkpoint_dir / TRAINING_STATE_DIR / TRAINING_STEP).get("num_processes")
|
||||
|
||||
|
||||
def load_training_batch_size(checkpoint_dir: Path) -> int | None:
|
||||
"""Per-process batch size recorded at checkpoint time, or None for older checkpoints."""
|
||||
return load_json(checkpoint_dir / TRAINING_STATE_DIR / TRAINING_STEP).get("batch_size")
|
||||
|
||||
|
||||
def update_last_checkpoint(checkpoint_dir: Path) -> Path:
|
||||
Args:
|
||||
checkpoint_dir (Path): The checkpoint step directory the `last` link should target.
|
||||
"""
|
||||
last_checkpoint_dir = checkpoint_dir.parent / LAST_CHECKPOINT_LINK
|
||||
if last_checkpoint_dir.is_symlink():
|
||||
last_checkpoint_dir.unlink()
|
||||
@@ -101,6 +129,68 @@ def update_last_checkpoint(checkpoint_dir: Path) -> Path:
|
||||
last_checkpoint_dir.symlink_to(relative_target)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# training_step.json
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
|
||||
|
||||
def save_training_metadata(step: int, save_dir: Path, cfg: TrainPipelineConfig) -> None:
|
||||
"""Record the step counter plus everything a resume needs to reason about topology changes.
|
||||
|
||||
`step` counts loop iterations (= micro-batches), so
|
||||
the sampler resume offset is `step x batch_size x dp_world_size` with no grad-accum factor.
|
||||
`grad_accum_steps` and the parallelism snapshot are recorded so a resume can warn precisely
|
||||
when the optimizer-update cadence or the sharding topology changed.
|
||||
|
||||
Args:
|
||||
step (int): The training step (micro-batch counter) to record.
|
||||
save_dir (Path): The `training_state/` directory to write `training_step.json` into.
|
||||
cfg (TrainPipelineConfig): The training config whose batch size, gradient-accumulation,
|
||||
and parallelism settings are snapshotted alongside the step.
|
||||
"""
|
||||
state: dict[str, Any] = {
|
||||
"step": step,
|
||||
"dp_world_size": cfg.parallelism.dp_world_size,
|
||||
"batch_size": cfg.batch_size,
|
||||
"grad_accum_steps": cfg.accelerator.gradient_accumulation.steps,
|
||||
"parallelism": {
|
||||
"dp_replicate": cfg.parallelism.dp_replicate,
|
||||
"dp_shard": cfg.parallelism.dp_shard,
|
||||
"ring_degree": cfg.parallelism.context_parallel.ring_degree,
|
||||
"ulysses_degree": cfg.parallelism.context_parallel.ulysses_degree,
|
||||
},
|
||||
}
|
||||
write_json(state, save_dir / TRAINING_STEP)
|
||||
|
||||
|
||||
def load_training_metadata(training_state_dir: Path) -> dict[str, Any]:
|
||||
"""Read everything `save_training_metadata` recorded, in a single pass.
|
||||
|
||||
Every key is always present: fields a checkpoint predates come back as None, so a caller
|
||||
reading `metadata["batch_size"]` gets a KeyError on a typo rather than a silent None.
|
||||
|
||||
Args:
|
||||
training_state_dir (Path): The checkpoint's `training_state/` directory.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: `step` plus the `dp_world_size`, `batch_size`, `grad_accum_steps` and
|
||||
`parallelism` snapshot recorded alongside it (None where not recorded).
|
||||
"""
|
||||
state = load_json(training_state_dir / TRAINING_STEP)
|
||||
return {
|
||||
"step": int(state["step"]),
|
||||
"dp_world_size": state.get("dp_world_size", state.get("num_processes")),
|
||||
"batch_size": state.get("batch_size"),
|
||||
"grad_accum_steps": state.get("grad_accum_steps"),
|
||||
"parallelism": state.get("parallelism"),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# Checkpoint save
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
|
||||
|
||||
def save_checkpoint(
|
||||
checkpoint_dir: Path,
|
||||
step: int,
|
||||
@@ -110,192 +200,301 @@ def save_checkpoint(
|
||||
scheduler: LRScheduler | None = None,
|
||||
preprocessor: PolicyProcessorPipeline | None = None,
|
||||
postprocessor: PolicyProcessorPipeline | None = None,
|
||||
num_processes: int | None = None,
|
||||
batch_size: int | None = None,
|
||||
model_state_dict: dict | None = None,
|
||||
optim_state_dict: dict | None = None,
|
||||
accelerator: "Accelerator | None" = None,
|
||||
) -> None:
|
||||
"""This function creates the following directory structure:
|
||||
|
||||
005000/ # training step at checkpoint
|
||||
├── pretrained_model/
|
||||
│ ├── config.json # policy config
|
||||
│ ├── model.safetensors # policy weights
|
||||
│ ├── model.safetensors # policy weights (checkpoint_format ∈ {safetensors, safetensors_dcp}, or any non-sharded run)
|
||||
│ ├── pytorch_model_fsdp_0/ # DCP model shards (checkpoint_format ∈ {dcp, safetensors_dcp})
|
||||
│ ├── train_config.json # train config
|
||||
│ ├── processor.json # processor config (if preprocessor provided)
|
||||
│ └── step_*.safetensors # processor state files (if any)
|
||||
│ ├── policy_preprocessor.json # preprocessor config (if preprocessor provided)
|
||||
│ ├── policy_preprocessor_step_*.safetensors # state of the stateful preprocessor steps
|
||||
│ ├── policy_postprocessor.json # postprocessor config (if postprocessor provided)
|
||||
│ └── policy_postprocessor_step_*.safetensors # state of the stateful postprocessor steps
|
||||
└── training_state/
|
||||
├── optimizer_param_groups.json # optimizer param groups
|
||||
├── optimizer_state.safetensors # optimizer state
|
||||
├── optimizer_param_groups.json # optimizer param groups (non-sharded runs)
|
||||
├── optimizer_state.safetensors # optimizer state (non-sharded runs)
|
||||
├── optimizer_0/ # DCP optimizer shards (sharded runs)
|
||||
├── rng_state.safetensors # rng states
|
||||
├── scheduler_state.json # scheduler state
|
||||
└── training_step.json # training step
|
||||
├── scheduler_state.json # scheduler state (if scheduler provided)
|
||||
└── training_step.json # training step + dp_world_size/batch_size/grad_accum + topology
|
||||
|
||||
Collective: MUST be called on every rank. Rank-0-only writes are gated internally, so the
|
||||
call site needs no rank branches.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The training config used for this run.
|
||||
checkpoint_dir (Path): The checkpoint step directory to write (e.g. `.../checkpoints/005000`).
|
||||
step (int): The training step at that checkpoint.
|
||||
cfg (TrainPipelineConfig): The training config used for this run.
|
||||
policy (PreTrainedPolicy): The policy to save.
|
||||
optimizer (Optimizer | None, optional): The optimizer to save the state from. Defaults to None.
|
||||
optimizer (Optimizer): The optimizer to save the state from.
|
||||
scheduler (LRScheduler | None, optional): The scheduler to save the state from. Defaults to None.
|
||||
preprocessor: The preprocessor/pipeline to save. Defaults to None.
|
||||
postprocessor: The postprocessor/pipeline to save. Defaults to None.
|
||||
num_processes (int | None, optional): Distributed world size to record for sample-exact
|
||||
resume. Defaults to None (not recorded).
|
||||
batch_size (int | None, optional): Per-process batch size to record for sample-exact
|
||||
resume. Defaults to None (not recorded).
|
||||
model_state_dict: Pre-gathered full (unsharded) model state dict. Required under FSDP,
|
||||
where `policy.state_dict()` would return sharded tensors; the caller gathers it via a
|
||||
cross-rank collective and passes it here so rank 0 can write it directly. It holds
|
||||
FSDP's fp32 master weights and is saved as-is (the loader casts to the policy dtype on
|
||||
read). When None (DDP / single-GPU), the model is saved the normal way. Defaults to None.
|
||||
optim_state_dict: Pre-gathered full (unsharded) optimizer state dict. Required under FSDP
|
||||
(gathered alongside `model_state_dict` via `gather_fsdp_state_dicts`); saved in the same
|
||||
safetensors format as the single-GPU path. When None, `optimizer.state_dict()` is used.
|
||||
preprocessor (PolicyProcessorPipeline | None, optional): The preprocessor/pipeline to save.
|
||||
Defaults to None.
|
||||
postprocessor (PolicyProcessorPipeline | None, optional): The postprocessor/pipeline to save.
|
||||
Defaults to None.
|
||||
accelerator (Accelerator | None, optional): The accelerator the policy was prepared with;
|
||||
used to unwrap the model and required on sharded runs, where it owns the DCP save
|
||||
channels. Defaults to None (plain single-process saves).
|
||||
"""
|
||||
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
|
||||
policy.save_pretrained(pretrained_dir, state_dict=model_state_dict)
|
||||
cfg.save_pretrained(pretrained_dir)
|
||||
fmt = cfg.checkpoint_format
|
||||
policy_to_save = accelerator.unwrap_model(policy) if accelerator is not None else policy
|
||||
sharded = is_sharded_module(policy_to_save)
|
||||
|
||||
# -- model artifact(s): the two collective-capable calls ----------------------------------
|
||||
if cfg.peft is not None:
|
||||
# When using PEFT, policy.save_pretrained will only write the adapter weights + config, not the
|
||||
# policy config which we need for loading the model. In this case we'll write it ourselves.
|
||||
policy.config.save_pretrained(pretrained_dir)
|
||||
if preprocessor is not None:
|
||||
preprocessor.save_pretrained(pretrained_dir)
|
||||
if postprocessor is not None:
|
||||
postprocessor.save_pretrained(pretrained_dir)
|
||||
# PeftModel.save_pretrained is an external API with no internal rank gate, and the
|
||||
# adapters are replicated (PEFT x sharded is rejected at validation): main rank writes.
|
||||
if is_main_process():
|
||||
policy_to_save.save_pretrained(pretrained_dir)
|
||||
elif fmt.wants_safetensors or not sharded:
|
||||
# Collective when sharded (full gather); writes happen on the main process only in all
|
||||
# multi-rank layouts (the gate lives inside _save_pretrained, next to its collective gather).
|
||||
policy_to_save.save_pretrained(pretrained_dir)
|
||||
if fmt.wants_dcp and sharded:
|
||||
save_sharded_model(accelerator, policy_to_save, pretrained_dir)
|
||||
|
||||
# -- sidecar configs: ONE gate for the whole contiguous rank-0-only region ----------------
|
||||
if is_main_process():
|
||||
if fmt.wants_dcp and not fmt.wants_safetensors:
|
||||
# save_pretrained did not run: keep the DCP-only checkpoint self-describing.
|
||||
policy_to_save.config.save_pretrained(pretrained_dir)
|
||||
cfg.save_pretrained(pretrained_dir)
|
||||
if cfg.peft is not None:
|
||||
# PEFT's save_pretrained writes only adapter weights + config; the policy config
|
||||
# needed to reload the base model is written explicitly.
|
||||
policy_to_save.config.save_pretrained(pretrained_dir)
|
||||
if preprocessor is not None:
|
||||
preprocessor.save_pretrained(pretrained_dir)
|
||||
if postprocessor is not None:
|
||||
postprocessor.save_pretrained(pretrained_dir)
|
||||
|
||||
save_training_state(
|
||||
checkpoint_dir,
|
||||
step,
|
||||
optimizer,
|
||||
scheduler,
|
||||
num_processes=num_processes,
|
||||
batch_size=batch_size,
|
||||
optim_state_dict=optim_state_dict,
|
||||
checkpoint_dir, step, cfg, optimizer, scheduler, accelerator, sharded=sharded, model=policy_to_save
|
||||
)
|
||||
if accelerator is not None:
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
|
||||
def save_training_state(
|
||||
checkpoint_dir: Path,
|
||||
train_step: int,
|
||||
optimizer: Optimizer | None = None,
|
||||
step: int,
|
||||
cfg: TrainPipelineConfig,
|
||||
optimizer: Optimizer | dict[str, Optimizer] | None = None,
|
||||
scheduler: LRScheduler | None = None,
|
||||
num_processes: int | None = None,
|
||||
batch_size: int | None = None,
|
||||
optim_state_dict: dict | None = None,
|
||||
accelerator: "Accelerator | None" = None,
|
||||
*,
|
||||
sharded: bool = False,
|
||||
model: PreTrainedPolicy | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Saves the training step, optimizer state, scheduler state, and rng state.
|
||||
"""Write training_state/. Collective under sharding: call on every rank.
|
||||
|
||||
Args:
|
||||
save_dir (Path): The directory to save artifacts to.
|
||||
train_step (int): Current training step.
|
||||
optimizer (Optimizer | None, optional): The optimizer from which to save the state_dict.
|
||||
checkpoint_dir (Path): The checkpoint step directory; `training_state/` is created inside it.
|
||||
step (int): The training step at that checkpoint.
|
||||
cfg (TrainPipelineConfig): The training config used for this run (its topology and
|
||||
accumulation settings are recorded in `training_step.json`).
|
||||
optimizer (Optimizer | dict[str, Optimizer] | None, optional): The optimizer(s) to save
|
||||
the state from. Defaults to None.
|
||||
scheduler (LRScheduler | None, optional): The scheduler to save the state from.
|
||||
Defaults to None.
|
||||
scheduler (LRScheduler | None, optional): The scheduler from which to save the state_dict.
|
||||
Defaults to None.
|
||||
num_processes (int | None, optional): Distributed world size to record. Defaults to None.
|
||||
batch_size (int | None, optional): Per-process batch size to record. Defaults to None.
|
||||
optim_state_dict: Pre-gathered full optimizer state dict (for FSDP). Saved instead of
|
||||
`optimizer.state_dict()` when provided. Defaults to None.
|
||||
accelerator (Accelerator | None, optional): Required when `sharded` is True — it owns
|
||||
the DCP optimizer save channel. Defaults to None.
|
||||
sharded (bool): The model's sharding state, computed once in `save_checkpoint` and
|
||||
threaded here so the two sites cannot disagree. Defaults to False.
|
||||
model (PreTrainedPolicy | None, optional): Required only for the sharded optimizer
|
||||
channel: torch's optimizer DCP APIs are model-coupled (the state dict is keyed by
|
||||
model FQNs), so accelerate's `save_fsdp_optimizer` needs the sharded module
|
||||
alongside the optimizer. Defaults to None.
|
||||
"""
|
||||
save_dir = checkpoint_dir / TRAINING_STATE_DIR
|
||||
# All ranks: the directory must exist before the DCP optimizer collective writes into it
|
||||
# (exist_ok makes the concurrent mkdir race-free on shared filesystems).
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_training_step(train_step, save_dir, num_processes=num_processes, batch_size=batch_size)
|
||||
save_rng_state(save_dir)
|
||||
if optimizer is not None:
|
||||
save_optimizer_state(optimizer, save_dir, optim_state_dict=optim_state_dict)
|
||||
if scheduler is not None:
|
||||
save_scheduler_state(scheduler, save_dir)
|
||||
|
||||
if optimizer is not None and sharded:
|
||||
if accelerator is None or model is None:
|
||||
raise ValueError("Saving a sharded optimizer state requires the accelerator and model.")
|
||||
# Collective — all ranks write their DCP shards into optimizer_0/.
|
||||
save_sharded_optimizer(accelerator, optimizer, model, save_dir)
|
||||
|
||||
if is_main_process(): # ONE grouped gate for the whole rank-0-only region
|
||||
save_training_metadata(step, save_dir, cfg)
|
||||
save_rng_state(save_dir)
|
||||
if scheduler is not None:
|
||||
save_scheduler_state(scheduler, save_dir)
|
||||
if optimizer is not None and not sharded:
|
||||
save_optimizer_state(optimizer, save_dir)
|
||||
|
||||
|
||||
def load_training_state(
|
||||
checkpoint_dir: Path, optimizer: Optimizer, scheduler: LRScheduler | None, load_optimizer: bool = True
|
||||
) -> tuple[int, Optimizer, LRScheduler | None]:
|
||||
"""
|
||||
Loads the training step, optimizer state, scheduler state, and rng state.
|
||||
This is used to resume a training run.
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# Two-phase resume
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
|
||||
|
||||
def resume_before_prepare(cfg: TrainPipelineConfig) -> int:
|
||||
"""Phase 1 — before `accelerator.prepare()`: restore RNG and return the step counter.
|
||||
|
||||
Pure loaders only. The sampler resume offset is *derived* from the returned step inside the
|
||||
dataloader factory, and everything bound to sharded objects (model DCP shards, optimizer,
|
||||
scheduler) loads in `resume_after_prepare`.
|
||||
|
||||
Args:
|
||||
checkpoint_dir (Path): The checkpoint directory. Should contain a 'training_state' dir.
|
||||
optimizer (Optimizer): The optimizer to load the state_dict to.
|
||||
scheduler (LRScheduler | None): The scheduler to load the state_dict to (can be None).
|
||||
load_optimizer (bool, optional): Whether to load the optimizer state from disk. Defaults to
|
||||
True. Set to False under FSDP, where the sharded optimizer state must be loaded after
|
||||
`accelerator.prepare()` via `load_fsdp_optimizer_state` (the optimizer is returned
|
||||
untouched here).
|
||||
cfg (TrainPipelineConfig): The resumed training config; `cfg.checkpoint_path` locates
|
||||
the checkpoint to restore from.
|
||||
|
||||
Returns:
|
||||
int: The training step recorded in the checkpoint (micro-batch counter).
|
||||
|
||||
Raises:
|
||||
NotADirectoryError: If 'checkpoint_dir' doesn't contain a 'training_state' dir
|
||||
|
||||
Returns:
|
||||
tuple[int, Optimizer, LRScheduler | None]: training step, optimizer and scheduler with their
|
||||
state_dict loaded.
|
||||
NotADirectoryError: If the checkpoint has no `training_state/` directory.
|
||||
ValueError: If the resumed topology crosses the sharded/non-sharded boundary relative
|
||||
to the one recorded in the checkpoint.
|
||||
"""
|
||||
training_state_dir = checkpoint_dir / TRAINING_STATE_DIR
|
||||
training_state_dir = cfg.checkpoint_path / TRAINING_STATE_DIR
|
||||
if not training_state_dir.is_dir():
|
||||
raise NotADirectoryError(training_state_dir)
|
||||
|
||||
metadata = load_training_metadata(training_state_dir)
|
||||
_guard_resume_changes(cfg, metadata)
|
||||
load_rng_state(training_state_dir)
|
||||
step = load_training_step(training_state_dir)
|
||||
if load_optimizer:
|
||||
optimizer = load_optimizer_state(optimizer, training_state_dir)
|
||||
return metadata["step"]
|
||||
|
||||
|
||||
def _guard_resume_changes(cfg: TrainPipelineConfig, metadata: dict[str, Any]) -> None:
|
||||
"""Check the resumed run settings against the ones recorded in the checkpoint.
|
||||
|
||||
Two tiers, both driven by the checkpoint's recorded parallelism snapshot:
|
||||
|
||||
- **Hard error** when the resume crosses the sharded/non-sharded boundary in either
|
||||
direction: the checkpoint's training-state artifacts only support resuming on the same
|
||||
kind of topology (resharding works across sizes, not across kinds). Checkpoints without
|
||||
a recorded snapshot skip this check.
|
||||
- **One warning** naming every other recorded setting that differs — those changes are
|
||||
legal (DCP reshards weights and optimizer state across topologies and the sampler offset
|
||||
adapts), but a changed ``grad_accum_steps`` shifts the optimizer-update cadence, so the
|
||||
resume says precisely what differs. The sampler-exactness warnings
|
||||
(``dp_world_size``/``batch_size``) live with the sampler math in the dataloader factory.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The resumed training config, compared against the settings
|
||||
recorded in the checkpoint.
|
||||
metadata (dict[str, Any]): The checkpoint's recorded training metadata, as returned by
|
||||
`load_training_metadata`.
|
||||
|
||||
Raises:
|
||||
ValueError: If the checkpoint records a sharded topology and the resumed run is
|
||||
non-sharded, or vice versa.
|
||||
"""
|
||||
snapshot = metadata["parallelism"]
|
||||
|
||||
if snapshot is not None:
|
||||
recorded_sharded = (
|
||||
snapshot.get("dp_shard", 1) != 1
|
||||
or snapshot.get("ring_degree", 1) * snapshot.get("ulysses_degree", 1) > 1
|
||||
)
|
||||
if recorded_sharded != cfg.parallelism.is_sharded:
|
||||
raise ValueError(
|
||||
f"Cannot resume: the checkpoint was written with a "
|
||||
f"{'sharded' if recorded_sharded else 'non-sharded'} topology "
|
||||
f"(dp_replicate={snapshot.get('dp_replicate')}, dp_shard={snapshot.get('dp_shard')}) "
|
||||
f"but this run is {'sharded' if cfg.parallelism.is_sharded else 'non-sharded'} "
|
||||
f"(dp_replicate={cfg.parallelism.dp_replicate}, dp_shard={cfg.parallelism.dp_shard})."
|
||||
)
|
||||
|
||||
recorded = {
|
||||
"grad_accum_steps": (
|
||||
metadata["grad_accum_steps"],
|
||||
cfg.accelerator.gradient_accumulation.steps,
|
||||
),
|
||||
}
|
||||
if snapshot is not None:
|
||||
recorded.update(
|
||||
{
|
||||
"dp_replicate": (snapshot.get("dp_replicate"), cfg.parallelism.dp_replicate),
|
||||
"dp_shard": (snapshot.get("dp_shard"), cfg.parallelism.dp_shard),
|
||||
"ring_degree": (
|
||||
snapshot.get("ring_degree"),
|
||||
cfg.parallelism.context_parallel.ring_degree,
|
||||
),
|
||||
"ulysses_degree": (
|
||||
snapshot.get("ulysses_degree"),
|
||||
cfg.parallelism.context_parallel.ulysses_degree,
|
||||
),
|
||||
}
|
||||
)
|
||||
changed = [f"{key}: {was} -> {now}" for key, (was, now) in recorded.items() if was not in (None, now)]
|
||||
if changed and is_main_process():
|
||||
logging.warning(
|
||||
"Resuming with settings that differ from the checkpoint: " + "; ".join(changed) + ". "
|
||||
"Topology changes reshard safely via DCP; a changed grad_accum_steps shifts the "
|
||||
"optimizer-update cadence (the step counter keeps counting micro-batches)."
|
||||
)
|
||||
|
||||
|
||||
def resume_after_prepare(
|
||||
cfg: TrainPipelineConfig,
|
||||
accelerator: "Accelerator",
|
||||
policy: PreTrainedPolicy,
|
||||
optimizer: Optimizer | dict[str, Optimizer],
|
||||
scheduler: LRScheduler | None,
|
||||
) -> None:
|
||||
"""Phase 2 — after `accelerator.prepare()`: model (DCP) -> optimizer -> scheduler.
|
||||
|
||||
Collective under sharding: call on every rank. The model-weight source follows the
|
||||
checkpoint's own recorded `checkpoint_format` (on resume, `cfg` was parsed from the
|
||||
checkpoint's train_config.json): DCP-bearing formats load shards here into the prepared
|
||||
model (whose construction skipped the safetensors load); the safetensors format was already
|
||||
loaded by `from_pretrained` before sharding — no model step here.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The resumed training config; `cfg.checkpoint_path` locates
|
||||
the checkpoint and `cfg.checkpoint_format` selects the model-weight source.
|
||||
accelerator (Accelerator): The accelerator the policy was prepared with; it unwraps the
|
||||
model and owns the DCP load channels.
|
||||
policy (PreTrainedPolicy): The prepared (possibly sharded) policy to load weights into.
|
||||
optimizer (Optimizer | dict[str, Optimizer]): The prepared optimizer(s) to restore.
|
||||
scheduler (LRScheduler | None): The scheduler to restore, or None if the run has none.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the checkpoint format declares DCP model shards but the shard
|
||||
directory is missing (e.g. it was pruned before upload).
|
||||
"""
|
||||
checkpoint_dir = cfg.checkpoint_path
|
||||
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
|
||||
training_state_dir = checkpoint_dir / TRAINING_STATE_DIR
|
||||
unwrapped = accelerator.unwrap_model(policy)
|
||||
sharded = is_sharded_module(unwrapped)
|
||||
|
||||
if cfg.checkpoint_format.wants_dcp:
|
||||
from accelerate.utils.constants import FSDP_MODEL_NAME
|
||||
|
||||
dcp_dir = pretrained_dir / f"{FSDP_MODEL_NAME}_0"
|
||||
if not dcp_dir.is_dir():
|
||||
raise FileNotFoundError(
|
||||
f"checkpoint_format={cfg.checkpoint_format.value} declares DCP model shards, "
|
||||
f"but {dcp_dir} is missing. If the shards were pruned, convert what remains "
|
||||
"with `lerobot-convert-dcp` or resume from a safetensors checkpoint."
|
||||
)
|
||||
load_sharded_model(accelerator, unwrapped, pretrained_dir)
|
||||
|
||||
if sharded:
|
||||
# Requires the prepared optimizer: FSDP2's prepare rebinds param groups to DTensors but
|
||||
# never migrates optimizer.state — DCP reshards it here (works across topology changes).
|
||||
load_sharded_optimizer(accelerator, optimizer, unwrapped, training_state_dir)
|
||||
else:
|
||||
load_optimizer_state(optimizer, training_state_dir)
|
||||
|
||||
if scheduler is not None:
|
||||
scheduler = load_scheduler_state(scheduler, training_state_dir)
|
||||
|
||||
return step, optimizer, scheduler
|
||||
load_scheduler_state(scheduler, training_state_dir)
|
||||
|
||||
|
||||
def gather_fsdp_state_dicts(model, optimizer) -> tuple[dict, dict]:
|
||||
"""Gather the full (unsharded) model and optimizer state dicts under FSDP.
|
||||
|
||||
`model.state_dict()` and `FSDP.optim_state_dict(...)` are cross-rank collectives, so this must be
|
||||
called on *every* rank with the prepared (FSDP-wrapped) `model` and `optimizer`. With
|
||||
`rank0_only=True` and `offload_to_cpu=True`, every rank runs the all-gather but only rank 0
|
||||
materializes the full dicts (the others get empty dicts) and they are kept on CPU to bound GPU
|
||||
memory. The returned optimizer state dict is keyed by parameter FQNs and is world-size
|
||||
independent; `load_fsdp_optimizer_state` reshards it on resume.
|
||||
|
||||
Returns:
|
||||
(model_state_dict, optim_state_dict): full dicts on rank 0, empty dicts on other ranks.
|
||||
"""
|
||||
from torch.distributed.fsdp import (
|
||||
FullOptimStateDictConfig,
|
||||
FullStateDictConfig,
|
||||
FullyShardedDataParallel as FSDP, # noqa F401
|
||||
StateDictType,
|
||||
)
|
||||
|
||||
state_cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
optim_cfg = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg):
|
||||
model_state_dict = model.state_dict()
|
||||
optim_state_dict = FSDP.optim_state_dict(model, optimizer)
|
||||
return model_state_dict, optim_state_dict
|
||||
|
||||
|
||||
def load_fsdp_optimizer_state(model, optimizer, checkpoint_dir: Path) -> None:
|
||||
"""Load the FSDP optimizer state (saved as safetensors) and reshard it into the optimizer.
|
||||
|
||||
This is a cross-rank collective and must be called on every rank *after* `accelerator.prepare()`
|
||||
with the prepared (FSDP-wrapped) `model` and `optimizer`. The saved state is the full,
|
||||
world-size-independent optimizer state (keyed by parameter FQNs); `FSDP.optim_state_dict_to_load`
|
||||
reshards it to the current FSDP topology, so resume on a different number of GPUs works.
|
||||
"""
|
||||
from torch.distributed.fsdp import (
|
||||
FullOptimStateDictConfig,
|
||||
FullStateDictConfig,
|
||||
FullyShardedDataParallel as FSDP, # noqa F401
|
||||
StateDictType,
|
||||
)
|
||||
|
||||
# Every rank reads the same full state from the (shared) checkpoint dir, so rank0_only=False.
|
||||
full_osd = load_optimizer_state_dict(checkpoint_dir / TRAINING_STATE_DIR)
|
||||
state_cfg = FullStateDictConfig(rank0_only=False)
|
||||
optim_cfg = FullOptimStateDictConfig(rank0_only=False)
|
||||
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg):
|
||||
sharded_osd = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=full_osd)
|
||||
optimizer.load_state_dict(sharded_osd)
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# Hub: checkpoint push (resume artifact) and publishing (distribution artifact)
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
|
||||
|
||||
def push_checkpoint_to_hub(
|
||||
@@ -311,6 +510,16 @@ def push_checkpoint_to_hub(
|
||||
The model repo is created idempotently, and the commit is tagged with the
|
||||
checkpoint step so a checkpoint can be recovered with
|
||||
--policy.pretrained_revision=<step> instead of a commit sha.
|
||||
|
||||
The directory is uploaded verbatim — including DCP shards under the DCP formats: this tree
|
||||
exists for *resume*, not distribution, and `resolve_resume_checkpoint` downloads it back
|
||||
symmetrically.
|
||||
|
||||
Args:
|
||||
checkpoint_dir (Path): The local checkpoint step directory to upload.
|
||||
repo_id (str): The Hub model repo to push to (created idempotently if missing).
|
||||
private (bool | None): Whether a newly created repo should be private. Defaults to
|
||||
None (public unless the organization's default is private).
|
||||
"""
|
||||
api = HfApi()
|
||||
api.create_repo(repo_id=repo_id, repo_type="model", private=private, exist_ok=True)
|
||||
@@ -338,6 +547,16 @@ def resolve_resume_checkpoint(repo_id: str, output_dir: Path) -> Path:
|
||||
into `output_dir/checkpoints/<step>/`, recreate the local `last` symlink, and return that local
|
||||
checkpoint dir. Used to resume training from the Hub on a machine (or HF Jobs pod) that does not
|
||||
have the original local run dir.
|
||||
|
||||
Args:
|
||||
repo_id (str): The Hub model repo holding `checkpoints/<step>/` subtrees.
|
||||
output_dir (Path): The local run directory to download the checkpoint into.
|
||||
|
||||
Returns:
|
||||
Path: The local checkpoint step directory, `output_dir/checkpoints/<step>`.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the repo contains no checkpoints under `checkpoints/`.
|
||||
"""
|
||||
latest = find_latest_hub_checkpoint(repo_id)
|
||||
if latest is None:
|
||||
@@ -354,3 +573,214 @@ def resolve_resume_checkpoint(repo_id: str, output_dir: Path) -> Path:
|
||||
checkpoint_dir = output_dir / latest
|
||||
update_last_checkpoint(checkpoint_dir)
|
||||
return checkpoint_dir
|
||||
|
||||
|
||||
def publish_trained_model(
|
||||
cfg: TrainPipelineConfig,
|
||||
model: "PreTrainedPolicy | PreTrainedRewardModel",
|
||||
preprocessor: PolicyProcessorPipeline | None,
|
||||
postprocessor: PolicyProcessorPipeline | None,
|
||||
dataset_meta: "LeRobotDatasetMetadata | None",
|
||||
*,
|
||||
peft_model: Any | None = None,
|
||||
) -> None:
|
||||
"""Publish the complete training bundle as a distributable model repo.
|
||||
|
||||
Collective-safe: call on ALL ranks — the model commit gathers sharded weights through
|
||||
`save_pretrained`; uploads happen on the main process only (gated inside
|
||||
`HubMixin.push_to_hub` and here). Commits, in order: (1) the model (skipped for PEFT —
|
||||
adapters replace full weights), (2) the preprocessor, (3) the postprocessor, (4) the bundle
|
||||
sidecar: README.md model card + train_config.json (+ adapter weights and the wrapped
|
||||
policy's config in the PEFT case). Every commit uploads a freshly assembled directory, so
|
||||
a published repo carries only the distributable artifacts.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and used
|
||||
to render the model card.
|
||||
model (PreTrainedPolicy | PreTrainedRewardModel): The trained model to publish; its
|
||||
config supplies the target repo id, visibility, license, and tags.
|
||||
preprocessor (PolicyProcessorPipeline | None): The preprocessor pipeline to publish
|
||||
alongside the model, if any.
|
||||
postprocessor (PolicyProcessorPipeline | None): The postprocessor pipeline to publish
|
||||
alongside the model, if any.
|
||||
dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata for the model card, if
|
||||
available.
|
||||
peft_model (Any | None): The PEFT wrapper when training adapters; its adapter weights
|
||||
replace the full model weights in the published repo. Defaults to None.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model config carries no repo id (`--policy.repo_id`).
|
||||
"""
|
||||
model_cfg = model.config
|
||||
repo_id = model_cfg.repo_id
|
||||
if not repo_id:
|
||||
raise ValueError("Publishing requires a repo id (--policy.repo_id).")
|
||||
ignore = ["*.tmp", "*.log"]
|
||||
|
||||
if peft_model is None:
|
||||
# Calls are made on the exact objects that own each method (never through PEFT's
|
||||
# attribute forwarding), so the peft branch below never touches this path.
|
||||
model.push_to_hub(repo_id, private=model_cfg.private, ignore_patterns=ignore)
|
||||
if preprocessor is not None:
|
||||
preprocessor.push_to_hub(repo_id, private=model_cfg.private)
|
||||
if postprocessor is not None:
|
||||
postprocessor.push_to_hub(repo_id, private=model_cfg.private)
|
||||
|
||||
if is_main_process():
|
||||
api = HfApi()
|
||||
repo_id = api.create_repo(repo_id=repo_id, private=model_cfg.private, exist_ok=True).repo_id
|
||||
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
|
||||
saved_path = Path(tmp) / repo_id
|
||||
saved_path.mkdir(parents=True, exist_ok=True)
|
||||
if peft_model is not None:
|
||||
peft_model.save_pretrained(saved_path) # adapter weights + adapter config
|
||||
model.config.save_pretrained(saved_path) # PEFT cannot write the policy config
|
||||
card = generate_model_card(model_cfg, cfg=cfg, dataset_meta=dataset_meta)
|
||||
card.save(str(saved_path / "README.md"))
|
||||
cfg.save_pretrained(saved_path) # train_config.json
|
||||
commit_info = api.upload_folder(
|
||||
repo_id=repo_id,
|
||||
repo_type="model",
|
||||
folder_path=saved_path,
|
||||
commit_message="Upload model card and train config",
|
||||
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
|
||||
ignore_patterns=ignore,
|
||||
)
|
||||
# Contract: lerobot.jobs.hf.submit_to_hf watches for this exact "Model pushed to <url>"
|
||||
# line to end a remote run early. Keep the wording and URL format in sync.
|
||||
logging.info(f"Model pushed to {commit_info.repo_url.url}")
|
||||
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# Model card
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
|
||||
_BASE_MODEL_MAPPING = {
|
||||
"smolvla": "lerobot/smolvla_base",
|
||||
"pi0": "lerobot/pi0_base",
|
||||
"pi05": "lerobot/pi05_base",
|
||||
"pi0_fast": "lerobot/pi0fast-base",
|
||||
"xvla": "lerobot/xvla-base",
|
||||
}
|
||||
|
||||
|
||||
def build_card_context(
|
||||
cfg: TrainPipelineConfig | None,
|
||||
dataset_meta: "LeRobotDatasetMetadata | None",
|
||||
input_features: dict | None,
|
||||
output_features: dict | None,
|
||||
) -> dict:
|
||||
"""Collect optional data for the model-card template.
|
||||
|
||||
Returns plain values only (no Markdown) — the template in
|
||||
``lerobot/templates/lerobot_modelcard_template.md`` decides how and whether to show
|
||||
each one. Everything is best-effort: anything unavailable is left empty/None and the
|
||||
template simply skips that section, so this never breaks a Hub push.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig | None): The training config supplying the training section,
|
||||
if available.
|
||||
dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata supplying the dataset,
|
||||
robot-type, and camera sections, if available.
|
||||
input_features (dict | None): The policy's input feature declarations, if any.
|
||||
output_features (dict | None): The policy's output feature declarations, if any.
|
||||
|
||||
Returns:
|
||||
dict: Template context with `training`, `input_features`, `output_features`,
|
||||
`dataset`, `robot_type`, and `cameras` entries; unavailable pieces stay
|
||||
empty/None.
|
||||
"""
|
||||
context = {
|
||||
"training": None,
|
||||
"input_features": input_features or {},
|
||||
"output_features": output_features or {},
|
||||
"dataset": None,
|
||||
"robot_type": None,
|
||||
"cameras": [],
|
||||
}
|
||||
|
||||
if cfg is not None:
|
||||
optimizer = getattr(cfg, "optimizer", None)
|
||||
context["training"] = {
|
||||
"steps": cfg.steps,
|
||||
"batch_size": cfg.batch_size,
|
||||
"seed": cfg.seed,
|
||||
"optimizer": getattr(optimizer, "type", None) if optimizer else None,
|
||||
"lr": getattr(optimizer, "lr", None) if optimizer else None,
|
||||
"lerobot_version": __version__,
|
||||
}
|
||||
|
||||
if dataset_meta is not None:
|
||||
context["dataset"] = {
|
||||
"repo_id": dataset_meta.repo_id,
|
||||
"episodes": dataset_meta.total_episodes,
|
||||
"frames": dataset_meta.total_frames,
|
||||
"fps": dataset_meta.fps,
|
||||
"tasks": [str(task) for task in dataset_meta.tasks.index],
|
||||
}
|
||||
context["robot_type"] = dataset_meta.robot_type
|
||||
context["cameras"] = [key.split(".")[-1] for key in dataset_meta.camera_keys]
|
||||
|
||||
return context
|
||||
|
||||
|
||||
def generate_model_card(
|
||||
model_cfg: PreTrainedConfig | RewardModelConfig,
|
||||
cfg: TrainPipelineConfig | None = None,
|
||||
dataset_meta: "LeRobotDatasetMetadata | None" = None,
|
||||
) -> ModelCard:
|
||||
"""Render the LeRobot model card for a trained policy or reward model.
|
||||
|
||||
A free function on purpose: every template variable comes from arguments — the model
|
||||
config, the training config, and the dataset metadata — none from a live model, so a card
|
||||
can also be rendered from a checkpoint's `config.json` alone (see `lerobot-convert-dcp`).
|
||||
The config type selects the template: reward models get the reward-model card, policies the
|
||||
policy card with the training/dataset sections.
|
||||
|
||||
Args:
|
||||
model_cfg (PreTrainedConfig | RewardModelConfig): The model config providing type,
|
||||
license, tags, repo id, and — for policies — the feature declarations.
|
||||
cfg (TrainPipelineConfig | None, optional): The training config for the training and
|
||||
dataset card sections. Defaults to None.
|
||||
dataset_meta (LeRobotDatasetMetadata | None, optional): Dataset metadata for the
|
||||
dataset card sections. Defaults to None.
|
||||
|
||||
Returns:
|
||||
ModelCard: The rendered and validated LeRobot model card.
|
||||
"""
|
||||
model_type = model_cfg.type
|
||||
base_model = _BASE_MODEL_MAPPING.get(model_type)
|
||||
|
||||
if isinstance(model_cfg, RewardModelConfig):
|
||||
tags = {"robotics", "lerobot", "reward-model", model_type}
|
||||
template_card = (
|
||||
files("lerobot.templates")
|
||||
.joinpath("lerobot_rewardmodel_modelcard_template.md")
|
||||
.read_text("utf-8")
|
||||
)
|
||||
context: dict[str, Any] = {} # the reward template renders from card_data alone
|
||||
else:
|
||||
tags = {"robotics", "lerobot", model_type}
|
||||
template_card = (
|
||||
files("lerobot.templates").joinpath("lerobot_modelcard_template.md").read_text("utf-8")
|
||||
)
|
||||
context = build_card_context(cfg, dataset_meta, model_cfg.input_features, model_cfg.output_features)
|
||||
# Used by the template to pre-fill commands and the "Fine-tuned from" line.
|
||||
context["policy_repo_id"] = model_cfg.repo_id
|
||||
context["base_model"] = base_model
|
||||
|
||||
card_data = ModelCardData(
|
||||
license=model_cfg.license or "apache-2.0",
|
||||
library_name="lerobot",
|
||||
pipeline_tag="robotics",
|
||||
tags=list(tags.union(model_cfg.tags or [])),
|
||||
model_name=model_type,
|
||||
datasets=cfg.dataset.repo_id if cfg is not None else None,
|
||||
base_model=base_model,
|
||||
)
|
||||
card = ModelCard.from_template(card_data, template_str=template_card, **context)
|
||||
card.validate()
|
||||
return card
|
||||
|
||||
@@ -0,0 +1,273 @@
|
||||
#!/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.
|
||||
"""Execution-runtime configuration: everything handed to (or applied by) the `Accelerator`.
|
||||
|
||||
Each sub-config mirrors the plain-typed subset of the corresponding accelerate object and
|
||||
builds it at runtime (the way ``OptimizerConfig.build()`` constructs a ``torch.optim.Optimizer``),
|
||||
so the whole tree round-trips through the CLI and ``train_config.json`` and parsing a config
|
||||
never imports accelerate.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from lerobot.configs.parallelism import ParallelismConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import (
|
||||
DistributedDataParallelKwargs,
|
||||
FullyShardedDataParallelPlugin,
|
||||
GradientAccumulationPlugin,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FSDPConfig:
|
||||
"""Mirror of the `FullyShardedDataParallelPlugin` subset LeRobot supports (FSDP2 only).
|
||||
|
||||
Exactly one wrap policy applies: `wrap_modules` (module *class names* forming the FSDP
|
||||
units — and, later, the activation-checkpointing units) or `min_num_params` (size-based).
|
||||
When both are None, the policy's own `_fsdp_wrap_modules` declaration is used; a run where
|
||||
no wrap source exists at all fails loudly rather than silently wrapping only the root.
|
||||
"""
|
||||
|
||||
reshard_after_forward: bool = True
|
||||
wrap_modules: list[str] | None = None
|
||||
min_num_params: int | None = None
|
||||
cpu_offload: bool = False
|
||||
# Regex matched against module FQNs to exclude their parameters from sharding.
|
||||
ignored_modules: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the wrap-policy fields.
|
||||
|
||||
Raises:
|
||||
ValueError: If both ``wrap_modules`` and ``min_num_params`` are set (they are
|
||||
mutually exclusive wrap policies), or if ``min_num_params`` is < 1.
|
||||
"""
|
||||
if self.wrap_modules is not None and self.min_num_params is not None:
|
||||
raise ValueError(
|
||||
"fsdp.wrap_modules and fsdp.min_num_params are mutually exclusive wrap policies."
|
||||
)
|
||||
if self.min_num_params is not None and self.min_num_params < 1:
|
||||
raise ValueError(f"fsdp.min_num_params must be >= 1, got {self.min_num_params}.")
|
||||
|
||||
def build_plugin(self) -> "FullyShardedDataParallelPlugin":
|
||||
"""Build the FSDP2 plugin for `Accelerator(fsdp_plugin=...)`.
|
||||
|
||||
Returns:
|
||||
FullyShardedDataParallelPlugin: FSDP2 (`fsdp_version=2`) plugin carrying the
|
||||
mirrored wrap policy, resharding, CPU-offload, and ignored-modules settings.
|
||||
"""
|
||||
from accelerate.utils import FullyShardedDataParallelPlugin
|
||||
|
||||
use_size_policy = self.min_num_params is not None
|
||||
return FullyShardedDataParallelPlugin(
|
||||
fsdp_version=2,
|
||||
reshard_after_forward=self.reshard_after_forward,
|
||||
auto_wrap_policy="size_based_wrap" if use_size_policy else "transformer_based_wrap",
|
||||
# May legitimately still be None here: the policy-declared default is applied right
|
||||
# before `accelerator.prepare()` (see lerobot.distributed.factory.set_fsdp_wrap_modules).
|
||||
transformer_cls_names_to_wrap=list(self.wrap_modules) if self.wrap_modules else None,
|
||||
min_num_params=self.min_num_params,
|
||||
cpu_offload=self.cpu_offload,
|
||||
ignored_modules=self.ignored_modules,
|
||||
# state_dict_type stays at the FSDP2 default (SHARDED_STATE_DICT) and is never
|
||||
# switched: full gathers go through torch's state-dict API, which does not consult
|
||||
# the plugin. activation_checkpointing stays False: AC is LeRobot-owned.
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DDPConfig:
|
||||
"""Mirror of the `DistributedDataParallelKwargs` subset LeRobot exposes."""
|
||||
|
||||
# Today's in-script default, kept for models with conditional computation.
|
||||
find_unused_parameters: bool = True
|
||||
gradient_as_bucket_view: bool = False
|
||||
static_graph: bool = False
|
||||
|
||||
def build_kwargs_handler(self) -> "DistributedDataParallelKwargs":
|
||||
"""Build the DDP kwargs handler for `Accelerator(kwargs_handlers=[...])`.
|
||||
|
||||
Returns:
|
||||
DistributedDataParallelKwargs: Handler carrying the mirrored DDP fields, applied
|
||||
by accelerate when it wraps the model in `DistributedDataParallel`.
|
||||
"""
|
||||
from accelerate.utils import DistributedDataParallelKwargs
|
||||
|
||||
return DistributedDataParallelKwargs(
|
||||
find_unused_parameters=self.find_unused_parameters,
|
||||
gradient_as_bucket_view=self.gradient_as_bucket_view,
|
||||
static_graph=self.static_graph,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GradientAccumulationConfig:
|
||||
"""Mirror of the `GradientAccumulationPlugin` subset LeRobot supports.
|
||||
|
||||
Only the step count is a knob. ``sync_with_dataloader`` is pinned to False by
|
||||
:meth:`build_plugin`: the training loop cycles a finite dataloader, so accelerate's default
|
||||
of syncing at every dataloader end would force an optimizer step at every dataset epoch
|
||||
boundary instead of every ``steps`` micro-batches.
|
||||
"""
|
||||
|
||||
steps: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the accumulation step count.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``steps`` is < 1.
|
||||
"""
|
||||
if self.steps < 1:
|
||||
raise ValueError(f"gradient_accumulation.steps must be >= 1, got {self.steps}.")
|
||||
|
||||
def build_plugin(self) -> "GradientAccumulationPlugin":
|
||||
"""Build the plugin for `Accelerator(gradient_accumulation_plugin=...)`.
|
||||
|
||||
A named plugin argument, not a `kwargs_handlers` entry: accelerate consumes this object
|
||||
through its dedicated constructor parameter — the `KwargsHandler` base class only lends
|
||||
it `to_kwargs()`, so the consumption site, not the inheritance, decides its role.
|
||||
|
||||
Returns:
|
||||
GradientAccumulationPlugin: Carrying the mirrored step count, with
|
||||
``sync_with_dataloader=False`` pinned (see the class docstring).
|
||||
"""
|
||||
from accelerate.utils import GradientAccumulationPlugin
|
||||
|
||||
return GradientAccumulationPlugin(num_steps=self.steps, sync_with_dataloader=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompileConfig:
|
||||
"""torch.compile knobs — a configured placeholder: wiring lands in a later round.
|
||||
|
||||
The setup-order contract it will follow is already fixed: compile applies
|
||||
after CP dispatch install and activation checkpointing, before `fully_shard`, regionally
|
||||
(per wrap unit) — the only combination proven with FSDP2.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
backend: str = "inductor"
|
||||
mode: str | None = None
|
||||
regional: bool = True
|
||||
|
||||
|
||||
class ActivationCheckpointingMode(str, Enum):
|
||||
NONE = "none"
|
||||
FULL = "full"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ActivationCheckpointingConfig:
|
||||
"""Activation-checkpointing knobs — a configured placeholder: wiring lands in a later round.
|
||||
|
||||
AC units will coincide with the FSDP wrap units (one declaration drives both), applied
|
||||
before torch.compile and `fully_shard` (the same ordering contract as CompileConfig).
|
||||
"""
|
||||
|
||||
mode: ActivationCheckpointingMode = ActivationCheckpointingMode.NONE
|
||||
|
||||
|
||||
@dataclass
|
||||
class AcceleratorConfig:
|
||||
"""Builds the `Accelerator` — the runtime counterpart of the `parallelism` topology.
|
||||
|
||||
`mixed_precision` selects accelerate-native AMP for DDP/single-GPU runs and the FSDP2
|
||||
`MixedPrecisionPolicy` for sharded runs (accelerate derives it). Sharded runs support
|
||||
"no" and "bf16" only; fp16's GradScaler-over-DTensor path is unverified and fails fast
|
||||
at config validation.
|
||||
"""
|
||||
|
||||
mixed_precision: str = "no"
|
||||
gradient_accumulation: GradientAccumulationConfig = field(default_factory=GradientAccumulationConfig)
|
||||
fsdp: FSDPConfig = field(default_factory=FSDPConfig)
|
||||
ddp: DDPConfig = field(default_factory=DDPConfig)
|
||||
compile: CompileConfig = field(default_factory=CompileConfig)
|
||||
activation_checkpointing: ActivationCheckpointingConfig = field(
|
||||
default_factory=ActivationCheckpointingConfig
|
||||
)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the accelerate-facing scalar fields.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``mixed_precision`` is not one of ``"no"``, ``"fp16"``, ``"bf16"``.
|
||||
"""
|
||||
if self.mixed_precision not in ("no", "fp16", "bf16"):
|
||||
raise ValueError(
|
||||
f"mixed_precision must be one of 'no', 'fp16', 'bf16', got {self.mixed_precision!r}."
|
||||
)
|
||||
|
||||
def build(self, parallelism: ParallelismConfig, *, cpu: bool = False) -> "Accelerator":
|
||||
"""Translate the mirrored fields into a ready `Accelerator` (call once per process).
|
||||
|
||||
`parallelism` must already be resolved against the world size. The degradation matrix
|
||||
is encoded here and nowhere else: sharded -> FSDP2 (+HSDP via the accelerate
|
||||
`ParallelismConfig` mesh), replicated-only -> DDP kwargs, single process -> plain.
|
||||
|
||||
Args:
|
||||
parallelism (ParallelismConfig): The resolved process topology; selects which
|
||||
accelerate path (FSDP2 mesh, DDP kwargs handler, or plain) is configured.
|
||||
cpu (bool): Force CPU execution even when CUDA is available. Defaults to False.
|
||||
|
||||
Returns:
|
||||
Accelerator: The configured accelerate entry point for this process.
|
||||
"""
|
||||
from accelerate import Accelerator
|
||||
|
||||
kwargs: dict = {
|
||||
# LeRobot steps its scheduler manually once per training step; accelerate must not
|
||||
# rescale scheduler stepping by num_processes.
|
||||
"step_scheduler_with_optimizer": False,
|
||||
"gradient_accumulation_plugin": self.gradient_accumulation.build_plugin(),
|
||||
"mixed_precision": self.mixed_precision,
|
||||
"cpu": cpu,
|
||||
}
|
||||
if parallelism.is_sharded:
|
||||
kwargs["fsdp_plugin"] = self.fsdp.build_plugin()
|
||||
kwargs["parallelism_config"] = _accelerate_parallelism_config(parallelism)
|
||||
elif parallelism.is_replicated_only:
|
||||
kwargs["kwargs_handlers"] = [self.ddp.build_kwargs_handler()]
|
||||
return Accelerator(**kwargs)
|
||||
|
||||
|
||||
def _accelerate_parallelism_config(parallelism: ParallelismConfig) -> object:
|
||||
"""LeRobot topology -> accelerate `ParallelismConfig`.
|
||||
|
||||
CP is declared honestly (`cp_size = ring x ulysses`) so accelerate builds the canonical
|
||||
mesh, folds CP into the FSDP shard group (`dp_shard_cp`), and duplicates batches within CP
|
||||
groups. The ring/ulysses sub-structure stays private to `lerobot.distributed.ParallelDims`.
|
||||
|
||||
Args:
|
||||
parallelism (ParallelismConfig): The resolved LeRobot topology to translate.
|
||||
|
||||
Returns:
|
||||
object: The accelerate `ParallelismConfig` mirroring `dp_replicate`, `dp_shard`, and
|
||||
the collapsed `cp_size` (annotated as `object` so importing this module never
|
||||
imports accelerate).
|
||||
"""
|
||||
from accelerate.parallelism_config import ParallelismConfig as AccelerateParallelismConfig
|
||||
|
||||
return AccelerateParallelismConfig(
|
||||
dp_replicate_size=parallelism.dp_replicate,
|
||||
dp_shard_size=parallelism.dp_shard,
|
||||
cp_size=parallelism.cp_size,
|
||||
)
|
||||
@@ -14,6 +14,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from lerobot.transforms import ImageTransformsConfig
|
||||
@@ -21,6 +22,8 @@ from lerobot.utils.import_utils import get_safe_default_video_backend
|
||||
|
||||
from .video import DEFAULT_DEPTH_UNIT, DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DatasetConfig:
|
||||
@@ -29,10 +32,15 @@ class DatasetConfig:
|
||||
# "dataset_index" into the returned item. The index mapping is made according to the order in which the
|
||||
# datasets are provided.
|
||||
repo_id: str
|
||||
# Hub repository type: "dataset" (default) or "bucket" for an HF Storage Bucket streamed over
|
||||
# hf://buckets/. Buckets are streaming-only, so "bucket" requires streaming=true.
|
||||
repo_type: str = "dataset"
|
||||
# Root directory for a concrete local dataset tree (e.g. 'dataset/path'). If None, local datasets are
|
||||
# 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
|
||||
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)
|
||||
revision: str | None = None
|
||||
use_imagenet_stats: bool = True
|
||||
@@ -48,6 +56,16 @@ class DatasetConfig:
|
||||
eval_split: float = 0.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.repo_type not in ("dataset", "bucket"):
|
||||
raise ValueError(f"repo_type must be 'dataset' or 'bucket', got {self.repo_type!r}")
|
||||
if self.repo_type == "bucket" and not self.streaming:
|
||||
raise ValueError(
|
||||
"repo_type='bucket' is streaming-only: set streaming=true to train from an HF Storage Bucket."
|
||||
)
|
||||
if self.repo_type == "bucket" and self.eval_split != 0.0:
|
||||
raise ValueError(
|
||||
"eval_split requires map-style datasets and is not supported with repo_type='bucket'."
|
||||
)
|
||||
if self.depth_output_unit not in (DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT):
|
||||
raise ValueError(
|
||||
f"depth_output_unit must be '{DEPTH_METER_UNIT}' or '{DEPTH_MILLIMETER_UNIT}', got {self.depth_output_unit!r}"
|
||||
@@ -62,6 +80,14 @@ class DatasetConfig:
|
||||
if len(self.episodes) != len(set(self.episodes)):
|
||||
duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1})
|
||||
raise ValueError(f"Episode indices contain duplicates: {duplicates}")
|
||||
if self.exclude_episodes is not None:
|
||||
negative_episodes = [episode for episode in self.exclude_episodes if episode < 0]
|
||||
if negative_episodes:
|
||||
logger.warning(
|
||||
"Ignoring negative exclude_episodes entries: %s",
|
||||
negative_episodes,
|
||||
)
|
||||
self.exclude_episodes = [episode for episode in self.exclude_episodes if episode >= 0]
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -93,9 +119,6 @@ class EvalConfig:
|
||||
recording_repo_id: str | None = None
|
||||
# Whether the pushed recording repositories should be private.
|
||||
recording_private: bool = False
|
||||
# Whether to save the policy's imagined/predicted video (world-model policies only) as mp4s.
|
||||
# Requests intermediate predictions from the policy each step; policies that produce none are unaffected.
|
||||
save_predicted_video: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.recording_repo_id is not None and not self.recording:
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
#!/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.
|
||||
"""Declarative process topology for distributed training and inference.
|
||||
|
||||
The mesh convention (canonical row-major rank layout, outermost first)::
|
||||
|
||||
(dp_replicate, dp_shard, ring, ulysses)
|
||||
|
||||
- ``dp_replicate x dp_shard`` is the data-parallel world: HSDP replicates over
|
||||
``dp_replicate`` and shards parameters over ``dp_shard``. FSDP2's actual shard
|
||||
group folds context parallelism in (``dp_shard x ring x ulysses``), matching
|
||||
accelerate's ``dp_shard_cp`` flattening and torchtitan's ``fsdp`` axis.
|
||||
- ``ring`` is the outer and ``ulysses`` the inner context-parallel dim
|
||||
(diffusers convention: ulysses all-to-all exchanges run over adjacent, typically
|
||||
NVLink-connected ranks).
|
||||
- ``cfg_parallel`` (classifier-free-guidance parallelism) is a branch-parallel,
|
||||
inference-only dim that sits between dp and the sequence dims. It never
|
||||
affects weight sharding or checkpoints.
|
||||
|
||||
This module is pure configuration: plain-typed dataclasses that draccus can
|
||||
round-trip through the CLI and ``train_config.json``. Runtime objects (device
|
||||
meshes, process groups) live in :mod:`lerobot.distributed`.
|
||||
"""
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContextParallelConfig:
|
||||
"""Ring x Ulysses context parallelism (sequence parallelism for attention).
|
||||
|
||||
Both degrees are configured placeholders in this release: the CP engine is not implemented
|
||||
yet, and enabling either degree > 1 fails fast at config validation. The fields exist now so
|
||||
that the CLI surface, checkpoint metadata, and mesh math are stable when the engine lands.
|
||||
"""
|
||||
|
||||
ring_degree: int = 1
|
||||
ulysses_degree: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the declared context-parallel degrees.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``ring_degree`` or ``ulysses_degree`` is < 1.
|
||||
"""
|
||||
if self.ring_degree < 1 or self.ulysses_degree < 1:
|
||||
raise ValueError(
|
||||
f"Context-parallel degrees must be >= 1, got ring_degree={self.ring_degree}, "
|
||||
f"ulysses_degree={self.ulysses_degree}."
|
||||
)
|
||||
|
||||
@property
|
||||
def size(self) -> int:
|
||||
"""Total number of ranks a full sequence is sharded across."""
|
||||
return self.ring_degree * self.ulysses_degree
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelismConfig:
|
||||
"""Degrees of every parallelism dim. Invariant: their product equals the world size.
|
||||
|
||||
Degradations are expressed purely through the degrees (no mode flags):
|
||||
|
||||
- single process: all degrees 1;
|
||||
- DDP: ``dp_replicate == world_size`` (auto-filled when every sharding field is left at its
|
||||
default — plain ``torchrun`` keeps today's out-of-the-box behavior);
|
||||
- FSDP: ``dp_shard > 1`` (or ``-1`` to fill the remaining world into the shard dim);
|
||||
- HSDP: ``dp_replicate > 1`` and ``dp_shard > 1``.
|
||||
|
||||
``resolve()`` turns the declared degrees into concrete ones once the world size is known and
|
||||
is the single place the world-size equation is enforced. It is called by
|
||||
:func:`lerobot.distributed.factory.make_accelerator`; the config is inert until then.
|
||||
"""
|
||||
|
||||
dp_replicate: int = 1
|
||||
# -1 is an explicit opt-in sentinel: shard over world_size // (dp_replicate * cp).
|
||||
dp_shard: int = 1
|
||||
context_parallel: ContextParallelConfig = field(default_factory=ContextParallelConfig)
|
||||
# Classifier-free-guidance parallelism — inference-only (cosmos/vllm-omni precedent:
|
||||
# cond/uncond branches on different ranks). Reserved for the serving round; training
|
||||
# validates it to 1. Meaningful values are 1 or 2 (Cosmos3 has two CFG branches).
|
||||
cfg_parallel: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the declared degrees (world-size-independent checks only).
|
||||
|
||||
Raises:
|
||||
ValueError: If ``dp_replicate`` is < 1, ``dp_shard`` is neither >= 1 nor the
|
||||
``-1`` infer sentinel, or ``cfg_parallel`` is not 1 or 2.
|
||||
"""
|
||||
if self.dp_replicate < 1:
|
||||
raise ValueError(f"dp_replicate must be >= 1, got {self.dp_replicate}.")
|
||||
if self.dp_shard < 1 and self.dp_shard != -1:
|
||||
raise ValueError(f"dp_shard must be >= 1, or -1 to infer, got {self.dp_shard}.")
|
||||
if self.cfg_parallel not in (1, 2):
|
||||
raise ValueError(f"cfg_parallel must be 1 or 2, got {self.cfg_parallel}.")
|
||||
|
||||
@property
|
||||
def cp_size(self) -> int:
|
||||
"""Total context-parallel size (``ring_degree * ulysses_degree``)."""
|
||||
return self.context_parallel.size
|
||||
|
||||
@property
|
||||
def is_sharded(self) -> bool:
|
||||
"""True when the run uses FSDP2 (parameters sharded); selects the sharded engine path."""
|
||||
return self.dp_shard != 1 or self.cp_size > 1
|
||||
|
||||
@property
|
||||
def is_replicated_only(self) -> bool:
|
||||
"""True for plain DDP (weights replicated, no sharding)."""
|
||||
return not self.is_sharded and self.dp_replicate > 1
|
||||
|
||||
@property
|
||||
def dp_world_size(self) -> int:
|
||||
"""Number of distinct data-parallel workers (batches are sharded this many ways).
|
||||
|
||||
Returns:
|
||||
int: ``dp_replicate * dp_shard``.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If accessed while ``dp_shard`` is still the ``-1`` sentinel, i.e.
|
||||
before :meth:`resolve` has bound the degrees to a world size.
|
||||
"""
|
||||
if self.dp_shard == -1:
|
||||
raise RuntimeError("dp_world_size is undefined before resolve() fills dp_shard=-1.")
|
||||
return self.dp_replicate * self.dp_shard
|
||||
|
||||
def resolve(self, world_size: int) -> None:
|
||||
"""Bind the declared degrees to a concrete world size (idempotent).
|
||||
|
||||
Fills the ``dp_shard=-1`` sentinel, auto-fills ``dp_replicate`` for the DDP degradation,
|
||||
and enforces ``dp_replicate * dp_shard * cp == world_size`` with every degree echoed on
|
||||
failure.
|
||||
|
||||
Args:
|
||||
world_size (int): Total number of launched processes (torchrun's ``WORLD_SIZE``).
|
||||
|
||||
Raises:
|
||||
ValueError: If a context-parallel degree is > 1 (the CP engine is not implemented
|
||||
yet), if ``dp_shard=-1`` cannot be inferred because ``world_size`` is not
|
||||
divisible by ``dp_replicate * cp``, or if the resolved degrees do not multiply
|
||||
to ``world_size``.
|
||||
"""
|
||||
if self.cp_size > 1:
|
||||
raise ValueError(
|
||||
"Context parallelism is not implemented yet: ring_degree and ulysses_degree "
|
||||
"must be 1. The fields are reserved for the CP engine round."
|
||||
)
|
||||
if self.is_sharded:
|
||||
if self.dp_shard == -1:
|
||||
self.dp_shard, remainder = divmod(world_size, self.dp_replicate * self.cp_size)
|
||||
if remainder or self.dp_shard < 1:
|
||||
raise ValueError(
|
||||
f"Cannot infer dp_shard: world_size={world_size} is not divisible by "
|
||||
f"dp_replicate={self.dp_replicate} * cp={self.cp_size}."
|
||||
)
|
||||
elif self.dp_replicate == 1:
|
||||
# Untouched config on a multi-process launch: fill the DDP degradation.
|
||||
self.dp_replicate = world_size
|
||||
total = self.dp_replicate * self.dp_shard * self.cp_size
|
||||
if total != world_size:
|
||||
raise ValueError(
|
||||
f"Parallelism degrees do not multiply to the world size: dp_replicate="
|
||||
f"{self.dp_replicate} * dp_shard={self.dp_shard} * ring="
|
||||
f"{self.context_parallel.ring_degree} * ulysses="
|
||||
f"{self.context_parallel.ulysses_degree} = {total} != WORLD_SIZE={world_size}."
|
||||
)
|
||||
|
||||
|
||||
def world_size_from_env() -> int:
|
||||
"""World size as set by torchrun (or 1 outside distributed launches).
|
||||
|
||||
Returns:
|
||||
int: The ``WORLD_SIZE`` environment variable, or 1 when unset.
|
||||
"""
|
||||
return int(os.environ.get("WORLD_SIZE", "1"))
|
||||
@@ -23,6 +23,7 @@ from typing import Any, Literal, get_args
|
||||
|
||||
MessageRole = Literal["user", "assistant", "system", "tool"]
|
||||
MessageStream = Literal["high_level", "low_level"]
|
||||
RecipeRoute = Literal["vqa"]
|
||||
|
||||
DEFAULT_BINDINGS = {
|
||||
"subtask": "active_at(t, style=subtask)",
|
||||
@@ -40,6 +41,7 @@ discovery (here) and rendered-message substitution (in ``language_render``)."""
|
||||
|
||||
_VALID_ROLES = frozenset(get_args(MessageRole))
|
||||
_VALID_STREAMS = frozenset(get_args(MessageStream))
|
||||
_VALID_ROUTES = frozenset(get_args(RecipeRoute))
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -78,7 +80,7 @@ class MessageTurn:
|
||||
raise ValueError(f"Unsupported message stream: {self.stream!r}")
|
||||
if self.content is None and self.tool_calls_from is None:
|
||||
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.")
|
||||
if isinstance(self.content, list):
|
||||
for block in self.content:
|
||||
@@ -99,13 +101,16 @@ class TrainingRecipe:
|
||||
|
||||
A recipe is either a *message recipe* (``messages`` plus optional
|
||||
``bindings``) or a *blend recipe* (``blend`` mapping names to weighted
|
||||
sub-recipes). ``weight`` is only meaningful inside a blend.
|
||||
sub-recipes). ``weight`` and ``route`` are only meaningful inside a blend;
|
||||
``route: vqa`` gives sparse VQA annotations priority over normal weighted
|
||||
selection.
|
||||
"""
|
||||
|
||||
messages: list[MessageTurn] | None = None
|
||||
bindings: dict[str, str] | None = None
|
||||
blend: dict[str, TrainingRecipe] | None = None
|
||||
weight: float | None = None
|
||||
route: RecipeRoute | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate that exactly one of ``messages`` or ``blend`` is set."""
|
||||
@@ -113,6 +118,10 @@ class TrainingRecipe:
|
||||
raise ValueError("TrainingRecipe must set only one of messages or blend.")
|
||||
if self.messages is None and self.blend is None:
|
||||
raise ValueError("TrainingRecipe must set one of messages or blend.")
|
||||
if self.route is not None and self.route not in _VALID_ROUTES:
|
||||
raise ValueError(f"Unsupported recipe route: {self.route!r}")
|
||||
if self.blend is not None and self.route is not None:
|
||||
raise ValueError("TrainingRecipe.route may only be set on a message recipe inside a blend.")
|
||||
|
||||
if self.messages is not None:
|
||||
self._validate_message_recipe()
|
||||
@@ -147,8 +156,9 @@ class TrainingRecipe:
|
||||
return cls.from_dict(data)
|
||||
|
||||
def _validate_message_recipe(self) -> None:
|
||||
"""Ensure every templated binding is known and at least one turn is a target."""
|
||||
assert self.messages is not None
|
||||
"""Validate bindings and require text or low-level action supervision."""
|
||||
if self.messages is None:
|
||||
raise ValueError("Cannot validate a message recipe without messages.")
|
||||
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
|
||||
|
||||
for turn in self.messages:
|
||||
@@ -156,12 +166,19 @@ class TrainingRecipe:
|
||||
if missing:
|
||||
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
|
||||
|
||||
if not any(turn.target for turn in self.messages):
|
||||
raise ValueError("Message recipes must contain at least one target turn.")
|
||||
has_target = any(turn.target for turn in self.messages)
|
||||
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:
|
||||
"""Ensure each blend component is a non-empty, weighted message recipe."""
|
||||
assert self.blend is not None
|
||||
if self.blend is None:
|
||||
raise ValueError("Cannot validate a blend recipe without blend components.")
|
||||
if not self.blend:
|
||||
raise ValueError("Blend recipes must contain at least one component.")
|
||||
|
||||
|
||||
@@ -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,72 @@
|
||||
# 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
|
||||
route: vqa
|
||||
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
|
||||
route: vqa
|
||||
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,6 +18,7 @@ import multiprocessing
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
@@ -26,6 +27,8 @@ from huggingface_hub import hf_hub_download
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
|
||||
from lerobot import envs
|
||||
from lerobot.configs.accelerator import AcceleratorConfig, ActivationCheckpointingMode
|
||||
from lerobot.configs.parallelism import ParallelismConfig
|
||||
from lerobot.optim import LRSchedulerConfig, OptimizerConfig
|
||||
from lerobot.utils.constants import PRETRAINED_MODEL_DIR
|
||||
from lerobot.utils.hub import HubMixin, find_latest_hub_checkpoint
|
||||
@@ -39,6 +42,34 @@ from .rewards import RewardModelConfig
|
||||
TRAIN_CONFIG_NAME = "train_config.json"
|
||||
|
||||
|
||||
class CheckpointFormat(str, Enum):
|
||||
"""Model-artifact format inside training checkpoints.
|
||||
|
||||
Selects only the *model* artifact; the training_state layout is format-independent (the
|
||||
optimizer channel is always DCP under sharded runs, safetensors+json otherwise).
|
||||
|
||||
- SAFETENSORS (default): a full `model.safetensors` — maximum compatibility, one gather per
|
||||
save under sharding.
|
||||
- DCP: sharded `pytorch_model_fsdp_0/*.distcp` only — fastest save/resume; convert with
|
||||
`lerobot-convert-dcp` before distributing.
|
||||
- SAFETENSORS_AND_DCP: both artifacts, written independently.
|
||||
"""
|
||||
|
||||
SAFETENSORS = "safetensors"
|
||||
DCP = "dcp"
|
||||
SAFETENSORS_AND_DCP = "safetensors_dcp"
|
||||
|
||||
@property
|
||||
def wants_safetensors(self) -> bool:
|
||||
"""True when a full `model.safetensors` artifact should be written."""
|
||||
return self in (CheckpointFormat.SAFETENSORS, CheckpointFormat.SAFETENSORS_AND_DCP)
|
||||
|
||||
@property
|
||||
def wants_dcp(self) -> bool:
|
||||
"""True when sharded DCP model shards (`pytorch_model_fsdp_0/`) should be written."""
|
||||
return self in (CheckpointFormat.DCP, CheckpointFormat.SAFETENSORS_AND_DCP)
|
||||
|
||||
|
||||
def _migrate_legacy_rabc_fields(config: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""Return migrated payload for legacy RA-BC fields, or None when no migration is needed."""
|
||||
legacy_fields = (
|
||||
@@ -121,9 +152,16 @@ class TrainPipelineConfig(HubMixin):
|
||||
# Checkpoint is saved every `save_freq` training iterations and after the last training step.
|
||||
# A non-positive value disables periodic saving, keeping only the final checkpoint.
|
||||
save_freq: int = 20_000
|
||||
# Model-artifact format inside checkpoints; non-default values require a sharded run.
|
||||
checkpoint_format: CheckpointFormat = CheckpointFormat.SAFETENSORS
|
||||
use_policy_training_preset: bool = True
|
||||
optimizer: OptimizerConfig | None = None
|
||||
scheduler: LRSchedulerConfig | None = None
|
||||
# Process topology: dp_replicate / dp_shard (HSDP) and context-parallel degree placeholders.
|
||||
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
|
||||
# Execution runtime handed to the Accelerator: mixed precision, gradient accumulation,
|
||||
# FSDP/DDP tuning knobs, compile & activation-checkpointing placeholders.
|
||||
accelerator: AcceleratorConfig = field(default_factory=AcceleratorConfig)
|
||||
eval: EvalConfig = field(default_factory=EvalConfig)
|
||||
wandb: WandBConfig = field(default_factory=WandBConfig)
|
||||
peft: PeftConfig | None = None
|
||||
@@ -291,6 +329,60 @@ class TrainPipelineConfig(HubMixin):
|
||||
if self.save_checkpoint_to_hub and not (self.policy is not None and self.policy.repo_id):
|
||||
raise ValueError("save_checkpoint_to_hub requires --policy.repo_id.")
|
||||
|
||||
self._validate_distributed()
|
||||
|
||||
def _validate_distributed(self) -> None:
|
||||
"""Fail-fasts for the distributed-training scope.
|
||||
|
||||
Raises:
|
||||
ValueError: If the config requests anything outside the verified scope: context
|
||||
parallelism or CFG parallelism (reserved placeholders), the compile or
|
||||
activation-checkpointing placeholders, a DCP checkpoint format on a
|
||||
non-sharded run, or — under sharded training — fp16 mixed precision, PEFT,
|
||||
reward-model training, in-training environment evaluation, or multi-optimizer
|
||||
configs.
|
||||
"""
|
||||
if self.parallelism.cp_size > 1:
|
||||
raise ValueError(
|
||||
"Context parallelism is not implemented yet: --parallelism.context_parallel.* "
|
||||
"degrees must be 1 (reserved for the CP engine round)."
|
||||
)
|
||||
if self.parallelism.cfg_parallel != 1:
|
||||
raise ValueError(
|
||||
"CFG parallelism is inference-only and must be 1 for training "
|
||||
"(cfg_parallel is reserved for the serving round)."
|
||||
)
|
||||
if self.accelerator.compile.enabled:
|
||||
raise ValueError("--accelerator.compile is a placeholder and not wired yet.")
|
||||
if self.accelerator.activation_checkpointing.mode is not ActivationCheckpointingMode.NONE:
|
||||
raise ValueError("--accelerator.activation_checkpointing is a placeholder and not wired yet.")
|
||||
if self.checkpoint_format is not CheckpointFormat.SAFETENSORS and not self.parallelism.is_sharded:
|
||||
raise ValueError(
|
||||
f"checkpoint_format={self.checkpoint_format.value} requires a sharded run "
|
||||
"(--parallelism.dp_shard != 1); non-sharded checkpoints are always safetensors."
|
||||
)
|
||||
if self.parallelism.is_sharded:
|
||||
if self.accelerator.mixed_precision == "fp16":
|
||||
raise ValueError(
|
||||
"fp16 is not supported under sharded training (GradScaler over DTensor "
|
||||
"gradients is unverified); use bf16 or full precision."
|
||||
)
|
||||
if self.peft is not None:
|
||||
raise ValueError("PEFT is not supported under sharded training yet.")
|
||||
if self.is_reward_model_training:
|
||||
raise ValueError(
|
||||
"Reward-model training is not supported under sharded training yet "
|
||||
"(reward models declare no FSDP wrap units and have no sharded save path)."
|
||||
)
|
||||
if self.env is not None and self.env_eval_freq > 0:
|
||||
raise ValueError(
|
||||
"In-training environment evaluation is not supported under sharded training "
|
||||
"(a rank-0-only rollout of a sharded model deadlocks on collectives); set "
|
||||
"--env_eval_freq=0 and evaluate with lerobot-eval on saved checkpoints."
|
||||
)
|
||||
if self.optimizer is not None and self.optimizer.builds_multiple_optimizers:
|
||||
raise ValueError("Multi-optimizer configs are not supported under sharded training.")
|
||||
|
||||
@classmethod
|
||||
def __get_path_fields__(cls) -> list[str]:
|
||||
"""Keys for draccus pretrained-path loading."""
|
||||
|
||||
@@ -76,7 +76,7 @@ import torch
|
||||
from pydantic import BaseModel, Field
|
||||
from transformers import AutoProcessor, Qwen3VLMoeForConditionalGeneration
|
||||
|
||||
from lerobot.datasets import LeRobotDataset
|
||||
from lerobot.datasets import LeRobotDataset, resolve_episode_indices
|
||||
|
||||
|
||||
# Pydantic Models for SARM Subtask Annotation
|
||||
@@ -1049,7 +1049,10 @@ def main():
|
||||
torch_dtype = {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}[args.dtype]
|
||||
|
||||
# Determine episodes
|
||||
episode_indices = args.episodes or list(range(dataset.meta.total_episodes))
|
||||
resolved_episodes = resolve_episode_indices(args.episodes, dataset.meta.total_episodes)
|
||||
episode_indices = (
|
||||
resolved_episodes if resolved_episodes is not None else list(range(dataset.meta.total_episodes))
|
||||
)
|
||||
|
||||
existing_annotations = load_annotations_from_dataset(dataset.root, prefix="sparse")
|
||||
if args.skip_existing:
|
||||
|
||||
@@ -52,7 +52,7 @@ from .pipeline_features import aggregate_pipeline_dataset_features, create_initi
|
||||
from .pyav_utils import check_video_encoder_parameters_pyav, detect_available_encoders_pyav
|
||||
from .sampler import EpisodeAwareSampler, compute_sampler_state
|
||||
from .streaming_dataset import StreamingLeRobotDataset
|
||||
from .utils import DEFAULT_EPISODES_PATH, create_lerobot_dataset_card
|
||||
from .utils import DEFAULT_EPISODES_PATH, create_lerobot_dataset_card, resolve_episode_indices
|
||||
from .video_utils import VideoEncodingManager
|
||||
|
||||
# NOTE: Low-level I/O functions (cast_stats_to_numpy, get_parquet_file_size_in_mb, etc.)
|
||||
@@ -97,6 +97,7 @@ __all__ = [
|
||||
"reencode_dataset",
|
||||
"remove_feature",
|
||||
"resolve_delta_timestamps",
|
||||
"resolve_episode_indices",
|
||||
"safe_stop_image_writer",
|
||||
"split_dataset",
|
||||
"write_stats",
|
||||
|
||||
@@ -22,6 +22,7 @@ from pathlib import Path
|
||||
from typing import Any, NotRequired, TypedDict
|
||||
|
||||
import datasets
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import tqdm
|
||||
|
||||
@@ -303,6 +304,46 @@ def update_meta_data(
|
||||
df["dataset_to_index"] = df["dataset_to_index"] + dst_meta.info.total_frames
|
||||
df["episode_index"] = df["episode_index"] + dst_meta.info.total_episodes
|
||||
|
||||
# Per-episode stats still describe the pre-merge values of the bookkeeping columns
|
||||
# reindexed above. index/episode_index shift by a constant; task_index is relabeled,
|
||||
# so recompute it from the episode's (stable) task strings via the unified tasks table.
|
||||
shift_stat_keys = ("min", "max", "mean", "q01", "q10", "q50", "q90", "q99")
|
||||
for name, offset in (
|
||||
("episode_index", dst_meta.info.total_episodes),
|
||||
("index", dst_meta.info.total_frames),
|
||||
):
|
||||
for stat in shift_stat_keys:
|
||||
col = f"stats/{name}/{stat}"
|
||||
if col in df.columns:
|
||||
df[col] = df[col] + offset
|
||||
|
||||
if any(c.startswith("stats/task_index/") for c in df.columns):
|
||||
quantiles = {"q01": 0.01, "q10": 0.10, "q50": 0.50, "q90": 0.90, "q99": 0.99}
|
||||
ids_per_row = [
|
||||
np.array([dst_meta.tasks.loc[t, "task_index"] for t in tasks], dtype=np.float64)
|
||||
for tasks in df["tasks"]
|
||||
]
|
||||
|
||||
def _task_stat(ids, stat):
|
||||
if stat == "min":
|
||||
return ids.min()
|
||||
if stat == "max":
|
||||
return ids.max()
|
||||
if stat == "std":
|
||||
return ids.std()
|
||||
if stat in quantiles:
|
||||
return np.quantile(ids, quantiles[stat])
|
||||
return ids.mean()
|
||||
|
||||
for stat in ("min", "max", "mean", "std", *quantiles):
|
||||
col = f"stats/task_index/{stat}"
|
||||
if col in df.columns:
|
||||
# np.full_like preserves each cell container and dtype so the parquet schema is unchanged.
|
||||
df[col] = [
|
||||
np.full_like(orig, _task_stat(ids, stat))
|
||||
for orig, ids in zip(df[col], ids_per_row, strict=True)
|
||||
]
|
||||
|
||||
return df
|
||||
|
||||
|
||||
|
||||
@@ -39,6 +39,7 @@ from .io_utils import (
|
||||
hf_transform_to_torch,
|
||||
load_nested_dataset,
|
||||
)
|
||||
from .utils import resolve_episode_indices
|
||||
from .video_utils import decode_video_frames
|
||||
|
||||
|
||||
@@ -83,7 +84,7 @@ class DatasetReader:
|
||||
"""
|
||||
self._meta = meta
|
||||
self.root = root
|
||||
self.episodes = episodes
|
||||
self.episodes = resolve_episode_indices(episodes, meta.total_episodes)
|
||||
self._tolerance_s = tolerance_s
|
||||
self._video_backend = video_backend
|
||||
if image_transforms is not None and not callable(image_transforms):
|
||||
@@ -163,10 +164,34 @@ class DatasetReader:
|
||||
def _load_hf_dataset(self) -> datasets.Dataset:
|
||||
"""hf_dataset contains all the observations, states, actions, rewards, etc."""
|
||||
features = get_hf_features_from_features(self._meta.features)
|
||||
self._validate_language_columns_declared(features)
|
||||
hf_dataset = load_nested_dataset(self.root / "data", features=features, episodes=self.episodes)
|
||||
hf_dataset.set_transform(hf_transform_to_torch)
|
||||
return hf_dataset
|
||||
|
||||
def _validate_language_columns_declared(self, features: datasets.Features) -> None:
|
||||
"""Require language columns stored in Parquet to be declared in metadata."""
|
||||
# Leave empty datasets to fail through the normal loading path.
|
||||
try:
|
||||
sample = next((self.root / "data").glob("*/*.parquet"))
|
||||
except StopIteration:
|
||||
return
|
||||
|
||||
from pyarrow import parquet as _pq # noqa: PLC0415
|
||||
|
||||
# LeRobot shards are schema-uniform, so one schema represents the dataset.
|
||||
schema_names = set(_pq.read_schema(sample).names)
|
||||
from .language import LANGUAGE_COLUMNS # noqa: PLC0415
|
||||
|
||||
missing = sorted(set(LANGUAGE_COLUMNS) & schema_names - set(features))
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"Dataset Parquet files contain language feature(s) missing from metadata: {missing}. "
|
||||
"Metadata must describe the stored data; add the entries returned by "
|
||||
"lerobot.datasets.language.language_feature_info() to meta/info.json['features'] "
|
||||
"or rerun the annotation pipeline's metadata synchronization."
|
||||
)
|
||||
|
||||
def _check_cached_episodes_sufficient(self) -> bool:
|
||||
"""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:
|
||||
|
||||
@@ -29,6 +29,7 @@ from .dataset_metadata import LeRobotDatasetMetadata
|
||||
from .lerobot_dataset import LeRobotDataset
|
||||
from .multi_dataset import MultiLeRobotDataset
|
||||
from .streaming_dataset import StreamingLeRobotDataset
|
||||
from .utils import resolve_episode_indices
|
||||
|
||||
|
||||
def resolve_delta_timestamps(
|
||||
@@ -84,14 +85,24 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
|
||||
if isinstance(cfg.dataset.repo_id, str):
|
||||
ds_meta = LeRobotDatasetMetadata(
|
||||
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
||||
cfg.dataset.repo_id,
|
||||
root=cfg.dataset.root,
|
||||
revision=cfg.dataset.revision,
|
||||
repo_type=cfg.dataset.repo_type,
|
||||
)
|
||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
|
||||
episodes = resolve_episode_indices(
|
||||
cfg.dataset.episodes, ds_meta.total_episodes, cfg.dataset.exclude_episodes
|
||||
)
|
||||
if not cfg.dataset.streaming:
|
||||
if cfg.dataset.repo_type == "bucket":
|
||||
raise ValueError(
|
||||
"repo_type='bucket' is streaming-only: set dataset.streaming=true to train from an HF Storage Bucket."
|
||||
)
|
||||
dataset = LeRobotDataset(
|
||||
cfg.dataset.repo_id,
|
||||
root=cfg.dataset.root,
|
||||
episodes=cfg.dataset.episodes,
|
||||
episodes=episodes,
|
||||
delta_timestamps=delta_timestamps,
|
||||
image_transforms=image_transforms,
|
||||
revision=cfg.dataset.revision,
|
||||
@@ -104,13 +115,14 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
dataset = StreamingLeRobotDataset(
|
||||
cfg.dataset.repo_id,
|
||||
root=cfg.dataset.root,
|
||||
episodes=cfg.dataset.episodes,
|
||||
episodes=episodes,
|
||||
delta_timestamps=delta_timestamps,
|
||||
image_transforms=image_transforms,
|
||||
revision=cfg.dataset.revision,
|
||||
max_num_shards=cfg.num_workers,
|
||||
tolerance_s=cfg.tolerance_s,
|
||||
return_uint8=True,
|
||||
repo_type=cfg.dataset.repo_type,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError("The MultiLeRobotDataset isn't supported for now.")
|
||||
|
||||
@@ -162,14 +162,32 @@ def render_sample(
|
||||
task: str | None = None,
|
||||
dataset_ctx: Any | None = None,
|
||||
) -> RenderedMessages | None:
|
||||
"""Render the chat-style messages for a single dataset sample.
|
||||
"""Render recipe-defined messages and supervision for one dataset sample.
|
||||
|
||||
Resolves the recipe's bindings against ``persistent`` and ``events`` rows
|
||||
at frame timestamp ``t``, then expands the recipe's message templates.
|
||||
Returns ``None`` if the resolved sample contains no target message.
|
||||
Resolves bindings against ``persistent`` and ``events`` at frame timestamp
|
||||
``t``. Blend recipes first route matching sparse VQA annotations, then use
|
||||
deterministic weighted selection for the remaining samples. Returns
|
||||
``None`` when the selected recipe provides no text or low-level action
|
||||
supervision for this sample.
|
||||
"""
|
||||
persistent_rows = _normalize_rows(persistent 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)
|
||||
bindings = _resolve_bindings(
|
||||
selected_recipe,
|
||||
@@ -183,6 +201,58 @@ def render_sample(
|
||||
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.
|
||||
"""
|
||||
if recipe.blend is None:
|
||||
return None
|
||||
renderable: list[tuple[float, RenderedMessages]] = []
|
||||
for component in recipe.blend.values():
|
||||
if component.route != "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:
|
||||
if component.weight is None:
|
||||
raise ValueError("Routed VQA blend components must define a weight.")
|
||||
renderable.append((component.weight, rendered))
|
||||
|
||||
if not renderable:
|
||||
return None
|
||||
if len(renderable) == 1:
|
||||
return renderable[0][1]
|
||||
|
||||
# Choose among matching cameras by their validated positive relative weights.
|
||||
total = sum(weight for weight, _ in 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 weight, rendered in renderable:
|
||||
cumulative += weight
|
||||
if draw < cumulative:
|
||||
return rendered
|
||||
return renderable[-1][1]
|
||||
|
||||
|
||||
def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
||||
"""Pick a deterministic blend component for ``sample_idx`` (or return ``recipe``)."""
|
||||
if recipe.blend is None:
|
||||
@@ -201,7 +271,8 @@ def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
||||
cumulative += component.weight or 0.0
|
||||
if draw < cumulative:
|
||||
return component
|
||||
assert last_component is not None
|
||||
if last_component is None:
|
||||
raise ValueError("Blend recipes must contain at least one component.")
|
||||
return last_component
|
||||
|
||||
|
||||
@@ -321,7 +392,8 @@ def _render_message_recipe(
|
||||
bindings: dict[str, LanguageRow | str | None],
|
||||
) -> RenderedMessages | None:
|
||||
"""Expand ``recipe.messages`` into rendered chat messages using ``bindings``."""
|
||||
assert recipe.messages is not None
|
||||
if recipe.messages is None:
|
||||
raise ValueError("Cannot render a blend recipe as a message recipe.")
|
||||
messages: list[dict[str, Any]] = []
|
||||
streams: list[str | None] = []
|
||||
target_indices: list[int] = []
|
||||
@@ -346,7 +418,9 @@ def _render_message_recipe(
|
||||
if turn.target:
|
||||
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
|
||||
|
||||
rendered = {
|
||||
@@ -403,14 +477,12 @@ def _validate_rendered(rendered: RenderedMessages) -> None:
|
||||
|
||||
if len(streams) != len(messages):
|
||||
raise ValueError("message_streams must be aligned with messages.")
|
||||
if not target_indices:
|
||||
raise ValueError("Rendered samples must contain at least one target message.")
|
||||
# Require text or low-level action supervision.
|
||||
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:
|
||||
if idx < 0 or idx >= len(messages):
|
||||
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(
|
||||
|
||||
@@ -18,6 +18,7 @@ import dataclasses
|
||||
import importlib.resources
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
@@ -98,6 +99,47 @@ VIDEO_DIR = "videos"
|
||||
|
||||
CHUNK_FILE_PATTERN = "chunk-{chunk_index:03d}/file-{file_index:03d}"
|
||||
IMAGE_FILE_PATTERN = "frame-{frame_index:06d}.png"
|
||||
|
||||
|
||||
def resolve_episode_indices(
|
||||
episodes: Sequence[int] | None,
|
||||
total_episodes: int,
|
||||
exclude_episodes: Sequence[int] | None = None,
|
||||
) -> list[int] | None:
|
||||
"""Resolve an optional episode allowlist and exclusion list against dataset bounds.
|
||||
|
||||
``None`` is preserved when no filtering is requested so callers can retain
|
||||
their native "all episodes" fast path. Invalid indices are ignored with a
|
||||
warning, and the input order is preserved.
|
||||
"""
|
||||
if total_episodes < 0:
|
||||
raise ValueError(f"total_episodes must be non-negative, got {total_episodes}")
|
||||
|
||||
if episodes is None and not exclude_episodes:
|
||||
return None
|
||||
|
||||
candidates = list(range(total_episodes)) if episodes is None else list(episodes)
|
||||
invalid = [episode for episode in candidates if not 0 <= episode < total_episodes]
|
||||
if invalid:
|
||||
logger.warning(
|
||||
"Ignoring episode indices outside the dataset range [0, %d): %s",
|
||||
total_episodes,
|
||||
invalid,
|
||||
)
|
||||
candidates = [episode for episode in candidates if 0 <= episode < total_episodes]
|
||||
|
||||
excluded = set(exclude_episodes or [])
|
||||
invalid_excluded = sorted(episode for episode in excluded if not 0 <= episode < total_episodes)
|
||||
if invalid_excluded:
|
||||
logger.warning(
|
||||
"Ignoring excluded episode indices outside the dataset range [0, %d): %s",
|
||||
total_episodes,
|
||||
invalid_excluded,
|
||||
)
|
||||
excluded = {episode for episode in excluded if 0 <= episode < total_episodes}
|
||||
return [episode for episode in candidates if episode not in excluded]
|
||||
|
||||
|
||||
DEPTH_FILE_PATTERN = "frame-{frame_index:06d}.tiff"
|
||||
DEFAULT_TASKS_PATH = "meta/tasks.parquet"
|
||||
DEFAULT_EPISODES_PATH = EPISODES_DIR + "/" + CHUNK_FILE_PATTERN + ".parquet"
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
#!/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.
|
||||
"""Distributed-training runtime for LeRobot.
|
||||
|
||||
This package owns everything that turns the declarative topology in
|
||||
:class:`lerobot.configs.parallelism.ParallelismConfig` into a running engine:
|
||||
mesh math (:class:`~lerobot.distributed.parallel_dims.ParallelDims`), the
|
||||
`Accelerator` factory (:func:`~lerobot.distributed.factory.make_accelerator`),
|
||||
sharding-aware checkpoint helpers, and small rank utilities.
|
||||
|
||||
Setup-order contract (normative):
|
||||
CP dispatch install -> activation checkpointing -> torch.compile ->
|
||||
``fully_shard``/DDP (via ``accelerator.prepare``) -> optimizer rebind.
|
||||
Only the last two steps are active today; CP/AC/compile are configured
|
||||
placeholders wired in later rounds.
|
||||
"""
|
||||
|
||||
from .factory import guard_against_env_interference, make_accelerator, set_fsdp_wrap_modules
|
||||
from .parallel_dims import ParallelDims
|
||||
from .utils import finalize_sharded_policy, is_main_process, strip_accelerate_cp_hooks
|
||||
|
||||
__all__ = [
|
||||
"ParallelDims",
|
||||
"finalize_sharded_policy",
|
||||
"guard_against_env_interference",
|
||||
"is_main_process",
|
||||
"make_accelerator",
|
||||
"set_fsdp_wrap_modules",
|
||||
"strip_accelerate_cp_hooks",
|
||||
]
|
||||
@@ -0,0 +1,195 @@
|
||||
#!/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.
|
||||
"""Sharding-aware checkpoint primitives.
|
||||
|
||||
Two artifact channels with distinct owners:
|
||||
|
||||
- the **distributable** ``model.safetensors``: produced by ``PreTrainedPolicy.save_pretrained``
|
||||
through :func:`full_model_state_dict` — a collective full gather when the model is sharded;
|
||||
- the **resume** channel (sharded runs): torch DCP directories written/read through accelerate's
|
||||
``save/load_fsdp_model`` and ``save/load_fsdp_optimizer`` (``pytorch_model_fsdp_0/`` and
|
||||
``optimizer_0/``, names imported from accelerate constants), which reshard on load across
|
||||
topology changes.
|
||||
|
||||
Every function that touches sharded state is a collective and must run on ALL ranks.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from accelerate import Accelerator
|
||||
|
||||
|
||||
def is_sharded_module(module: nn.Module) -> bool:
|
||||
"""True when `fully_shard` owns this module's parameters (FSDP2's in-place class swap).
|
||||
|
||||
Args:
|
||||
module (nn.Module): The module to inspect (a torch.compile wrapper is looked through
|
||||
via `_orig_mod`).
|
||||
|
||||
Returns:
|
||||
bool: True when the module (or its compiled `_orig_mod`) is an `FSDPModule`.
|
||||
"""
|
||||
from torch.distributed.fsdp import FSDPModule
|
||||
|
||||
if isinstance(module, FSDPModule):
|
||||
return True
|
||||
# torch.compile wraps the sharded module; mirror accelerate's `_orig_mod` check.
|
||||
orig_mod = getattr(module, "_orig_mod", None)
|
||||
return orig_mod is not None and isinstance(orig_mod, FSDPModule)
|
||||
|
||||
|
||||
def full_model_state_dict(module: nn.Module) -> dict[str, torch.Tensor]:
|
||||
"""The module's full (unsharded) state dict, however its parameters are laid out.
|
||||
|
||||
Sharded modules gather through torch's DCP state-dict API: a COLLECTIVE that must run on
|
||||
every rank; with ``cpu_offload=True`` the full dict materializes on the main rank only and
|
||||
every other rank receives a literal ``{}`` (runtime-verified — a
|
||||
rank-0-gated call deadlocks). Plain modules return ``module.state_dict()`` on every rank.
|
||||
|
||||
Args:
|
||||
module (nn.Module): The (possibly sharded) module to read the state dict from.
|
||||
|
||||
Returns:
|
||||
dict[str, torch.Tensor]: The full state dict — on the main rank only (``{}``
|
||||
elsewhere) when the module is sharded, on every rank otherwise.
|
||||
"""
|
||||
if not is_sharded_module(module):
|
||||
return module.state_dict()
|
||||
|
||||
from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict
|
||||
|
||||
return get_model_state_dict(module, options=StateDictOptions(full_state_dict=True, cpu_offload=True))
|
||||
|
||||
|
||||
def _fsdp_plugin(accelerator: "Accelerator") -> object:
|
||||
"""The accelerator's FSDP plugin, required by every DCP save/load helper below.
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator that prepared the sharded model.
|
||||
|
||||
Returns:
|
||||
object: The FSDP plugin held by `accelerator.state`.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the accelerator was not configured with an FSDP plugin.
|
||||
"""
|
||||
plugin = getattr(accelerator.state, "fsdp_plugin", None)
|
||||
if plugin is None:
|
||||
raise RuntimeError("Sharded checkpointing requires an FSDP-prepared Accelerator.")
|
||||
return plugin
|
||||
|
||||
|
||||
def save_sharded_model(accelerator: "Accelerator", model: nn.Module, output_dir: Path) -> None:
|
||||
"""Write the DCP model shards (`pytorch_model_fsdp_0/`). Collective: call on all ranks.
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator that prepared the sharded model.
|
||||
model (nn.Module): The prepared (sharded) model to save.
|
||||
output_dir (Path): The directory the shard subdirectory is created in.
|
||||
"""
|
||||
from accelerate.utils import save_fsdp_model
|
||||
|
||||
# accelerate 1.14's DCP helpers do string containment checks on the path:
|
||||
# always hand them str, never Path.
|
||||
save_fsdp_model(_fsdp_plugin(accelerator), accelerator, model, str(output_dir))
|
||||
|
||||
|
||||
def load_sharded_model(accelerator: "Accelerator", model: nn.Module, input_dir: Path) -> None:
|
||||
"""Load DCP model shards into the prepared (sharded) model. Collective: call on all ranks.
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator that prepared the sharded model.
|
||||
model (nn.Module): The prepared (sharded) model to load into.
|
||||
input_dir (Path): The directory containing the `pytorch_model_fsdp_0/` shard
|
||||
subdirectory.
|
||||
"""
|
||||
from accelerate.utils import load_fsdp_model
|
||||
from accelerate.utils.constants import FSDP_MODEL_NAME
|
||||
|
||||
# Pass the exact shard directory: accelerate's load resolves it with a substring check
|
||||
# ("pytorch_model_fsdp" in the path -> use as-is), which misfires on run paths that happen
|
||||
# to contain the marker; the exact dir makes the check deterministic.
|
||||
load_fsdp_model(_fsdp_plugin(accelerator), accelerator, model, str(input_dir / f"{FSDP_MODEL_NAME}_0"))
|
||||
|
||||
|
||||
def save_sharded_optimizer(
|
||||
accelerator: "Accelerator", optimizer: torch.optim.Optimizer, model: nn.Module, output_dir: Path
|
||||
) -> None:
|
||||
"""Write the DCP optimizer shards (`optimizer_0/`). Collective: call on all ranks.
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator that prepared the model and optimizer.
|
||||
optimizer (torch.optim.Optimizer): The prepared optimizer to save the state from.
|
||||
model (nn.Module): The prepared (sharded) model the optimizer state is keyed by.
|
||||
output_dir (Path): The directory the shard subdirectory is created in.
|
||||
"""
|
||||
from accelerate.utils import save_fsdp_optimizer
|
||||
|
||||
save_fsdp_optimizer(_fsdp_plugin(accelerator), accelerator, optimizer, model, str(output_dir))
|
||||
|
||||
|
||||
def load_sharded_optimizer(
|
||||
accelerator: "Accelerator", optimizer: torch.optim.Optimizer, model: nn.Module, input_dir: Path
|
||||
) -> None:
|
||||
"""Load DCP optimizer shards into the prepared optimizer. Collective: call on all ranks.
|
||||
|
||||
Must run AFTER ``accelerator.prepare()``: FSDP2's prepare rebinds the optimizer's param
|
||||
groups to sharded DTensors but never migrates ``optimizer.state`` — the resharding load is
|
||||
the only correct way to restore it.
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator that prepared the model and optimizer.
|
||||
optimizer (torch.optim.Optimizer): The prepared optimizer to restore the state into.
|
||||
model (nn.Module): The prepared (sharded) model the optimizer state is keyed by.
|
||||
input_dir (Path): The directory containing the `optimizer_0/` shard subdirectory.
|
||||
"""
|
||||
from accelerate.utils import load_fsdp_optimizer
|
||||
from accelerate.utils.constants import OPTIMIZER_NAME
|
||||
|
||||
# Exact shard directory for the same reason as load_sharded_model: accelerate's substring
|
||||
# check ("optimizer" in the path) would misread e.g. --job_name=optimizer_sweep run paths.
|
||||
load_fsdp_optimizer(
|
||||
_fsdp_plugin(accelerator), accelerator, optimizer, model, str(input_dir / f"{OPTIMIZER_NAME}_0")
|
||||
)
|
||||
|
||||
|
||||
def dcp_to_safetensors(dcp_dir: Path, output_dir: Path, *, delete_dcp: bool = False) -> Path:
|
||||
"""Merge a DCP shard directory into a single `model.safetensors` (offline, single process).
|
||||
|
||||
Thin wrapper over `accelerate.utils.merge_fsdp_weights`, which loads the shards without a
|
||||
process group, writes safetensors directly, and — when asked — removes the merged shard
|
||||
directory itself, only on the main process and only once the merge has succeeded.
|
||||
|
||||
Args:
|
||||
dcp_dir (Path): The DCP shard directory to merge (e.g. `.../pytorch_model_fsdp_0`).
|
||||
output_dir (Path): The directory the merged `model.safetensors` is written into.
|
||||
delete_dcp (bool): Whether to remove the shard directory once it has been merged.
|
||||
Defaults to False.
|
||||
|
||||
Returns:
|
||||
Path: The written `model.safetensors` file's path.
|
||||
"""
|
||||
from accelerate.utils import merge_fsdp_weights
|
||||
|
||||
merge_fsdp_weights(
|
||||
str(dcp_dir), str(output_dir), safe_serialization=True, remove_checkpoint_dir=delete_dcp
|
||||
)
|
||||
return output_dir / "model.safetensors"
|
||||
@@ -0,0 +1,147 @@
|
||||
#!/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.
|
||||
"""The `Accelerator` factory — the only place accelerate gets configured.
|
||||
|
||||
`torchrun` is the launcher; every accelerate parameter comes from `TrainPipelineConfig`
|
||||
(`cfg.parallelism` + `cfg.accelerator`) so a run is reproducible from its `train_config.json`
|
||||
alone. `accelerate launch` without a `--config_file` remains equivalent (it only sets rendezvous
|
||||
env vars in that mode); the yaml flow is superseded.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from lerobot.configs.parallelism import world_size_from_env
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from accelerate import Accelerator
|
||||
|
||||
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||
|
||||
# Env vars through which `accelerate launch --config_file` (or a stray shell) would configure
|
||||
# accelerate behind the config system's back. Plugin `__post_init__`s read these silently as
|
||||
# field fallbacks (ACCELERATE_DYNAMO_* enables torch.compile through the default
|
||||
# TorchDynamoPlugin; ACCELERATE_GRADIENT_ACCUMULATION_STEPS overrides the explicitly passed
|
||||
# value inside Accelerator.__init__), which would make train_config.json lie about what ran.
|
||||
_ACCELERATE_ENV_PREFIXES = ("FSDP_", "PARALLELISM_CONFIG_", "ACCELERATE_DYNAMO_")
|
||||
_ACCELERATE_ENV_VARS = (
|
||||
"ACCELERATE_USE_FSDP",
|
||||
"ACCELERATE_USE_PARALLELISM_CONFIG",
|
||||
"ACCELERATE_GRADIENT_ACCUMULATION_STEPS",
|
||||
)
|
||||
_ENV_OVERRIDE = "LEROBOT_ALLOW_ACCELERATE_ENV"
|
||||
|
||||
|
||||
def guard_against_env_interference() -> None:
|
||||
"""Hard-error when accelerate-configuring env vars are set.
|
||||
|
||||
A silently env-overridden "reproducible" config is worse than a stop: users migrating from
|
||||
the old `accelerate launch --config_file fsdp.yaml` flow get a precise error instead of a
|
||||
config that lies. Set LEROBOT_ALLOW_ACCELERATE_ENV=1 to acknowledge and proceed.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If any accelerate-configuring environment variable is set and the
|
||||
LEROBOT_ALLOW_ACCELERATE_ENV override is not.
|
||||
"""
|
||||
if os.environ.get(_ENV_OVERRIDE):
|
||||
return
|
||||
offending = sorted(
|
||||
name
|
||||
for name in os.environ
|
||||
if name in _ACCELERATE_ENV_VARS or name.startswith(_ACCELERATE_ENV_PREFIXES)
|
||||
)
|
||||
if offending:
|
||||
raise RuntimeError(
|
||||
f"Accelerate-configuring environment variables are set: {', '.join(offending)}. "
|
||||
"LeRobot manages accelerate exclusively through TrainPipelineConfig "
|
||||
"(--parallelism.* / --accelerator.*); launch with plain torchrun and remove these "
|
||||
"variables (the `accelerate launch --config_file` flow is superseded), or set "
|
||||
f"{_ENV_OVERRIDE}=1 to acknowledge that they may override your config."
|
||||
)
|
||||
|
||||
|
||||
def make_accelerator(cfg: TrainPipelineConfig) -> "Accelerator":
|
||||
"""Resolve the topology against the launched world and build the `Accelerator`.
|
||||
|
||||
Must run once per process, before any other component needs the device or the process
|
||||
group (`Accelerator.__init__` initializes both and builds the device mesh).
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The full training config; `cfg.parallelism` is resolved in
|
||||
place against the launched world size and `cfg.accelerator` builds the result.
|
||||
|
||||
Returns:
|
||||
Accelerator: The configured accelerator, with device and process group initialized.
|
||||
|
||||
Raises:
|
||||
ValueError: If `cfg.checkpoint_format` requires DCP but the topology resolved to a
|
||||
non-sharded run.
|
||||
"""
|
||||
guard_against_env_interference()
|
||||
cfg.parallelism.resolve(world_size_from_env())
|
||||
# The parse-time format check ran against the declared degrees, where the dp_shard=-1
|
||||
# sentinel counts as sharded; it may resolve to an unsharded run (e.g. -1 at world size 1).
|
||||
# Re-check against the concrete degrees so the recorded format never lies about the
|
||||
# artifacts a checkpoint will actually contain.
|
||||
if cfg.checkpoint_format.wants_dcp and not cfg.parallelism.is_sharded:
|
||||
raise ValueError(
|
||||
f"checkpoint_format={cfg.checkpoint_format.value} requires a sharded run, but the "
|
||||
f"topology resolved to a non-sharded one (dp_replicate={cfg.parallelism.dp_replicate}, "
|
||||
f"dp_shard={cfg.parallelism.dp_shard}); non-sharded checkpoints are always safetensors."
|
||||
)
|
||||
return cfg.accelerator.build(
|
||||
cfg.parallelism,
|
||||
cpu=cfg.trainable_config.device == "cpu",
|
||||
)
|
||||
|
||||
|
||||
def set_fsdp_wrap_modules(accelerator: "Accelerator", policy: "PreTrainedPolicy") -> None:
|
||||
"""Resolve the FSDP wrap-unit class names onto the plugin before `accelerator.prepare()`.
|
||||
|
||||
Resolution order: user override (`--accelerator.fsdp.wrap_modules`, already on the plugin)
|
||||
-> the policy's `_fsdp_wrap_modules` declaration -> hard error. Root-only wrapping — the
|
||||
silent default when no wrap source exists — is never accepted: it quietly forfeits all
|
||||
sharding memory savings.
|
||||
|
||||
No-op for the size-based policy (`--accelerator.fsdp.min_num_params`), which needs no class
|
||||
names, and for non-sharded runs (no fsdp plugin).
|
||||
|
||||
Args:
|
||||
accelerator (Accelerator): The accelerator whose FSDP plugin receives the wrap-unit
|
||||
class names.
|
||||
policy (PreTrainedPolicy): The trainable whose class may declare `_fsdp_wrap_modules`.
|
||||
|
||||
Raises:
|
||||
ValueError: If sharded class-based wrapping is configured but neither a user override
|
||||
nor a policy declaration supplies wrap-unit class names.
|
||||
"""
|
||||
plugin = getattr(accelerator.state, "fsdp_plugin", None)
|
||||
if plugin is None or plugin.min_num_params:
|
||||
return
|
||||
if plugin.transformer_cls_names_to_wrap: # user override, set at build time
|
||||
return
|
||||
# getattr, not attribute access: non-policy trainables (no `_fsdp_wrap_modules` attribute)
|
||||
# must reach the actionable error below, not an AttributeError.
|
||||
declared = getattr(type(policy), "_fsdp_wrap_modules", None)
|
||||
if not declared:
|
||||
raise ValueError(
|
||||
f"Policy '{type(policy).__name__}' declares no FSDP wrap units. Sharded training "
|
||||
"requires wrap-unit class names: set --accelerator.fsdp.wrap_modules='[\"MyBlock\"]' "
|
||||
"(or --accelerator.fsdp.min_num_params for a size-based policy), or declare "
|
||||
"`_fsdp_wrap_modules` on the policy class."
|
||||
)
|
||||
plugin.transformer_cls_names_to_wrap = list(declared)
|
||||
@@ -0,0 +1,112 @@
|
||||
#!/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.
|
||||
"""Runtime mesh math derived from the declarative :class:`ParallelismConfig`.
|
||||
|
||||
`ParallelDims` is the training script's single source of truth for topology-derived numbers
|
||||
(data-parallel world size and rank, sample accounting inputs) and — once the CP engine lands —
|
||||
the owner of LeRobot's private ``(dp_replicate, dp_shard, ring, ulysses)`` mesh. It is a runtime
|
||||
object and is never serialized (the config it derives from is what lands in
|
||||
``train_config.json``).
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch.distributed as dist
|
||||
|
||||
from lerobot.configs.parallelism import ParallelismConfig
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ParallelDims:
|
||||
"""Concrete parallelism degrees bound to a world size (canonical row-major rank layout)."""
|
||||
|
||||
dp_replicate: int
|
||||
dp_shard: int
|
||||
ring: int
|
||||
ulysses: int
|
||||
world_size: int
|
||||
device_type: str
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, cfg: ParallelismConfig, world_size: int, device_type: str) -> "ParallelDims":
|
||||
"""Bind a *resolved* config to the actual runtime world size (cross-checked here).
|
||||
|
||||
Args:
|
||||
cfg (ParallelismConfig): The declarative topology, already resolved via
|
||||
`ParallelismConfig.resolve(world_size)`.
|
||||
world_size (int): The launched world size the declared degrees must multiply to.
|
||||
device_type (str): The accelerator device type backing the mesh (e.g. "cuda").
|
||||
|
||||
Returns:
|
||||
ParallelDims: The concrete parallelism degrees bound to this world.
|
||||
|
||||
Raises:
|
||||
ValueError: If the config is unresolved (`dp_shard == -1`) or its degrees do not
|
||||
multiply to `world_size`.
|
||||
"""
|
||||
total = cfg.dp_replicate * cfg.dp_shard * cfg.cp_size
|
||||
if cfg.dp_shard == -1 or total != world_size:
|
||||
raise ValueError(
|
||||
f"ParallelismConfig is not resolved against this world: dp_replicate="
|
||||
f"{cfg.dp_replicate} * dp_shard={cfg.dp_shard} * cp={cfg.cp_size} != "
|
||||
f"world_size={world_size}. Call ParallelismConfig.resolve(world_size) first "
|
||||
"(make_accelerator does this)."
|
||||
)
|
||||
return cls(
|
||||
dp_replicate=cfg.dp_replicate,
|
||||
dp_shard=cfg.dp_shard,
|
||||
ring=cfg.context_parallel.ring_degree,
|
||||
ulysses=cfg.context_parallel.ulysses_degree,
|
||||
world_size=world_size,
|
||||
device_type=device_type,
|
||||
)
|
||||
|
||||
@property
|
||||
def cp_size(self) -> int:
|
||||
"""Total context-parallel degree (`ring * ulysses`)."""
|
||||
return self.ring * self.ulysses
|
||||
|
||||
@property
|
||||
def is_sharded(self) -> bool:
|
||||
"""Whether parameters are sharded (`dp_shard > 1` or any context parallelism)."""
|
||||
return self.dp_shard > 1 or self.cp_size > 1
|
||||
|
||||
@property
|
||||
def dp_world_size(self) -> int:
|
||||
"""Number of distinct data-parallel workers — the divisor for all sample accounting."""
|
||||
return self.dp_replicate * self.dp_shard
|
||||
|
||||
@property
|
||||
def dp_rank(self) -> int:
|
||||
"""This process's data-parallel coordinate (CP peers share one dp_rank).
|
||||
|
||||
With the canonical row-major layout and (ring, ulysses) innermost, CP peers are
|
||||
contiguous global ranks, so the dp coordinate is the integer quotient by cp_size —
|
||||
the same arithmetic accelerate's mesh-aware dataloader applies.
|
||||
"""
|
||||
global_rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
return global_rank // self.cp_size
|
||||
|
||||
def cp_mesh(self) -> None:
|
||||
"""Private (ring, ulysses) mesh for the CP engine — reserved for the CP round.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Always — context parallelism is not implemented yet.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Context parallelism is not implemented yet; ParallelDims.cp_mesh is reserved for "
|
||||
"the CP engine round (a private mesh aligned with accelerate's cp block)."
|
||||
)
|
||||
@@ -0,0 +1,94 @@
|
||||
#!/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.
|
||||
"""Rank utilities and post-`prepare()` sharding finalization."""
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch.distributed as dist
|
||||
from torch import nn
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.distributed.parallel_dims import ParallelDims
|
||||
|
||||
|
||||
def is_main_process() -> bool:
|
||||
"""True on the process that owns rank-0-only side effects (file writes, uploads, logging).
|
||||
|
||||
Torch-native on purpose: persistence code must not depend on an `Accelerator` handle —
|
||||
`_save_pretrained` and the hub publishers run in contexts that have none. Outside
|
||||
distributed runs every process is the main process.
|
||||
|
||||
Returns:
|
||||
bool: True when this process is rank 0 or no process group is initialized.
|
||||
"""
|
||||
return not dist.is_initialized() or dist.get_rank() == 0
|
||||
|
||||
|
||||
def strip_accelerate_cp_hooks(model: nn.Module) -> int:
|
||||
"""Remove accelerate's context-parallel forward-pre-hooks from every module.
|
||||
|
||||
When `cp_size > 1` is declared, `accelerator.prepare()` unconditionally attaches hooks that
|
||||
silently replace any `attention_mask` kwarg of `*self_attn` modules with `is_causal=True`
|
||||
(`accelerate.big_modeling._attach_context_parallel_hooks`) — mask corruption for policies
|
||||
with non-causal attention. LeRobot implements CP itself and never enters accelerate's CP
|
||||
context, so these hooks are pure hazard. Deterministically identified by their defining
|
||||
module; a version canary pins that identity.
|
||||
|
||||
Args:
|
||||
model (nn.Module): The prepared model to strip the hooks from (all submodules are
|
||||
visited).
|
||||
|
||||
Returns:
|
||||
int: The number of hooks removed.
|
||||
"""
|
||||
removed = 0
|
||||
for module in model.modules():
|
||||
for hook_id, hook in list(module._forward_pre_hooks.items()):
|
||||
if getattr(hook, "__module__", None) == "accelerate.big_modeling":
|
||||
del module._forward_pre_hooks[hook_id]
|
||||
module._forward_pre_hooks_with_kwargs.pop(hook_id, None)
|
||||
removed += 1
|
||||
return removed
|
||||
|
||||
|
||||
def finalize_sharded_policy(policy: nn.Module, parallel_dims: "ParallelDims") -> None:
|
||||
"""Sharding correctness protocol, applied once, immediately after `accelerator.prepare()`.
|
||||
|
||||
1. Strip accelerate's CP mask hooks (only attached when cp > 1 was declared).
|
||||
2. Register the policy's non-`forward` entry points (`_fsdp_forward_methods`) so FSDP2
|
||||
unshards parameters around `select_action` & co. — without this, any inference-style
|
||||
call on a sharded policy crashes on mixed Tensor/DTensor.
|
||||
|
||||
No-op for DDP/single-process runs.
|
||||
|
||||
Args:
|
||||
policy (nn.Module): The policy as returned by `accelerator.prepare()`.
|
||||
parallel_dims (ParallelDims): The run's resolved topology; decides whether the protocol
|
||||
applies.
|
||||
"""
|
||||
if not parallel_dims.is_sharded:
|
||||
return
|
||||
if parallel_dims.cp_size > 1:
|
||||
removed = strip_accelerate_cp_hooks(policy)
|
||||
logging.info("Stripped %d accelerate context-parallel attention-mask hooks.", removed)
|
||||
|
||||
from torch.distributed.fsdp import FSDPModule, register_fsdp_forward_method
|
||||
|
||||
if isinstance(policy, FSDPModule):
|
||||
for method_name in getattr(type(policy), "_fsdp_forward_methods", ()):
|
||||
if callable(getattr(policy, method_name, None)):
|
||||
register_fsdp_forward_method(policy, method_name)
|
||||
@@ -328,6 +328,7 @@ class LiberoEnv(EnvConfig):
|
||||
render_mode: str = "rgb_array"
|
||||
camera_name: str = "agentview_image,robot0_eye_in_hand_image"
|
||||
init_states: bool = True
|
||||
hard_reset: bool = True
|
||||
camera_name_mapping: dict[str, str] | None = None
|
||||
observation_height: int = 360
|
||||
observation_width: int = 360
|
||||
@@ -356,6 +357,8 @@ class LiberoEnv(EnvConfig):
|
||||
def __post_init__(self):
|
||||
if self.fps <= 0:
|
||||
raise ValueError(f"fps must be positive, got {self.fps}")
|
||||
if not self.hard_reset and not self.init_states:
|
||||
raise ValueError("hard_reset=False requires init_states=True")
|
||||
|
||||
if self.obs_type == "pixels":
|
||||
self.features[LIBERO_KEY_PIXELS_AGENTVIEW] = PolicyFeature(
|
||||
@@ -416,6 +419,7 @@ class LiberoEnv(EnvConfig):
|
||||
"observation_height": self.observation_height,
|
||||
"observation_width": self.observation_width,
|
||||
"control_freq": self.fps,
|
||||
"hard_reset": self.hard_reset,
|
||||
}
|
||||
if self.task_ids is not None:
|
||||
kwargs["task_ids"] = self.task_ids
|
||||
|
||||
@@ -128,10 +128,13 @@ class LiberoEnv(gym.Env):
|
||||
control_freq: int = 20,
|
||||
control_mode: str = "relative",
|
||||
is_libero_plus: bool = False,
|
||||
hard_reset: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
if control_freq <= 0:
|
||||
raise ValueError(f"control_freq must be positive, got {control_freq}")
|
||||
if not hard_reset and not init_states:
|
||||
raise ValueError("hard_reset=False requires init_states=True")
|
||||
self.task_id = task_id
|
||||
self.is_libero_plus = is_libero_plus
|
||||
self.obs_type = obs_type
|
||||
@@ -158,6 +161,7 @@ class LiberoEnv(gym.Env):
|
||||
self.camera_name_mapping = camera_name_mapping
|
||||
self.num_steps_wait = num_steps_wait
|
||||
self.control_freq = control_freq
|
||||
self.hard_reset = hard_reset
|
||||
self.episode_index = episode_index
|
||||
self.episode_length = episode_length
|
||||
# Load once and keep
|
||||
@@ -265,6 +269,9 @@ class LiberoEnv(gym.Env):
|
||||
camera_heights=self.observation_height,
|
||||
camera_widths=self.observation_width,
|
||||
control_freq=self.control_freq,
|
||||
# Soft resets skip LIBERO's model and renderer rebuild. They are opt-in
|
||||
# because settle steps can make their observations differ from hard resets.
|
||||
hard_reset=self.hard_reset,
|
||||
)
|
||||
env.reset()
|
||||
self._env = env
|
||||
@@ -377,8 +384,9 @@ class LiberoEnv(gym.Env):
|
||||
}
|
||||
)
|
||||
observation = self._format_raw_obs(raw_obs)
|
||||
if terminated:
|
||||
self.reset()
|
||||
# Return the terminal observation unchanged. The caller owns resetting after
|
||||
# termination; vector envs created below use NEXT_STEP autoreset. Resetting here
|
||||
# would therefore reset twice and skip an initial state.
|
||||
truncated = False
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
@@ -476,6 +484,7 @@ def create_libero_envs(
|
||||
print(f"Restricting to task_ids={task_ids_filter}")
|
||||
|
||||
is_async = env_cls is gym.vector.AsyncVectorEnv
|
||||
is_sync = env_cls is gym.vector.SyncVectorEnv
|
||||
|
||||
out: dict[str, dict[int, Any]] = defaultdict(dict)
|
||||
for suite_name in suite_names:
|
||||
@@ -512,6 +521,10 @@ def create_libero_envs(
|
||||
cached_act_space = lazy.action_space
|
||||
cached_metadata = lazy.metadata
|
||||
out[suite_name][tid] = lazy
|
||||
elif is_sync:
|
||||
out[suite_name][tid] = gym.vector.SyncVectorEnv(
|
||||
fns, autoreset_mode=gym.vector.AutoresetMode.NEXT_STEP
|
||||
)
|
||||
else:
|
||||
out[suite_name][tid] = env_cls(fns)
|
||||
print(f"Built vec env | suite={suite_name} | task_id={tid} | n_envs={n_envs}")
|
||||
|
||||
@@ -177,6 +177,76 @@ def _sub_env_has_attr(env: gym.vector.VectorEnv, attr: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
# Passed in `reset(options=...)` by `rollout()` to mark the start of a new rollout.
|
||||
# FreezeAfterEpisodeEnd thaws only on this, so Gymnasium's argument-less autoreset
|
||||
# cannot be mistaken for a genuine new episode.
|
||||
NEW_ROLLOUT_OPTION = "lerobot_new_rollout"
|
||||
|
||||
|
||||
class FreezeAfterEpisodeEnd(gym.Wrapper):
|
||||
"""Stop doing simulator work once a sub-env's episode has ended.
|
||||
|
||||
`rollout()` runs `while not np.all(done)` with `done` latched, so a sub-env that
|
||||
terminates early keeps being stepped -- physics and offscreen rendering included --
|
||||
until the slowest sub-env in the batch finishes. The batch runs for
|
||||
`max(episode_lengths)` iterations to complete work that only needs
|
||||
`mean(episode_lengths)`.
|
||||
|
||||
This caches the terminal transition and replays it for any further `step()` or
|
||||
autoreset, so a finished sub-env costs nothing. The rollout already ignores those
|
||||
transitions.
|
||||
|
||||
The freeze survives Gymnasium's autoreset deliberately. Under
|
||||
`AutoresetMode.NEXT_STEP` the vector env resets a terminated sub-env on the
|
||||
following step and runs it through an entire extra episode that the rollout
|
||||
discards, because `done` stays latched. Absorbing that reset is most of the saving.
|
||||
|
||||
Only an explicit reset carrying `NEW_ROLLOUT_OPTION` thaws it, so the signal is
|
||||
explicit rather than inferred: Gymnasium's autoreset calls `reset()` with no
|
||||
arguments, but so would a caller passing `seeds=None`, and confusing the two would
|
||||
strand an env frozen for a whole rollout.
|
||||
|
||||
`AutoresetMode.DISABLED` is not an alternative here — Gymnasium asserts that no
|
||||
terminated env is ever stepped in that mode, so the wrapper is never reached.
|
||||
"""
|
||||
|
||||
def __init__(self, env: gym.Env):
|
||||
super().__init__(env)
|
||||
self._frozen: tuple | None = None
|
||||
|
||||
def reset(self, *, seed=None, options=None):
|
||||
if self._frozen is not None and not (options or {}).get(NEW_ROLLOUT_OPTION):
|
||||
# Gymnasium's autoreset for a sub-env the rollout has already finished with.
|
||||
# Replay the terminal observation instead of rebuilding the simulation.
|
||||
obs, _, _, _, info = self._frozen
|
||||
return obs, info
|
||||
self._frozen = None
|
||||
return self.env.reset(seed=seed, options=options)
|
||||
|
||||
def step(self, action):
|
||||
if self._frozen is not None:
|
||||
return self._frozen
|
||||
obs, reward, terminated, truncated, info = self.env.step(action)
|
||||
if terminated or truncated:
|
||||
# Zero the reward on replay so a frozen sub-env cannot inflate a return if a
|
||||
# caller sums rewards over the padded tail.
|
||||
self._frozen = (obs, 0.0, terminated, truncated, info)
|
||||
return obs, reward, terminated, truncated, info
|
||||
|
||||
@property
|
||||
def is_frozen(self) -> bool:
|
||||
return self._frozen is not None
|
||||
|
||||
|
||||
def freeze_after_episode_end(env_fn: Callable[[], gym.Env]) -> Callable[[], gym.Env]:
|
||||
"""Wrap an env factory so the built env freezes once its episode ends."""
|
||||
|
||||
def _fn() -> gym.Env:
|
||||
return FreezeAfterEpisodeEnd(env_fn())
|
||||
|
||||
return _fn
|
||||
|
||||
|
||||
class _LazyAsyncVectorEnv:
|
||||
"""Defers AsyncVectorEnv creation until first use.
|
||||
|
||||
@@ -212,7 +282,12 @@ class _LazyAsyncVectorEnv:
|
||||
|
||||
def _ensure(self) -> None:
|
||||
if self._env is None:
|
||||
self._env = gym.vector.AsyncVectorEnv(self._env_fns, context="forkserver", shared_memory=True)
|
||||
self._env = gym.vector.AsyncVectorEnv(
|
||||
[freeze_after_episode_end(fn) for fn in self._env_fns],
|
||||
context="forkserver",
|
||||
shared_memory=True,
|
||||
autoreset_mode=gym.vector.AutoresetMode.NEXT_STEP,
|
||||
)
|
||||
|
||||
@property
|
||||
def unwrapped(self):
|
||||
|
||||
@@ -432,7 +432,7 @@ def submit_to_hf(cfg: TrainPipelineConfig) -> None:
|
||||
|
||||
# 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
|
||||
# exact log line emitted by PreTrainedPolicy.push_model_to_hub — the two must stay
|
||||
# exact log line emitted by lerobot.common.train_utils.publish_trained_model — the two must stay
|
||||
# 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.
|
||||
success_marker = f"Model pushed to https://huggingface.co/{repo_id}"
|
||||
|
||||
@@ -314,11 +314,16 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
To find the port, you can run our utility script:
|
||||
```bash
|
||||
lerobot-find-port.py
|
||||
>>> Finding all available ports for the MotorsBus.
|
||||
>>> ["/dev/tty.usbmodem575E0032081", "/dev/tty.usbmodem575E0031751"]
|
||||
>>> Remove the usb cable from your MotorsBus and press Enter when done.
|
||||
>>> The port of this MotorsBus is /dev/tty.usbmodem575E0031751.
|
||||
>>> Reconnect the usb cable.
|
||||
```
|
||||
|
||||
which prints:
|
||||
|
||||
```
|
||||
Finding all available ports for the MotorsBus.
|
||||
["/dev/tty.usbmodem575E0032081", "/dev/tty.usbmodem575E0031751"]
|
||||
Remove the usb cable from your MotorsBus and press Enter when done.
|
||||
The port of this MotorsBus is /dev/tty.usbmodem575E0031751.
|
||||
Reconnect the usb cable.
|
||||
```
|
||||
|
||||
Example of usage for 1 Feetech sts3215 motor connected to the bus:
|
||||
@@ -595,7 +600,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
ID, and finally programs the bus' default baud-rate.
|
||||
|
||||
Args:
|
||||
motor (str): Key of the motor in :pyattr:`motors`.
|
||||
motor (str): Key of the motor in `motors`.
|
||||
initial_baudrate (int | None, optional): Current baud-rate (skips scanning when provided).
|
||||
Defaults to None.
|
||||
initial_id (int | None, optional): Current ID (skips scanning when provided). Defaults to None.
|
||||
@@ -666,7 +671,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
"""Enable torque on selected motors.
|
||||
|
||||
Args:
|
||||
motors (int | str | list[str] | None, optional): Same semantics as :pymeth:`disable_torque`.
|
||||
motors (int | str | list[str] | None, optional): Same semantics as [`~motors.motors_bus.MotorsBus.disable_torque`].
|
||||
Defaults to `None`.
|
||||
num_retry (int, optional): Number of additional retry attempts on communication failure.
|
||||
Defaults to 0.
|
||||
@@ -679,10 +684,12 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
This helper is useful to temporarily disable torque when configuring motors.
|
||||
|
||||
Examples:
|
||||
>>> with bus.torque_disabled():
|
||||
Example:
|
||||
```python
|
||||
>>> with bus.torque_disabled(): # doctest: +SKIP
|
||||
... # Safe operations here
|
||||
... pass
|
||||
```
|
||||
"""
|
||||
self.disable_torque(motors)
|
||||
try:
|
||||
@@ -695,7 +702,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
Args:
|
||||
timeout_ms (int | None, optional): Timeout in *milliseconds*. If `None` (default) the method falls
|
||||
back to :pyattr:`default_timeout`.
|
||||
back to `default_timeout`.
|
||||
"""
|
||||
timeout_ms = timeout_ms if timeout_ms is not None else self.default_timeout
|
||||
self.port_handler.setPacketTimeoutMillis(timeout_ms)
|
||||
@@ -746,8 +753,8 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
Args:
|
||||
calibration_dict (dict[str, MotorCalibration]): Calibration obtained from
|
||||
:pymeth:`read_calibration` or crafted by the user.
|
||||
cache (bool, optional): Save the calibration to :pyattr:`calibration`. Defaults to True.
|
||||
[`~motors.motors_bus.MotorsBus.read_calibration`] or crafted by the user.
|
||||
cache (bool, optional): Save the calibration to `calibration`. Defaults to True.
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -755,7 +762,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
"""Restore factory calibration for the selected motors.
|
||||
|
||||
Homing offset is set to ``0`` and min/max position limits are set to the full usable range.
|
||||
The in-memory :pyattr:`calibration` is cleared.
|
||||
The in-memory `calibration` is cleared.
|
||||
|
||||
Args:
|
||||
motors (NameOrID | Sequence[NameOrID] | None, optional): Selection of motors. `None` (default)
|
||||
@@ -1069,9 +1076,9 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
) -> None:
|
||||
"""Write a value to a single motor's register.
|
||||
|
||||
Contrary to :pymeth:`sync_write`, this expects a response status packet emitted by the motor, which
|
||||
Contrary to [`~motors.motors_bus.MotorsBus.sync_write`], this expects a response status packet emitted by the motor, which
|
||||
provides a guarantee that the value was written to the register successfully. In consequence, it is
|
||||
slower than :pymeth:`sync_write` but it is more reliable. It should typically be used when configuring
|
||||
slower than [`~motors.motors_bus.MotorsBus.sync_write`] but it is more reliable. It should typically be used when configuring
|
||||
motors.
|
||||
|
||||
Args:
|
||||
@@ -1228,8 +1235,8 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
) -> None:
|
||||
"""Write the same register on multiple motors.
|
||||
|
||||
Contrary to :pymeth:`write`, this *does not* expects a response status packet emitted by the motor, which
|
||||
can allow for lost packets. It is faster than :pymeth:`write` and should typically be used when
|
||||
Contrary to [`~motors.motors_bus.MotorsBus.write`], this *does not* expects a response status packet emitted by the motor, which
|
||||
can allow for lost packets. It is faster than [`~motors.motors_bus.MotorsBus.write`] and should typically be used when
|
||||
frequency matters and losing some packets is acceptable (e.g. teleoperation loops).
|
||||
|
||||
Args:
|
||||
|
||||
@@ -20,7 +20,6 @@ from .optimizers import (
|
||||
SGDConfig as SGDConfig,
|
||||
XVLAAdamWConfig as XVLAAdamWConfig,
|
||||
load_optimizer_state,
|
||||
load_optimizer_state_dict,
|
||||
save_optimizer_state,
|
||||
)
|
||||
from .schedulers import (
|
||||
@@ -51,7 +50,6 @@ __all__ = [
|
||||
"VQBeTSchedulerConfig",
|
||||
# State management
|
||||
"load_optimizer_state",
|
||||
"load_optimizer_state_dict",
|
||||
"load_scheduler_state",
|
||||
"save_optimizer_state",
|
||||
"save_scheduler_state",
|
||||
|
||||
@@ -27,7 +27,7 @@ from lerobot.utils.constants import (
|
||||
OPTIMIZER_PARAM_GROUPS,
|
||||
OPTIMIZER_STATE,
|
||||
)
|
||||
from lerobot.utils.io_utils import deserialize_json_into_object, load_json, write_json
|
||||
from lerobot.utils.io_utils import deserialize_json_into_object, write_json
|
||||
from lerobot.utils.utils import flatten_dict, unflatten_dict
|
||||
|
||||
# Type alias for parameters accepted by optimizer build() methods.
|
||||
@@ -52,6 +52,11 @@ class OptimizerConfig(draccus.ChoiceRegistry, abc.ABC):
|
||||
def type(self) -> str:
|
||||
return self.get_choice_name(self.__class__)
|
||||
|
||||
@property
|
||||
def builds_multiple_optimizers(self) -> bool:
|
||||
"""True when build() returns a dict of optimizers (unsupported under sharded training)."""
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def default_choice_name(cls) -> str | None:
|
||||
return "adam"
|
||||
@@ -245,6 +250,10 @@ class MultiAdamConfig(OptimizerConfig):
|
||||
grad_clip_norm: float = 10.0
|
||||
optimizer_groups: dict[str, dict[str, Any]] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def builds_multiple_optimizers(self) -> bool:
|
||||
return True
|
||||
|
||||
def build(self, params: OptimizerParams) -> dict[str, torch.optim.Optimizer]:
|
||||
"""Build multiple Adam optimizers.
|
||||
|
||||
@@ -283,35 +292,27 @@ class MultiAdamConfig(OptimizerConfig):
|
||||
def save_optimizer_state(
|
||||
optimizer: torch.optim.Optimizer | dict[str, torch.optim.Optimizer],
|
||||
save_dir: Path,
|
||||
optim_state_dict: dict | None = None,
|
||||
) -> None:
|
||||
"""Save optimizer state to disk.
|
||||
"""Save optimizer state to disk (non-sharded runs; sharded runs use the DCP channel).
|
||||
|
||||
Args:
|
||||
optimizer: Either a single optimizer or a dictionary of optimizers.
|
||||
save_dir: Directory to save the optimizer state.
|
||||
optim_state_dict: Pre-gathered optimizer state dict (for FSDP, where the sharded state must
|
||||
be gathered across ranks first). If provided, it is saved directly instead of calling
|
||||
``optimizer.state_dict()``. Only supported for a single optimizer. Defaults to None.
|
||||
"""
|
||||
if isinstance(optimizer, dict):
|
||||
# Handle dictionary of optimizers
|
||||
if optim_state_dict is not None:
|
||||
raise ValueError("optim_state_dict is not supported for a dict of optimizers")
|
||||
for name, opt in optimizer.items():
|
||||
optimizer_dir = save_dir / name
|
||||
optimizer_dir.mkdir(exist_ok=True, parents=True)
|
||||
_save_single_optimizer_state(opt, optimizer_dir)
|
||||
else:
|
||||
# Handle single optimizer
|
||||
_save_single_optimizer_state(optimizer, save_dir, optim_state_dict=optim_state_dict)
|
||||
_save_single_optimizer_state(optimizer, save_dir)
|
||||
|
||||
|
||||
def _save_single_optimizer_state(
|
||||
optimizer: torch.optim.Optimizer, save_dir: Path, optim_state_dict: dict | None = None
|
||||
) -> None:
|
||||
def _save_single_optimizer_state(optimizer: torch.optim.Optimizer, save_dir: Path) -> None:
|
||||
"""Save a single optimizer's state to disk."""
|
||||
state = dict(optim_state_dict) if optim_state_dict is not None else optimizer.state_dict()
|
||||
state = optimizer.state_dict()
|
||||
param_groups = state.pop("param_groups")
|
||||
flat_state = flatten_dict(state)
|
||||
save_file(flat_state, save_dir / OPTIMIZER_STATE)
|
||||
@@ -365,19 +366,3 @@ def _load_single_optimizer_state(optimizer: torch.optim.Optimizer, save_dir: Pat
|
||||
|
||||
optimizer.load_state_dict(loaded_state_dict)
|
||||
return optimizer
|
||||
|
||||
|
||||
def load_optimizer_state_dict(save_dir: Path) -> dict:
|
||||
"""Read a saved optimizer state dict (safetensors + json) back into a plain dict.
|
||||
|
||||
Unlike `load_optimizer_state`, this does not load into an optimizer and preserves the original
|
||||
``state`` keys verbatim (e.g. FSDP parameter FQNs, which are not integer-castable). It is used by
|
||||
the FSDP resume path, where the full state must be resharded via `FSDP.optim_state_dict_to_load`
|
||||
before being loaded into the (sharded) optimizer.
|
||||
"""
|
||||
flat_state = load_file(save_dir / OPTIMIZER_STATE)
|
||||
state = unflatten_dict(flat_state)
|
||||
return {
|
||||
"state": state.get("state", {}),
|
||||
"param_groups": load_json(save_dir / OPTIMIZER_PARAM_GROUPS),
|
||||
}
|
||||
|
||||
@@ -40,44 +40,93 @@ class ACTConfig(PreTrainedConfig):
|
||||
- "action" is required as an output key.
|
||||
|
||||
Args:
|
||||
n_obs_steps: Number of environment steps worth of observations to pass to the policy (takes the
|
||||
current step and additional steps going back).
|
||||
chunk_size: The size of the action prediction "chunks" in units of environment steps.
|
||||
n_action_steps: The number of action steps to run in the environment for one invocation of the policy.
|
||||
This should be no greater than the chunk size. For example, if the chunk size size 100, you may
|
||||
set this to 50. This would mean that the model predicts 100 steps worth of actions, runs 50 in the
|
||||
environment, and throws the other 50 out.
|
||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
||||
vision_backbone: Name of the torchvision resnet backbone to use for encoding images.
|
||||
pretrained_backbone_weights: Pretrained weights from torchvision to initialize the backbone.
|
||||
`None` means no pretrained weights.
|
||||
replace_final_stride_with_dilation: Whether to replace the ResNet's final 2x2 stride with a dilated
|
||||
convolution.
|
||||
pre_norm: Whether to use "pre-norm" in the transformer blocks.
|
||||
dim_model: The transformer blocks' main hidden dimension.
|
||||
n_heads: The number of heads to use in the transformer blocks' multi-head attention.
|
||||
dim_feedforward: The dimension to expand the transformer's hidden dimension to in the feed-forward
|
||||
layers.
|
||||
feedforward_activation: The activation to use in the transformer block's feed-forward layers.
|
||||
n_encoder_layers: The number of transformer layers to use for the transformer encoder.
|
||||
n_decoder_layers: The number of transformer layers to use for the transformer decoder.
|
||||
use_vae: Whether to use a variational objective during training. This introduces another transformer
|
||||
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||
Number of environment steps of observation to pass to the policy (the current step and
|
||||
additional steps going back). ACT only supports a value of 1; anything else raises in
|
||||
`__post_init__`.
|
||||
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
device (`str | None`, *optional*):
|
||||
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||
best available device if unset or unavailable.
|
||||
use_amp (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||
repo_id (`str | None`, *optional*):
|
||||
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||
private (`bool | None`, *optional*):
|
||||
Whether to create/push the Hub repository as private.
|
||||
tags (`list[str] | None`, *optional*):
|
||||
Tags to attach to the policy's Hub model card.
|
||||
license (`str | None`, *optional*):
|
||||
License identifier to add to the policy's Hub model card.
|
||||
pretrained_path (`Path | None`, *optional*):
|
||||
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||
is initialized from scratch.
|
||||
pretrained_revision (`str | None`, *optional*):
|
||||
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||
chunk_size (`int`, *optional*, defaults to 100):
|
||||
The size of the action prediction "chunks" in units of environment steps.
|
||||
n_action_steps (`int`, *optional*, defaults to 100):
|
||||
The number of action steps to run in the environment for one invocation of the policy. This
|
||||
should be no greater than `chunk_size`. For example, if the chunk size is 100, you may set this
|
||||
to 50: the model predicts 100 steps worth of actions, runs 50 in the environment, and throws
|
||||
the other 50 out.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||
it. Defaults to mean/std normalization for visual, state, and action features.
|
||||
vision_backbone (`str`, *optional*, defaults to `"resnet18"`):
|
||||
Name of the torchvision resnet backbone to use for encoding images.
|
||||
pretrained_backbone_weights (`str | None`, *optional*, defaults to `"ResNet18_Weights.IMAGENET1K_V1"`):
|
||||
Pretrained weights from torchvision to initialize the backbone. `None` means no pretrained
|
||||
weights.
|
||||
replace_final_stride_with_dilation (`int`, *optional*, defaults to `False`):
|
||||
Whether to replace the ResNet's final 2x2 stride with a dilated convolution.
|
||||
pre_norm (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use "pre-norm" in the transformer blocks.
|
||||
dim_model (`int`, *optional*, defaults to 512):
|
||||
The transformer blocks' main hidden dimension.
|
||||
n_heads (`int`, *optional*, defaults to 8):
|
||||
The number of heads to use in the transformer blocks' multi-head attention.
|
||||
dim_feedforward (`int`, *optional*, defaults to 3200):
|
||||
The dimension to expand the transformer's hidden dimension to in the feed-forward layers.
|
||||
feedforward_activation (`str`, *optional*, defaults to `"relu"`):
|
||||
The activation to use in the transformer block's feed-forward layers.
|
||||
n_encoder_layers (`int`, *optional*, defaults to 4):
|
||||
The number of transformer layers to use for the transformer encoder.
|
||||
n_decoder_layers (`int`, *optional*, defaults to 1):
|
||||
The number of transformer layers to use for the transformer decoder.
|
||||
use_vae (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use a variational objective during training. This introduces another transformer
|
||||
which is used as the VAE's encoder (not to be confused with the transformer encoder - see
|
||||
documentation in the policy class).
|
||||
latent_dim: The VAE's latent dimension.
|
||||
n_vae_encoder_layers: The number of transformer layers to use for the VAE's encoder.
|
||||
temporal_ensemble_coeff: Coefficient for the exponential weighting scheme to apply for temporal
|
||||
ensembling. Defaults to None which means temporal ensembling is not used. `n_action_steps` must be
|
||||
1 when using this feature, as inference needs to happen at every step to form an ensemble. For
|
||||
more information on how ensembling works, please see `ACTTemporalEnsembler`.
|
||||
dropout: Dropout to use in the transformer layers (see code for details).
|
||||
kl_weight: The weight to use for the KL-divergence component of the loss if the variational objective
|
||||
is enabled. Loss is then calculated as: `reconstruction_loss + kl_weight * kld_loss`.
|
||||
latent_dim (`int`, *optional*, defaults to 32):
|
||||
The VAE's latent dimension.
|
||||
n_vae_encoder_layers (`int`, *optional*, defaults to 4):
|
||||
The number of transformer layers to use for the VAE's encoder.
|
||||
temporal_ensemble_coeff (`float | None`, *optional*):
|
||||
Coefficient for the exponential weighting scheme to apply for temporal ensembling. `None` (the
|
||||
default) means temporal ensembling is not used. `n_action_steps` must be 1 when using this
|
||||
feature, as inference needs to happen at every step to form an ensemble. For more information
|
||||
on how ensembling works, see `ACTTemporalEnsembler`.
|
||||
dropout (`float`, *optional*, defaults to 0.1):
|
||||
Dropout to use in the transformer layers (see code for details).
|
||||
kl_weight (`float`, *optional*, defaults to 10.0):
|
||||
The weight to use for the KL-divergence component of the loss if the variational objective is
|
||||
enabled. Loss is then calculated as: `reconstruction_loss + kl_weight * kld_loss`.
|
||||
optimizer_lr (`float`, *optional*, defaults to 1e-05):
|
||||
Learning rate for the AdamW optimizer preset.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 0.0001):
|
||||
Weight decay for the AdamW optimizer preset.
|
||||
optimizer_lr_backbone (`float`, *optional*, defaults to 1e-05):
|
||||
Learning rate for the vision backbone's parameters in the AdamW optimizer preset.
|
||||
"""
|
||||
|
||||
# Input / output structure.
|
||||
@@ -128,9 +177,9 @@ class ACTConfig(PreTrainedConfig):
|
||||
optimizer_lr_backbone: float = 1e-5
|
||||
|
||||
def __post_init__(self):
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates `vision_backbone`, `temporal_ensemble_coeff`/`n_action_steps`, `n_action_steps`/`chunk_size`, and `n_obs_steps`."""
|
||||
super().__post_init__()
|
||||
|
||||
"""Input validation (not exhaustive)."""
|
||||
if not self.vision_backbone.startswith("resnet"):
|
||||
raise ValueError(
|
||||
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
||||
@@ -151,26 +200,32 @@ class ACTConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
return AdamWConfig(
|
||||
lr=self.optimizer_lr,
|
||||
weight_decay=self.optimizer_weight_decay,
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return None
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||
if not self.image_features and not self.env_state_feature:
|
||||
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return list(range(self.chunk_size))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -40,23 +40,25 @@ from .configuration_act import ACTConfig
|
||||
|
||||
|
||||
class ACTPolicy(PreTrainedPolicy):
|
||||
"""
|
||||
Action Chunking Transformer Policy as per Learning Fine-Grained Bimanual Manipulation with Low-Cost
|
||||
"""Action Chunking Transformer Policy as per Learning Fine-Grained Bimanual Manipulation with Low-Cost
|
||||
Hardware (paper: https://huggingface.co/papers/2304.13705, code: https://github.com/tonyzhaozh/act)
|
||||
"""
|
||||
|
||||
config_class = ACTConfig
|
||||
name = "act"
|
||||
# FSDP2 wrap units: one unit per transformer layer of both stacks.
|
||||
_fsdp_wrap_modules = ["ACTEncoderLayer", "ACTDecoderLayer"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: ACTConfig,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
"""Build the ACT model (and, if enabled, the temporal ensembler) from `config`.
|
||||
|
||||
Args:
|
||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
||||
the configuration class is used.
|
||||
config (`ACTConfig`):
|
||||
Policy configuration.
|
||||
"""
|
||||
super().__init__(config)
|
||||
config.validate_features()
|
||||
@@ -70,6 +72,11 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
self.reset()
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Splits parameters into two groups: the vision backbone, trained at `optimizer_lr_backbone`, and
|
||||
everything else, trained at the base `optimizer_lr`.
|
||||
"""
|
||||
# TODO(aliberts, rcadene): As of now, lr_backbone == lr
|
||||
# Should we remove this and just `return self.parameters()`?
|
||||
return [
|
||||
@@ -91,7 +98,11 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
]
|
||||
|
||||
def reset(self):
|
||||
"""This should be called whenever the environment is reset."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||
|
||||
Resets the `ACTTemporalEnsembler` when temporal ensembling is enabled, otherwise clears the action
|
||||
queue consumed by `select_action`.
|
||||
"""
|
||||
if self.config.temporal_ensemble_coeff is not None:
|
||||
self.temporal_ensembler.reset()
|
||||
else:
|
||||
@@ -99,11 +110,11 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Select a single action given environment observations.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
This method wraps `select_actions` in order to return one action at a time for execution in the
|
||||
environment. It works by managing the actions in a queue and only calling `select_actions` when the
|
||||
queue is empty.
|
||||
Returns one action at a time from a queue populated by `predict_action_chunk`, refilling it once
|
||||
it runs dry. When temporal ensembling is enabled, the queue is bypassed and the action is instead
|
||||
produced by combining chunks via `ACTTemporalEnsembler`.
|
||||
"""
|
||||
self.eval() # keeping the policy in eval mode as it could be set to train mode while queue is consumed
|
||||
|
||||
@@ -124,7 +135,7 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Predict a chunk of actions given environment observations."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]."""
|
||||
self.eval()
|
||||
|
||||
if self.config.image_features:
|
||||
@@ -135,7 +146,11 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
return actions
|
||||
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||
"""Run the batch through the model and compute the loss for training or validation."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
The loss is an L1 reconstruction loss between the predicted and target actions, plus (when
|
||||
`use_vae` is enabled) a KL-divergence term weighted by `kl_weight`.
|
||||
"""
|
||||
if self.config.image_features:
|
||||
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
||||
batch[OBS_IMAGES] = [batch[key] for key in self.config.image_features]
|
||||
@@ -219,8 +234,7 @@ class ACTTemporalEnsembler:
|
||||
self.ensembled_actions_count = None
|
||||
|
||||
def update(self, actions: Tensor) -> Tensor:
|
||||
"""
|
||||
Takes a (batch, chunk_size, action_dim) sequence of actions, update the temporal ensemble for all
|
||||
"""Takes a (batch, chunk_size, action_dim) sequence of actions, update the temporal ensemble for all
|
||||
time steps, and pop/return the next batch of actions in the sequence.
|
||||
"""
|
||||
self.ensemble_weights = self.ensemble_weights.to(device=actions.device)
|
||||
@@ -624,13 +638,13 @@ class ACTDecoderLayer(nn.Module):
|
||||
decoder_pos_embed: Tensor | None = None,
|
||||
encoder_pos_embed: Tensor | None = None,
|
||||
) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
x: (Decoder Sequence, Batch, Channel) tensor of input tokens.
|
||||
encoder_out: (Encoder Sequence, B, C) output features from the last layer of the encoder we are
|
||||
cross-attending with.
|
||||
encoder_pos_embed: (ES, 1, C) positional embedding for keys (from the encoder).
|
||||
decoder_pos_embed: (DS, 1, C) positional embedding for the queries (from the decoder).
|
||||
|
||||
Returns:
|
||||
(DS, B, C) tensor of decoder output features.
|
||||
"""
|
||||
@@ -669,9 +683,11 @@ def create_sinusoidal_pos_embedding(num_positions: int, dimension: int) -> Tenso
|
||||
"""1D sinusoidal positional embeddings as in Attention is All You Need.
|
||||
|
||||
Args:
|
||||
num_positions: Number of token positions required.
|
||||
Returns: (num_positions, dimension) position embeddings (the first dimension is the batch dimension).
|
||||
num_positions (`int`): Number of positions to embed (the sequence length).
|
||||
dimension (`int`): The embedding dimension.
|
||||
|
||||
Returns:
|
||||
`(num_positions, dimension)` position embeddings (the first dimension is the batch dimension).
|
||||
"""
|
||||
|
||||
def get_position_angle_vec(position):
|
||||
@@ -691,9 +707,8 @@ class ACTSinusoidalPositionEmbedding2d(nn.Module):
|
||||
"""
|
||||
|
||||
def __init__(self, dimension: int):
|
||||
"""
|
||||
Args:
|
||||
dimension: The desired dimension of the embeddings.
|
||||
"""Args:
|
||||
dimension: The desired dimension of the embeddings.
|
||||
"""
|
||||
super().__init__()
|
||||
self.dimension = dimension
|
||||
@@ -703,9 +718,9 @@ class ACTSinusoidalPositionEmbedding2d(nn.Module):
|
||||
self._temperature = 10000
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
x: A (B, C, H, W) batch of 2D feature map to generate the embeddings for.
|
||||
|
||||
Returns:
|
||||
A (1, C, H, W) batch of corresponding sinusoidal positional embeddings.
|
||||
"""
|
||||
|
||||
@@ -40,7 +40,7 @@ def make_act_pre_post_processors(
|
||||
|
||||
Args:
|
||||
config (ACTConfig): The ACT policy configuration object.
|
||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None): A dictionary containing dataset
|
||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): A dictionary containing dataset
|
||||
statistics (e.g., mean and std) used for normalization. Defaults to None.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -41,63 +41,135 @@ class DiffusionConfig(PreTrainedConfig):
|
||||
- "action" is required as an output key.
|
||||
|
||||
Args:
|
||||
n_obs_steps: Number of environment steps worth of observations to pass to the policy (takes the
|
||||
current step and additional steps going back).
|
||||
horizon: Diffusion model action prediction size as detailed in `DiffusionPolicy.select_action`.
|
||||
n_action_steps: The number of action steps to run in the environment for one invocation of the policy.
|
||||
See `DiffusionPolicy.select_action` for more details.
|
||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
||||
vision_backbone: Name of the torchvision resnet backbone to use for encoding images.
|
||||
resize_shape: (H, W) shape to resize images to as a preprocessing step for the vision
|
||||
backbone. If None, no resizing is done and the original image resolution is used.
|
||||
crop_ratio: Ratio in (0, 1] used to derive the crop size from resize_shape
|
||||
(crop_h = int(resize_shape[0] * crop_ratio), likewise for width).
|
||||
Set to 1.0 to disable cropping. Only takes effect when resize_shape is not None.
|
||||
crop_shape: (H, W) shape to crop images to. When resize_shape is set and crop_ratio < 1.0,
|
||||
this is computed automatically. Can also be set directly for legacy configs that use
|
||||
crop-only (without resize). If None and no derivation applies, no cropping is done.
|
||||
crop_is_random: Whether the crop should be random at training time (it's always a center
|
||||
crop in eval mode).
|
||||
pretrained_backbone_weights: Pretrained weights from torchvision to initialize the backbone.
|
||||
`None` means no pretrained weights.
|
||||
use_group_norm: Whether to replace batch normalization with group normalization in the backbone.
|
||||
The group sizes are set to be about 16 (to be precise, feature_dim // 16).
|
||||
spatial_softmax_num_keypoints: Number of keypoints for SpatialSoftmax.
|
||||
use_separate_rgb_encoder_per_camera: Whether to use a separate RGB encoder for each camera view.
|
||||
down_dims: Feature dimension for each stage of temporal downsampling in the diffusion modeling Unet.
|
||||
You may provide a variable number of dimensions, therefore also controlling the degree of
|
||||
n_obs_steps (`int`, *optional*, defaults to 2):
|
||||
Number of environment steps of observation to pass to the policy (the current step and
|
||||
additional steps going back).
|
||||
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
device (`str | None`, *optional*):
|
||||
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||
best available device if unset or unavailable.
|
||||
use_amp (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||
repo_id (`str | None`, *optional*):
|
||||
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||
private (`bool | None`, *optional*):
|
||||
Whether to create/push the Hub repository as private.
|
||||
tags (`list[str] | None`, *optional*):
|
||||
Tags to attach to the policy's Hub model card.
|
||||
license (`str | None`, *optional*):
|
||||
License identifier to add to the policy's Hub model card.
|
||||
pretrained_path (`Path | None`, *optional*):
|
||||
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||
is initialized from scratch.
|
||||
pretrained_revision (`str | None`, *optional*):
|
||||
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||
horizon (`int`, *optional*, defaults to 64):
|
||||
Diffusion model action prediction size as detailed in `DiffusionPolicy.select_action`.
|
||||
n_action_steps (`int`, *optional*, defaults to 32):
|
||||
The number of action steps to run in the environment for one invocation of the policy. See
|
||||
`DiffusionPolicy.select_action` for more details.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||
it. Defaults to mean/std normalization for visual features and min/max normalization for
|
||||
state and action features.
|
||||
drop_n_last_frames (`int`, *optional*, defaults to 7):
|
||||
Number of frames dropped from the end of each episode when sampling training windows, which
|
||||
avoids excessive padding. Should track `horizon - n_action_steps - n_obs_steps + 1`.
|
||||
vision_backbone (`str`, *optional*, defaults to `"resnet18"`):
|
||||
Name of the torchvision resnet backbone to use for encoding images.
|
||||
resize_shape (`tuple[int, int] | None`, *optional*):
|
||||
(H, W) shape to resize images to as a preprocessing step for the vision backbone. `None`
|
||||
disables resizing, so the original image resolution is used.
|
||||
crop_ratio (`float`, *optional*, defaults to 1.0):
|
||||
Ratio in (0, 1] used to derive the crop size from `resize_shape` (`crop_h =
|
||||
int(resize_shape[0] * crop_ratio)`, likewise for width). Set to 1.0 to disable cropping. Only
|
||||
takes effect when `resize_shape` is not `None`.
|
||||
crop_shape (`tuple[int, int] | None`, *optional*):
|
||||
(H, W) shape to crop images to. Computed automatically when `resize_shape` is set and
|
||||
`crop_ratio` < 1.0. Can also be set directly for legacy configs that use crop-only (without
|
||||
resize). `None`, with no derivation applying, means no cropping.
|
||||
crop_is_random (`bool`, *optional*, defaults to `True`):
|
||||
Whether the crop should be random at training time (it's always a center crop in eval mode).
|
||||
pretrained_backbone_weights (`str | None`, *optional*, defaults to `"ResNet18_Weights.IMAGENET1K_V1"`):
|
||||
Pretrained weights from torchvision to initialize the backbone. `None` means no pretrained
|
||||
weights.
|
||||
use_group_norm (`bool`, *optional*, defaults to `False`):
|
||||
Whether to replace batch normalization with group normalization in the backbone. The group
|
||||
sizes are set to be about 16 (`feature_dim // 16`).
|
||||
spatial_softmax_num_keypoints (`int`, *optional*, defaults to 32):
|
||||
Number of keypoints for SpatialSoftmax.
|
||||
use_separate_rgb_encoder_per_camera (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use a separate RGB encoder for each camera view.
|
||||
down_dims (`tuple[int, ...]`, *optional*, defaults to `(512, 1024, 2048)`):
|
||||
Feature dimension for each stage of temporal downsampling in the diffusion modeling Unet. You
|
||||
may provide a variable number of dimensions, therefore also controlling the degree of
|
||||
downsampling.
|
||||
kernel_size: The convolutional kernel size of the diffusion modeling Unet.
|
||||
n_groups: Number of groups used in the group norm of the Unet's convolutional blocks.
|
||||
diffusion_step_embed_dim: The Unet is conditioned on the diffusion timestep via a small non-linear
|
||||
network. This is the output dimension of that network, i.e., the embedding dimension.
|
||||
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
|
||||
kernel_size (`int`, *optional*, defaults to 5):
|
||||
The convolutional kernel size of the diffusion modeling Unet.
|
||||
n_groups (`int`, *optional*, defaults to 8):
|
||||
Number of groups used in the group norm of the Unet's convolutional blocks.
|
||||
diffusion_step_embed_dim (`int`, *optional*, defaults to 128):
|
||||
The Unet is conditioned on the diffusion timestep via a small non-linear network. This is the
|
||||
output dimension of that network, i.e. the embedding dimension.
|
||||
use_film_scale_modulation (`bool`, *optional*, defaults to `True`):
|
||||
FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning. Bias
|
||||
modulation is used by default, while this parameter indicates whether to also use scale
|
||||
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"].
|
||||
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_start: Beta value for the first forward-diffusion step.
|
||||
beta_end: Beta value for the last forward-diffusion step.
|
||||
prediction_type: The type of prediction that the diffusion modeling Unet makes. Choose from "epsilon"
|
||||
or "sample". These have equivalent outcomes from a latent variable modeling perspective, but
|
||||
"epsilon" has been shown to work better in many deep neural network settings.
|
||||
clip_sample: Whether to clip the sample to [-`clip_sample_range`, +`clip_sample_range`] for each
|
||||
denoising step at inference time. WARNING: you will need to make sure your action-space is
|
||||
normalized to fit within this range.
|
||||
clip_sample_range: The magnitude of the clipping range as described above.
|
||||
num_inference_steps: Number of reverse diffusion steps to use at inference time (steps are evenly
|
||||
spaced). If not provided, this defaults to be the same as `num_train_timesteps`.
|
||||
do_mask_loss_for_padding: Whether to mask the loss when there are copy-padded actions. See
|
||||
`LeRobotDataset` and `load_previous_and_future_frames` for more information. Note, this defaults
|
||||
to False as the original Diffusion Policy implementation does the same.
|
||||
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||
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 (`str`, *optional*, defaults to `"DDPM"`):
|
||||
Name of the noise scheduler to use. Supported options: `"DDPM"`, `"DDIM"`.
|
||||
num_train_timesteps (`int`, *optional*, defaults to 100):
|
||||
Number of diffusion steps for the forward diffusion schedule.
|
||||
beta_schedule (`str`, *optional*, defaults to `"squaredcos_cap_v2"`):
|
||||
Name of the diffusion beta schedule as per `DDPMScheduler` from Hugging Face diffusers.
|
||||
beta_start (`float`, *optional*, defaults to 0.0001):
|
||||
Beta value for the first forward-diffusion step.
|
||||
beta_end (`float`, *optional*, defaults to 0.02):
|
||||
Beta value for the last forward-diffusion step.
|
||||
prediction_type (`str`, *optional*, defaults to `"epsilon"`):
|
||||
The type of prediction that the diffusion modeling Unet makes. Choose from `"epsilon"` or
|
||||
`"sample"`. These have equivalent outcomes from a latent variable modeling perspective, but
|
||||
`"epsilon"` has been shown to work better in many deep neural network settings.
|
||||
clip_sample (`bool`, *optional*, defaults to `True`):
|
||||
Whether to clip the sample to `[-clip_sample_range, +clip_sample_range]` for each denoising
|
||||
step at inference time. This requires the action space to be normalized to fit within that
|
||||
range.
|
||||
clip_sample_range (`float`, *optional*, defaults to 1.0):
|
||||
The magnitude of the clipping range described above.
|
||||
num_inference_steps (`int | None`, *optional*):
|
||||
Number of reverse diffusion steps to use at inference time (steps are evenly spaced). If not
|
||||
provided, defaults to the same value as `num_train_timesteps`.
|
||||
compile_model (`bool`, *optional*, defaults to `False`):
|
||||
Whether to compile the Unet with `torch.compile`.
|
||||
compile_mode (`str`, *optional*, defaults to `"reduce-overhead"`):
|
||||
`torch.compile` mode to use when `compile_model` is enabled.
|
||||
do_mask_loss_for_padding (`bool`, *optional*, defaults to `False`):
|
||||
Whether to mask the loss when there are copy-padded actions. See `LeRobotDataset` and
|
||||
`load_previous_and_future_frames` for more information. This defaults to `False` as the
|
||||
original Diffusion Policy implementation does the same.
|
||||
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||
Learning rate for the Adam optimizer preset.
|
||||
optimizer_betas (`tuple`, *optional*, defaults to `(0.95, 0.999)`):
|
||||
Adam optimizer's beta coefficients.
|
||||
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||
Adam optimizer's epsilon for numerical stability.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 1e-06):
|
||||
Weight decay for the Adam optimizer preset.
|
||||
scheduler_name (`str`, *optional*, defaults to `"cosine"`):
|
||||
Name of the LR scheduler preset to use.
|
||||
scheduler_warmup_steps (`int`, *optional*, defaults to 500):
|
||||
Number of warmup steps for the LR scheduler preset.
|
||||
"""
|
||||
|
||||
# Inputs / output structure.
|
||||
@@ -164,9 +236,9 @@ class DiffusionConfig(PreTrainedConfig):
|
||||
scheduler_warmup_steps: int = 500
|
||||
|
||||
def __post_init__(self):
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates image/state feature presence and normalization-mode compatibility with the configured vision backbone."""
|
||||
super().__post_init__()
|
||||
|
||||
"""Input validation (not exhaustive)."""
|
||||
if not self.vision_backbone.startswith("resnet"):
|
||||
raise ValueError(
|
||||
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
||||
@@ -213,6 +285,7 @@ class DiffusionConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_optimizer_preset(self) -> AdamConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
return AdamConfig(
|
||||
lr=self.optimizer_lr,
|
||||
betas=self.optimizer_betas,
|
||||
@@ -221,12 +294,14 @@ class DiffusionConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self) -> DiffuserSchedulerConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return DiffuserSchedulerConfig(
|
||||
name=self.scheduler_name,
|
||||
num_warmup_steps=self.scheduler_warmup_steps,
|
||||
)
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||
if len(self.image_features) == 0 and self.env_state_feature is None:
|
||||
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
||||
|
||||
@@ -249,12 +324,15 @@ class DiffusionConfig(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
return list(range(1 - self.n_obs_steps, 1))
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return list(range(1 - self.n_obs_steps, 1 - self.n_obs_steps + self.horizon))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -54,8 +54,7 @@ from .configuration_diffusion import DiffusionConfig
|
||||
|
||||
|
||||
class DiffusionPolicy(PreTrainedPolicy):
|
||||
"""
|
||||
Diffusion Policy as per "Diffusion Policy: Visuomotor Policy Learning via Action Diffusion"
|
||||
"""Diffusion Policy as per "Diffusion Policy: Visuomotor Policy Learning via Action Diffusion"
|
||||
(paper: https://huggingface.co/papers/2303.04137, code: https://github.com/real-stanford/diffusion_policy).
|
||||
"""
|
||||
|
||||
@@ -67,12 +66,11 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
config: DiffusionConfig,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
"""Build the diffusion model from `config`.
|
||||
|
||||
Args:
|
||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
||||
the configuration class is used.
|
||||
dataset_stats: Dataset statistics to be used for normalization. If not passed here, it is expected
|
||||
that they will be passed with a call to `load_state_dict` before the policy is used.
|
||||
config (`DiffusionConfig`):
|
||||
Policy configuration.
|
||||
"""
|
||||
require_package("diffusers", extra="diffusion")
|
||||
super().__init__(config)
|
||||
@@ -87,10 +85,14 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
self.reset()
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||
return self.diffusion.parameters()
|
||||
|
||||
def reset(self):
|
||||
"""Clear observation and action queues. Should be called on `env.reset()`"""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||
|
||||
Clears the observation and action queues consumed by `select_action`.
|
||||
"""
|
||||
self._queues = {
|
||||
OBS_STATE: deque(maxlen=self.config.n_obs_steps),
|
||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||
@@ -102,7 +104,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||
"""Predict a chunk of actions given environment observations.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Supports two modes:
|
||||
- Online (queues populated via select_action): stacks observations from internal queues.
|
||||
@@ -123,7 +125,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||
"""Select a single action given environment observations.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
This method handles caching a history of observations and an action trajectory generated by the
|
||||
underlying diffusion model. Here's how it works:
|
||||
@@ -161,7 +163,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
return action
|
||||
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, None]:
|
||||
"""Run the batch through the model and compute the loss for training or validation."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]."""
|
||||
if self.config.image_features:
|
||||
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
||||
for key in self.config.image_features:
|
||||
@@ -174,8 +176,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
||||
|
||||
|
||||
def _make_noise_scheduler(name: str, **kwargs: dict):
|
||||
"""
|
||||
Factory for noise scheduler instances of the requested type. All kwargs are passed
|
||||
"""Factory for noise scheduler instances of the requested type. All kwargs are passed
|
||||
to the scheduler.
|
||||
"""
|
||||
require_package("diffusers", extra="diffusion")
|
||||
@@ -306,8 +307,7 @@ class DiffusionModel(nn.Module):
|
||||
return torch.cat(global_cond_feats, dim=-1).flatten(start_dim=1)
|
||||
|
||||
def generate_actions(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||
"""
|
||||
This function expects `batch` to have:
|
||||
"""This function expects `batch` to have:
|
||||
{
|
||||
"observation.state": (B, n_obs_steps, state_dim)
|
||||
|
||||
@@ -333,8 +333,7 @@ class DiffusionModel(nn.Module):
|
||||
return actions
|
||||
|
||||
def compute_loss(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""
|
||||
This function expects `batch` to have (at least):
|
||||
"""This function expects `batch` to have (at least):
|
||||
{
|
||||
"observation.state": (B, n_obs_steps, state_dim)
|
||||
|
||||
@@ -401,8 +400,7 @@ class DiffusionModel(nn.Module):
|
||||
|
||||
|
||||
class SpatialSoftmax(nn.Module):
|
||||
"""
|
||||
Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
||||
"""Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
||||
(https://huggingface.co/papers/1509.06113). A minimal port of the robomimic implementation.
|
||||
|
||||
At a high level, this takes 2D feature maps (from a convnet/ViT) and returns the "center of mass"
|
||||
@@ -424,10 +422,9 @@ class SpatialSoftmax(nn.Module):
|
||||
"""
|
||||
|
||||
def __init__(self, input_shape, num_kp=None):
|
||||
"""
|
||||
Args:
|
||||
input_shape (list): (C, H, W) input feature map shape.
|
||||
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
||||
"""Args:
|
||||
input_shape (list): (C, H, W) input feature map shape.
|
||||
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
@@ -450,9 +447,9 @@ class SpatialSoftmax(nn.Module):
|
||||
self.register_buffer("pos_grid", torch.cat([pos_x, pos_y], dim=1))
|
||||
|
||||
def forward(self, features: Tensor) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
features: (B, C, H, W) input feature maps.
|
||||
|
||||
Returns:
|
||||
(B, K, 2) image-space coordinates of keypoints.
|
||||
"""
|
||||
@@ -536,9 +533,9 @@ class DiffusionRgbEncoder(nn.Module):
|
||||
self.relu = nn.ReLU()
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
x: (B, C, H, W) image tensor with pixel values in [0, 1].
|
||||
|
||||
Returns:
|
||||
(B, D) image feature.
|
||||
"""
|
||||
@@ -562,11 +559,11 @@ class DiffusionRgbEncoder(nn.Module):
|
||||
def _replace_submodules(
|
||||
root_module: nn.Module, predicate: Callable[[nn.Module], bool], func: Callable[[nn.Module], nn.Module]
|
||||
) -> nn.Module:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
root_module: The module for which the submodules need to be replaced
|
||||
predicate: Takes a module as an argument and must return True if the that module is to be replaced.
|
||||
func: Takes a module as an argument and returns a new module to replace it with.
|
||||
|
||||
Returns:
|
||||
The root module with its submodules replaced.
|
||||
"""
|
||||
@@ -708,12 +705,12 @@ class DiffusionConditionalUnet1d(nn.Module):
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor, timestep: Tensor | int, global_cond=None) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
x: (B, T, input_dim) tensor for input to the Unet.
|
||||
timestep: (B,) tensor of (timestep_we_are_denoising_from - 1).
|
||||
global_cond: (B, global_cond_dim)
|
||||
output: (B, T, input_dim)
|
||||
|
||||
Returns:
|
||||
(B, T, input_dim) diffusion model prediction.
|
||||
"""
|
||||
@@ -798,10 +795,10 @@ class DiffusionConditionalResidualBlock1d(nn.Module):
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor, cond: Tensor) -> Tensor:
|
||||
"""
|
||||
Args:
|
||||
"""Args:
|
||||
x: (B, in_channels, T)
|
||||
cond: (B, cond_dim)
|
||||
|
||||
Returns:
|
||||
(B, out_channels, T)
|
||||
"""
|
||||
|
||||
@@ -34,8 +34,7 @@ def make_diffusion_pre_post_processors(
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""
|
||||
Constructs pre-processor and post-processor pipelines for a diffusion policy.
|
||||
"""Constructs pre-processor and post-processor pipelines for a diffusion policy.
|
||||
|
||||
The pre-processing pipeline prepares the input data for the model by:
|
||||
1. Renaming features.
|
||||
@@ -48,10 +47,8 @@ def make_diffusion_pre_post_processors(
|
||||
2. Unnormalizing the output features to their original scale.
|
||||
|
||||
Args:
|
||||
config: The configuration object for the diffusion policy,
|
||||
containing feature definitions, normalization mappings, and device information.
|
||||
dataset_stats: A dictionary of statistics used for normalization.
|
||||
Defaults to None.
|
||||
config (`DiffusionConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
|
||||
@@ -42,7 +42,117 @@ else:
|
||||
@PreTrainedConfig.register_subclass("eo1")
|
||||
@dataclass
|
||||
class EO1Config(PreTrainedConfig):
|
||||
"""Configuration for native EO1 policy integration in LeRobot."""
|
||||
"""Configuration for native EO1 policy integration in LeRobot.
|
||||
|
||||
EO1 wraps a Qwen2.5-VL vision-language backbone with a flow-matching action head: the backbone attends
|
||||
over interleaved vision/language/state/action tokens, and the head denoises an action chunk from noise
|
||||
via Euler integration.
|
||||
|
||||
Args:
|
||||
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||
Number of environment steps of observation to pass to the policy.
|
||||
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Input feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Output feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||
device (`str`, *optional*):
|
||||
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||
unavailable.
|
||||
use_amp (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub.
|
||||
repo_id (`str`, *optional*):
|
||||
Hub repository id to push the policy to.
|
||||
private (`bool`, *optional*):
|
||||
Whether the pushed Hub repository is private.
|
||||
tags (`list[str]`, *optional*):
|
||||
Tags to attach to the policy on the Hub.
|
||||
license (`str`, *optional*):
|
||||
License identifier for the policy on the Hub.
|
||||
pretrained_path (`Path`, *optional*):
|
||||
Repo id or local directory of pretrained weights saved with `save_pretrained`. Left unset to
|
||||
initialize from scratch.
|
||||
pretrained_revision (`str`, *optional*):
|
||||
Hub revision to pin when loading `pretrained_path`.
|
||||
vlm_base (`str`, *optional*, defaults to `"Qwen/Qwen2.5-VL-3B-Instruct"`):
|
||||
Hugging Face Hub id of the Qwen2.5-VL backbone used to initialize the vision-language model.
|
||||
vlm_config (`dict`, *optional*):
|
||||
Serialized Qwen2.5-VL backbone config. Populated automatically from `vlm_base` in
|
||||
`__post_init__` when left unset.
|
||||
image_min_pixels (`int`, *optional*, defaults to 50176):
|
||||
Minimum number of pixels the vision processor resizes an image down to.
|
||||
image_max_pixels (`int`, *optional*, defaults to 100352):
|
||||
Maximum number of pixels the vision processor resizes an image up to.
|
||||
use_fast_processor (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use the Hugging Face "fast" image processor.
|
||||
chunk_size (`int`, *optional*, defaults to 8):
|
||||
Number of actions predicted per flow-matching sampling call.
|
||||
n_action_steps (`int`, *optional*, defaults to 8):
|
||||
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||
policy. Must not exceed `chunk_size`.
|
||||
max_state_dim (`int`, *optional*, defaults to 32):
|
||||
Padded dimensionality of the state vector fed to the flow-matching head.
|
||||
max_action_dim (`int`, *optional*, defaults to 32):
|
||||
Padded dimensionality of the action vector fed to the flow-matching head.
|
||||
num_denoise_steps (`int`, *optional*, defaults to 10):
|
||||
Number of Euler integration steps used to sample an action chunk.
|
||||
num_action_layers (`int`, *optional*, defaults to 2):
|
||||
Number of linear layers in the action output projector MLP.
|
||||
action_act (`str`, *optional*, defaults to `"linear"`):
|
||||
Activation used between the action output projector's layers.
|
||||
time_sampling_beta_alpha (`float`, *optional*, defaults to 1.5):
|
||||
Alpha parameter of the Beta distribution used to sample the flow-matching timestep during
|
||||
training.
|
||||
time_sampling_beta_beta (`float`, *optional*, defaults to 1.0):
|
||||
Beta parameter of the same Beta distribution.
|
||||
time_sampling_scale (`float`, *optional*, defaults to 0.999):
|
||||
Scale applied to the sampled Beta timestep.
|
||||
time_sampling_offset (`float`, *optional*, defaults to 0.001):
|
||||
Offset added to the scaled Beta timestep.
|
||||
min_period (`float`, *optional*, defaults to 0.004):
|
||||
Minimum period of the sinusoidal timestep embedding.
|
||||
max_period (`float`, *optional*, defaults to 4.0):
|
||||
Maximum period of the sinusoidal timestep embedding.
|
||||
supervise_padding_action_dims (`bool`, *optional*, defaults to `True`):
|
||||
Whether the flow-matching loss also supervises the padded action dimensions that lie beyond
|
||||
the dataset's real action size.
|
||||
supervise_padding_actions (`bool`, *optional*, defaults to `True`):
|
||||
Whether the flow-matching loss also supervises padded action timesteps. Padded timesteps are
|
||||
marked by `action_is_pad`.
|
||||
dtype (`str`, *optional*, defaults to `"auto"`):
|
||||
Dtype requested for the Qwen backbone. `"auto"` follows the backbone checkpoint's default
|
||||
dtype (bf16 for Qwen2.5-VL); the flow-matching head always keeps its own parameters in fp32
|
||||
regardless. Other supported values are `"bfloat16"` and `"float32"`.
|
||||
force_fp32_autocast (`bool`, *optional*, defaults to `True`):
|
||||
Whether to disable autocast around the flow-matching head so its projections run in fp32 even
|
||||
when the backbone runs under bf16 autocast.
|
||||
attn_implementation (`str`, *optional*):
|
||||
Attention backend requested for the Qwen backbone, e.g. `"sdpa"` or `"flash_attention_2"`.
|
||||
Left unset to use the backbone's default.
|
||||
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||
Whether to enable gradient checkpointing on the Qwen backbone to reduce memory usage.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||
Peak learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`):
|
||||
Adam beta coefficients for the default optimizer preset.
|
||||
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||
Adam epsilon for the default optimizer preset.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 0.1):
|
||||
Weight decay for the default optimizer preset.
|
||||
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||
Gradient-norm clipping threshold for the default optimizer preset.
|
||||
scheduler_warmup_steps (`int`, *optional*, defaults to 900):
|
||||
Number of warmup steps for the default cosine-decay-with-warmup scheduler preset.
|
||||
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||
Number of decay steps for the default scheduler preset.
|
||||
scheduler_decay_lr (`float`, *optional*, defaults to 0.0):
|
||||
Learning rate reached at the end of the default scheduler's decay.
|
||||
"""
|
||||
|
||||
vlm_base: str = "Qwen/Qwen2.5-VL-3B-Instruct"
|
||||
vlm_config: dict | None = None
|
||||
@@ -112,6 +222,7 @@ class EO1Config(PreTrainedConfig):
|
||||
scheduler_decay_lr: float = 0.0
|
||||
|
||||
def __post_init__(self):
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the VLM backbone/tokenizer configuration."""
|
||||
super().__post_init__()
|
||||
|
||||
if self.n_action_steps > self.chunk_size:
|
||||
@@ -126,6 +237,7 @@ class EO1Config(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def vlm_backbone_config(self) -> Qwen2_5_VLConfig:
|
||||
"""Build the Qwen2.5-VL backbone config from `vlm_config`, applying `attn_implementation` if set."""
|
||||
require_package("transformers", extra="eo1")
|
||||
config_dict = deepcopy(self.vlm_config)
|
||||
if self.attn_implementation is not None:
|
||||
@@ -134,10 +246,12 @@ class EO1Config(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def text_config(self) -> Qwen2_5_VLTextConfig:
|
||||
"""The text-tower sub-config of `vlm_backbone_config`."""
|
||||
return self.vlm_backbone_config.text_config
|
||||
|
||||
@property
|
||||
def vision_config(self) -> Qwen2_5_VLVisionConfig:
|
||||
"""The vision-tower sub-config of `vlm_backbone_config`."""
|
||||
return self.vlm_backbone_config.vision_config
|
||||
|
||||
def validate_features(self) -> None:
|
||||
@@ -164,6 +278,7 @@ class EO1Config(PreTrainedConfig):
|
||||
self.output_features[ACTION] = action_feature
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
return AdamWConfig(
|
||||
lr=self.optimizer_lr,
|
||||
betas=self.optimizer_betas,
|
||||
@@ -173,6 +288,7 @@ class EO1Config(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self):
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return CosineDecayWithWarmupSchedulerConfig(
|
||||
peak_lr=self.optimizer_lr,
|
||||
decay_lr=self.scheduler_decay_lr,
|
||||
@@ -182,12 +298,15 @@ class EO1Config(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list[int]:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return list(range(self.chunk_size))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -54,6 +54,14 @@ class EO1Policy(PreTrainedPolicy):
|
||||
name = "eo1"
|
||||
|
||||
def __init__(self, config: EO1Config, **kwargs):
|
||||
"""Build the Qwen2.5-VL backbone and the flow-matching action head.
|
||||
|
||||
Args:
|
||||
config (`EO1Config`):
|
||||
Policy configuration. Also drives whether the Qwen backbone is loaded from
|
||||
`config.vlm_base` (fresh initialization) or reconstructed from `config.vlm_backbone_config`
|
||||
(resuming from `config.pretrained_path`).
|
||||
"""
|
||||
require_package("transformers", extra="eo1")
|
||||
super().__init__(config)
|
||||
config.validate_features()
|
||||
@@ -80,6 +88,7 @@ class EO1Policy(PreTrainedPolicy):
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||
|
||||
@staticmethod
|
||||
@@ -87,6 +96,11 @@ class EO1Policy(PreTrainedPolicy):
|
||||
return {key: value for key, value in batch.items() if key not in excluded_keys}
|
||||
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
Computes the flow-matching loss: the mean squared error between the noise-minus-action target and
|
||||
the velocity predicted by the Qwen backbone plus flow head at a sampled timestep.
|
||||
"""
|
||||
state = self.prepare_state(batch[OBS_STATE])
|
||||
actions = self.prepare_action(batch[ACTION])
|
||||
model_inputs = self._get_model_inputs(batch, {OBS_STATE, ACTION})
|
||||
@@ -97,6 +111,11 @@ class EO1Policy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Samples the chunk by Euler-integrating the flow-matching head from noise, then slices it back down
|
||||
to the dataset's real action dimensionality (undoing the `max_action_dim` padding).
|
||||
"""
|
||||
self.eval()
|
||||
|
||||
states = self.prepare_state(batch[OBS_STATE])
|
||||
@@ -107,13 +126,16 @@ class EO1Policy(PreTrainedPolicy):
|
||||
return actions[:, :, :original_action_dim]
|
||||
|
||||
def prepare_state(self, state: Tensor) -> Tensor:
|
||||
"""Zero-pad a state tensor up to `config.max_state_dim` for the flow-matching head."""
|
||||
return pad_vector(state, self.config.max_state_dim)
|
||||
|
||||
def prepare_action(self, action: Tensor) -> Tensor:
|
||||
"""Zero-pad an action tensor up to `config.max_action_dim` for the flow-matching head."""
|
||||
return pad_vector(action, self.config.max_action_dim)
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Uses an action queue populated by `predict_action_chunk`."""
|
||||
self.eval()
|
||||
|
||||
if len(self._action_queue) == 0:
|
||||
@@ -123,6 +145,7 @@ class EO1Policy(PreTrainedPolicy):
|
||||
return self._action_queue.popleft()
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]. Trains every policy parameter with a single learning rate."""
|
||||
return self.parameters()
|
||||
|
||||
|
||||
@@ -358,7 +381,6 @@ class EO1VisionFlowMatchingModel(nn.Module):
|
||||
**kwargs,
|
||||
) -> Tensor:
|
||||
"""Run the EO1 training forward pass and compute the flow-matching loss."""
|
||||
|
||||
# 1. Build the EO1 prefix with state placeholders resolved.
|
||||
inputs_embeds = self.embed_prefix(
|
||||
input_ids,
|
||||
|
||||
@@ -31,6 +31,155 @@ logger = logging.getLogger(__name__)
|
||||
@PreTrainedConfig.register_subclass("evo1")
|
||||
@dataclass
|
||||
class Evo1Config(PreTrainedConfig):
|
||||
"""Configuration for the EVO1 vision-language-action policy.
|
||||
|
||||
EVO1 pairs an InternVL3 vision-language backbone with a flow-matching action head. Training proceeds
|
||||
in two stages (`training_stage`): stage 1 freezes the VLM and trains only the action head, stage 2
|
||||
fine-tunes the whole model.
|
||||
|
||||
Args:
|
||||
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||
Number of environment steps of observation to pass to the policy.
|
||||
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Input feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Output feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||
device (`str`, *optional*):
|
||||
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||
unavailable.
|
||||
use_amp (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use Automatic Mixed Precision. EVO1 also manages its own bfloat16 autocast around
|
||||
its forward passes independently of this flag; see `dtype`-related fields below.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub.
|
||||
repo_id (`str`, *optional*):
|
||||
Hub repository id to push the policy to.
|
||||
private (`bool`, *optional*):
|
||||
Whether the pushed Hub repository is private.
|
||||
tags (`list[str]`, *optional*):
|
||||
Tags to attach to the policy on the Hub.
|
||||
license (`str`, *optional*):
|
||||
License identifier for the policy on the Hub.
|
||||
pretrained_path (`Path`, *optional*):
|
||||
Repo id or local directory of pretrained weights saved with `save_pretrained`. Left unset to
|
||||
initialize from scratch.
|
||||
pretrained_revision (`str`, *optional*):
|
||||
Hub revision to pin when loading `pretrained_path`.
|
||||
training_stage (`str`, *optional*, defaults to `"stage1"`):
|
||||
Either `"stage1"` (VLM frozen, only the action head trains) or `"stage2"` (the whole model
|
||||
trains). Drives the default `finetune_*` flags unless they are set explicitly and
|
||||
`apply_training_stage_defaults` is `False`.
|
||||
chunk_size (`int`, *optional*, defaults to 50):
|
||||
Number of actions predicted by the flow-matching head per inference call.
|
||||
n_action_steps (`int`, *optional*, defaults to 50):
|
||||
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||
policy. Must not exceed `chunk_size`.
|
||||
max_state_dim (`int`, *optional*, defaults to 24):
|
||||
Padded dimensionality of the state vector fed to the action head.
|
||||
max_action_dim (`int`, *optional*, defaults to 24):
|
||||
Padded dimensionality of the action vector fed to the action head.
|
||||
max_views (`int`, *optional*, defaults to 3):
|
||||
Maximum number of camera streams the policy accepts.
|
||||
image_resolution (`tuple[int, int]`, *optional*, defaults to `(448, 448)`):
|
||||
Target resolution images are resized to before the InternVL3 embedder. Must be square.
|
||||
empty_cameras (`int`, *optional*, defaults to 0):
|
||||
Number of placeholder, always-masked-out camera views added to `input_features` so the batch
|
||||
has a fixed number of views regardless of how many real cameras the dataset provides.
|
||||
postprocess_action_dim (`int`, *optional*):
|
||||
Overrides the action dimensionality the postprocessor crops predictions down to. Falls back to
|
||||
the dataset's action feature width, or `max_action_dim` if that is unavailable.
|
||||
binarize_gripper (`bool`, *optional*, defaults to `False`):
|
||||
Whether the postprocessor snaps the gripper action channel to one of two fixed values instead
|
||||
of passing through the continuous prediction.
|
||||
gripper_index (`int`, *optional*, defaults to 6):
|
||||
Index of the gripper channel within the action vector, used when `binarize_gripper` is `True`.
|
||||
gripper_threshold (`float`, *optional*, defaults to 0.5):
|
||||
Decision threshold applied to the gripper channel when `binarize_gripper` is `True`.
|
||||
gripper_below_threshold_value (`float`, *optional*, defaults to 1.0):
|
||||
Value written to the gripper channel when it is at or below `gripper_threshold`.
|
||||
gripper_above_threshold_value (`float`, *optional*, defaults to -1.0):
|
||||
Value written to the gripper channel when it is above `gripper_threshold`.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||
vlm_model_name (`str`, *optional*, defaults to `"OpenGVLab/InternVL3-1B-hf"`):
|
||||
Hugging Face Hub id of the InternVL3 vision-language backbone.
|
||||
vlm_num_layers (`int`, *optional*, defaults to 14):
|
||||
Number of transformer layers kept from the InternVL3 language model. `None` keeps all of them.
|
||||
vlm_dtype (`str`, *optional*, defaults to `"bfloat16"`):
|
||||
Dtype the InternVL3 backbone is loaded in.
|
||||
max_text_length (`int`, *optional*, defaults to 1024):
|
||||
Maximum token length for the tokenized (image placeholders + instruction) prompt. Longer
|
||||
prompts are right-truncated.
|
||||
use_flash_attn (`bool`, *optional*, defaults to `True`):
|
||||
Whether to request FlashAttention in the InternVL3 backbone.
|
||||
action_head (`str`, *optional*, defaults to `"flowmatching"`):
|
||||
Identifier of the action-generation head architecture.
|
||||
embed_dim (`int`, *optional*, defaults to 896):
|
||||
Dimensionality of the fused vision-language token embeddings consumed by the action head.
|
||||
hidden_dim (`int`, *optional*, defaults to 1024):
|
||||
Hidden width of the action head's transformer layers.
|
||||
state_hidden_dim (`int`, *optional*, defaults to 1024):
|
||||
Hidden width of the state encoder inside the action head.
|
||||
num_heads (`int`, *optional*, defaults to 8):
|
||||
Number of attention heads in the action head's transformer layers.
|
||||
num_layers (`int`, *optional*, defaults to 8):
|
||||
Number of transformer layers in the action head.
|
||||
dropout (`float`, *optional*, defaults to 0.0):
|
||||
Dropout probability applied inside the action head.
|
||||
num_inference_timesteps (`int`, *optional*, defaults to 32):
|
||||
Number of integration steps used to sample an action chunk from the flow-matching head.
|
||||
num_categories (`int`, *optional*, defaults to 1):
|
||||
Number of embodiment categories the action head conditions on.
|
||||
return_cls_only (`bool`, *optional*, defaults to `False`):
|
||||
Whether the action head is conditioned on a single pooled VL token (the last non-padding token
|
||||
of the causal decoder) instead of the full fused token sequence.
|
||||
enable_gradient_checkpointing (`bool`, *optional*, defaults to `True`):
|
||||
Whether to enable gradient checkpointing on the VLM backbone to reduce memory usage.
|
||||
gradient_checkpointing_use_reentrant (`bool`, *optional*, defaults to `False`):
|
||||
Whether gradient checkpointing uses the reentrant autograd variant.
|
||||
finetune_vlm (`bool`, *optional*):
|
||||
Whether the whole VLM backbone is trainable. Defaulted from `training_stage` unless set
|
||||
explicitly with `apply_training_stage_defaults=False`. Must agree with the union of
|
||||
`finetune_language_model` and `finetune_vision_model` when those are set explicitly.
|
||||
finetune_language_model (`bool`, *optional*):
|
||||
Whether the VLM's language branch is trainable. Defaulted from `training_stage` unless set
|
||||
explicitly with `apply_training_stage_defaults=False`.
|
||||
finetune_vision_model (`bool`, *optional*):
|
||||
Whether the VLM's vision branch is trainable. Defaulted from `training_stage` unless set
|
||||
explicitly with `apply_training_stage_defaults=False`.
|
||||
finetune_action_head (`bool`, *optional*):
|
||||
Whether the flow-matching action head is trainable. Defaulted from `training_stage` unless set
|
||||
explicitly with `apply_training_stage_defaults=False`.
|
||||
apply_training_stage_defaults (`bool`, *optional*, defaults to `True`):
|
||||
Whether to reapply the `training_stage` defaults to the `finetune_*` flags after loading a
|
||||
checkpoint config, so a stage-2 run cannot silently inherit a stage-1 checkpoint's frozen-VLM
|
||||
flags. Set `False` to keep explicit finetuning flags.
|
||||
task_field (`str`, *optional*, defaults to `"task"`):
|
||||
Batch key holding the language instruction(s) passed to the VLM.
|
||||
embodiment_id_field (`str`, *optional*):
|
||||
Batch key holding an explicit per-sample embodiment id. Falls back to `"embodiment_id"`, then
|
||||
to `default_embodiment_id`, when unset or absent from the batch.
|
||||
default_embodiment_id (`int`, *optional*, defaults to 0):
|
||||
Embodiment id used when the batch carries none. Must be in `[0, num_categories)`.
|
||||
rtc_config (`RTCConfig`, *optional*):
|
||||
Real-Time Chunking guidance for asynchronous inference. `None` disables RTC.
|
||||
`lerobot-rollout --inference.type=rtc` sets this and calls `init_rtc_processor()`.
|
||||
optimizer_lr (`float`, *optional*, defaults to 1e-05):
|
||||
Learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`):
|
||||
Adam beta coefficients for the default optimizer preset.
|
||||
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||
Adam epsilon for the default optimizer preset.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 1e-05):
|
||||
Weight decay applied to the decayed parameter group in the default optimizer preset.
|
||||
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||
Gradient-norm clipping threshold for the default optimizer preset.
|
||||
scheduler_warmup_steps (`int`, *optional*, defaults to 300):
|
||||
Number of warmup steps for the default cosine-annealing-with-warmup scheduler preset.
|
||||
"""
|
||||
|
||||
training_stage: str = "stage1"
|
||||
# When True and the policy runs on CUDA, EVO1 wraps its own forward passes (training and
|
||||
# inference) in a bfloat16 autocast block, so its numerics do not depend on the dtype of any
|
||||
@@ -108,6 +257,7 @@ class Evo1Config(PreTrainedConfig):
|
||||
scheduler_warmup_steps: int = 300
|
||||
|
||||
def __post_init__(self):
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the VLM backbone/tokenizer configuration."""
|
||||
super().__post_init__()
|
||||
if self.training_stage not in {"stage1", "stage2"}:
|
||||
raise ValueError(
|
||||
@@ -200,6 +350,7 @@ class Evo1Config(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||
if self.input_features is None:
|
||||
self.input_features = {}
|
||||
if self.output_features is None:
|
||||
@@ -226,6 +377,7 @@ class Evo1Config(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
return AdamWConfig(
|
||||
lr=self.optimizer_lr,
|
||||
betas=self.optimizer_betas,
|
||||
@@ -235,18 +387,22 @@ class Evo1Config(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self):
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return CosineAnnealingWithWarmupSchedulerConfig(
|
||||
num_warmup_steps=self.scheduler_warmup_steps,
|
||||
)
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list[int]:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
return [0]
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list[int]:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return list(range(self.chunk_size))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -33,19 +33,43 @@ from .evo1_model import Evo1Model
|
||||
|
||||
|
||||
class ActionSelectKwargs(TypedDict, total=False):
|
||||
"""Extra keyword arguments accepted by EVO1's `select_action`/`predict_action_chunk` for RTC inference.
|
||||
|
||||
**Attributes**:
|
||||
- **inference_delay** (`int | None`) -- Number of environment steps the previous inference call
|
||||
took, used by the RTC processor to blend the new chunk with `prev_chunk_left_over`.
|
||||
- **prev_chunk_left_over** (`Tensor | None`) -- Unconsumed tail of the previously predicted action
|
||||
chunk, blended with the new prediction for a smooth handoff.
|
||||
- **execution_horizon** (`int | None`) -- Number of steps of the new chunk that will actually be
|
||||
executed before the next inference call, used to weight the RTC blend.
|
||||
"""
|
||||
|
||||
inference_delay: int | None
|
||||
prev_chunk_left_over: Tensor | None
|
||||
execution_horizon: int | None
|
||||
|
||||
|
||||
class Evo1Policy(PreTrainedPolicy):
|
||||
"""EVO1 vision-language-action policy: an InternVL3 backbone with a flow-matching action head."""
|
||||
|
||||
config_class = Evo1Config
|
||||
name = "evo1"
|
||||
|
||||
def supports_rtc(self) -> bool:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. EVO1 supports Real-Time Chunking."""
|
||||
return True
|
||||
|
||||
def __init__(self, config: Evo1Config, *, vlm_hub_kwargs: dict | None = None, **kwargs):
|
||||
"""Build the InternVL3 vision-language embedder and the flow-matching action head.
|
||||
|
||||
Args:
|
||||
config (`Evo1Config`):
|
||||
Policy configuration.
|
||||
vlm_hub_kwargs (`dict`, *optional*):
|
||||
Hub download options (`token`, `cache_dir`, `local_files_only`, `proxies`) forwarded to the
|
||||
VLM backbone's own `from_pretrained` call, as distinct from the ones used to load this
|
||||
policy's own checkpoint.
|
||||
"""
|
||||
super().__init__(config)
|
||||
config.validate_features()
|
||||
|
||||
@@ -93,6 +117,12 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
strict: bool | None = None,
|
||||
**kwargs,
|
||||
) -> T:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||
|
||||
Defaults `strict` to `True` instead of `False`, and additionally forwards `vlm_hub_kwargs` (or
|
||||
derives them from `token`, `cache_dir`, `local_files_only`, and `proxies`) to the InternVL3
|
||||
backbone's own `from_pretrained` call.
|
||||
"""
|
||||
if strict is None:
|
||||
strict = True
|
||||
vlm_hub_kwargs = kwargs.pop("vlm_hub_kwargs", None)
|
||||
@@ -170,6 +200,11 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
return nullcontext()
|
||||
|
||||
def get_optim_params(self) -> list[dict]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Splits parameters into a weight-decayed group and a no-decay group (biases and 1D/normalization
|
||||
parameters).
|
||||
"""
|
||||
decay, no_decay = [], []
|
||||
for name, param in self.named_parameters():
|
||||
if not param.requires_grad:
|
||||
@@ -186,6 +221,7 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
]
|
||||
|
||||
def reset(self):
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||
self._action_queue = deque([], maxlen=self.config.n_action_steps)
|
||||
|
||||
def _normalize_task_batch(self, batch: dict[str, Tensor | list[str] | str]) -> list[str]:
|
||||
@@ -362,6 +398,12 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
embedder.eval()
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
"""Set training mode, keeping the VLM embedder in eval mode when its weights are frozen.
|
||||
|
||||
Args:
|
||||
mode (`bool`, *optional*, defaults to `True`):
|
||||
Whether to set training (`True`) or evaluation (`False`) mode.
|
||||
"""
|
||||
super().train(mode)
|
||||
self._keep_frozen_embedder_eval()
|
||||
return self
|
||||
@@ -452,6 +494,12 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
return sq_error.sum() / active.sum()
|
||||
|
||||
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
Computes the flow-matching velocity-regression loss (squared error between the predicted and
|
||||
target velocity), masked to the active state/action dimensions and averaged per sample. Set
|
||||
`reduction="none"` to get the per-sample loss instead of the batch mean.
|
||||
"""
|
||||
prompts = self._normalize_task_batch(batch)
|
||||
image_batches, image_masks = self._collect_image_batches(batch)
|
||||
states, _state_mask = self._prepare_state(batch)
|
||||
@@ -486,6 +534,12 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Accepts `ActionSelectKwargs`'s RTC-specific arguments (`inference_delay`, `prev_chunk_left_over`,
|
||||
`execution_horizon`), which are rejected unless `config.rtc_config` is set and
|
||||
`init_rtc_processor()` has been called.
|
||||
"""
|
||||
inference_delay = kwargs.get("inference_delay")
|
||||
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
||||
execution_horizon = kwargs.get("execution_horizon")
|
||||
@@ -522,6 +576,11 @@ class Evo1Policy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
Uses an action queue populated by `predict_action_chunk`. Real-Time Chunking is not supported
|
||||
here; use `predict_action_chunk` directly when `config.rtc_config` is enabled.
|
||||
"""
|
||||
assert not self._rtc_enabled(), (
|
||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||
)
|
||||
|
||||
@@ -381,6 +381,25 @@ def make_evo1_pre_post_processors(
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""Build the pre/post-processor pipelines for EVO1.
|
||||
|
||||
The preprocessor pads observation state and training actions to EVO1's fixed `max_state_dim` /
|
||||
`max_action_dim` widths (tracking the padding with an `action_mask`) before normalizing and moving the
|
||||
batch to `config.device`. The postprocessor unnormalizes predicted actions, crops them back down to the
|
||||
real action dimensionality, optionally binarizes the gripper channel, and moves the result to CPU.
|
||||
|
||||
Args:
|
||||
config (`Evo1Config`):
|
||||
EVO1 policy configuration.
|
||||
dataset_stats (`dict[str, dict[str, torch.Tensor]]`, *optional*):
|
||||
Per-feature normalization statistics, as produced by `LeRobotDatasetMetadata.stats`. Padded to
|
||||
`max_state_dim`/`max_action_dim` before being handed to the (un)normalizer steps.
|
||||
|
||||
Returns:
|
||||
`tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]`: The preprocessor (batch of raw
|
||||
observations/actions -> model input) and postprocessor (model output -> environment action)
|
||||
pipelines.
|
||||
"""
|
||||
normalization_features = _evo1_normalization_features(config)
|
||||
action_features = _evo1_action_features(config)
|
||||
normalization_stats = _pad_evo1_stats(config, dataset_stats)
|
||||
|
||||
@@ -31,10 +31,8 @@ from lerobot.envs import EnvConfig, env_to_policy_features
|
||||
from lerobot.lerobot_types import PolicyAction
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
RelativeActionsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
batch_to_transition,
|
||||
policy_action_to_transition,
|
||||
transition_to_batch,
|
||||
@@ -78,55 +76,8 @@ def _reconnect_relative_absolute_steps(
|
||||
step.relative_step = relative_step
|
||||
|
||||
|
||||
def _ensure_relative_actions(
|
||||
preprocessor: PolicyProcessorPipeline, postprocessor: PolicyProcessorPipeline, policy_cfg
|
||||
) -> None:
|
||||
"""Enable (or inject) the relative/absolute action steps in a loaded pipeline.
|
||||
|
||||
When loading from a pretrained checkpoint, the saved processor is authoritative. If the base
|
||||
predates the relative-action feature (e.g. FastWAM/LingBot bases) its pipeline has no
|
||||
RelativeActionsProcessorStep, so lerobot-train's override cannot enable one — those override
|
||||
keys are popped before `from_pretrained` (else it raises) and we reconstruct the steps here.
|
||||
Bases that DO ship the (disabled) steps (e.g. pi0/pi05) are simply flipped on. No-op unless
|
||||
``policy_cfg.use_relative_actions`` is set, so non-relative runs are untouched.
|
||||
"""
|
||||
if not getattr(policy_cfg, "use_relative_actions", False):
|
||||
return
|
||||
|
||||
exclude_joints = list(getattr(policy_cfg, "relative_exclude_joints", []) or [])
|
||||
action_names = getattr(policy_cfg, "action_feature_names", None)
|
||||
|
||||
pre_steps = list(preprocessor.steps)
|
||||
relative_step = next((s for s in pre_steps if isinstance(s, RelativeActionsProcessorStep)), None)
|
||||
if relative_step is None:
|
||||
relative_step = RelativeActionsProcessorStep(
|
||||
enabled=True, exclude_joints=exclude_joints, action_names=action_names
|
||||
)
|
||||
# Insert right before the normalizer (raw -> relative -> normalize); fall back to the front.
|
||||
idx = next((i for i, s in enumerate(pre_steps) if isinstance(s, NormalizerProcessorStep)), 0)
|
||||
pre_steps.insert(idx, relative_step)
|
||||
preprocessor.steps = pre_steps
|
||||
else:
|
||||
relative_step.enabled = True
|
||||
relative_step.exclude_joints = exclude_joints
|
||||
relative_step.action_names = action_names
|
||||
|
||||
post_steps = list(postprocessor.steps)
|
||||
absolute_step = next((s for s in post_steps if isinstance(s, AbsoluteActionsProcessorStep)), None)
|
||||
if absolute_step is None:
|
||||
absolute_step = AbsoluteActionsProcessorStep(enabled=True, relative_step=relative_step)
|
||||
# Insert right after the unnormalizer (unnormalize -> absolute); fall back to the front.
|
||||
idx = next((i for i, s in enumerate(post_steps) if isinstance(s, UnnormalizerProcessorStep)), -1)
|
||||
post_steps.insert(idx + 1, absolute_step)
|
||||
postprocessor.steps = post_steps
|
||||
else:
|
||||
absolute_step.enabled = True
|
||||
absolute_step.relative_step = relative_step
|
||||
|
||||
|
||||
def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||
"""
|
||||
Retrieves a policy class by its registered name.
|
||||
"""Retrieves a policy class by its registered name.
|
||||
|
||||
Resolution is convention-based: the draccus-registered config class of ``name`` is
|
||||
looked up, its ``configuration_*`` module path is rewritten to ``modeling_*``, and
|
||||
@@ -136,7 +87,8 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||
``@PreTrainedConfig.register_subclass``).
|
||||
|
||||
Args:
|
||||
name: The registered name of the policy (e.g. "act", "diffusion", "pi0").
|
||||
name (`str`): The registered name of the policy (e.g. "act", "diffusion", "pi0").
|
||||
|
||||
Returns:
|
||||
The policy class corresponding to the given name.
|
||||
|
||||
@@ -148,16 +100,15 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||
|
||||
|
||||
def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
"""
|
||||
Instantiates a policy configuration object based on the policy type.
|
||||
"""Instantiates a policy configuration object based on the policy type.
|
||||
|
||||
This factory function simplifies the creation of policy configuration objects by
|
||||
mapping a string identifier to the corresponding config class.
|
||||
|
||||
Args:
|
||||
policy_type: The registered type of the policy (any name registered via
|
||||
``@PreTrainedConfig.register_subclass``, e.g. "act", "diffusion", "pi0").
|
||||
**kwargs: Keyword arguments to be passed to the configuration class constructor.
|
||||
policy_type (`str`): The registered type of the policy (any name registered via
|
||||
`@PreTrainedConfig.register_subclass`, e.g. "act", "diffusion", "pi0").
|
||||
kwargs (`Any`, *optional*): Keyword arguments to be passed to the configuration class constructor.
|
||||
|
||||
Returns:
|
||||
An instance of a `PreTrainedConfig` subclass.
|
||||
@@ -173,18 +124,21 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
|
||||
|
||||
class ProcessorConfigKwargs(TypedDict, total=False):
|
||||
"""
|
||||
A TypedDict defining the keyword arguments for processor configuration.
|
||||
"""A TypedDict defining the keyword arguments for processor configuration.
|
||||
|
||||
This provides type hints for the optional arguments passed to `make_pre_post_processors`,
|
||||
improving code clarity and enabling static analysis.
|
||||
|
||||
Attributes:
|
||||
preprocessor_config_filename: The filename for the preprocessor configuration.
|
||||
postprocessor_config_filename: The filename for the postprocessor configuration.
|
||||
preprocessor_overrides: A dictionary of overrides for the preprocessor configuration.
|
||||
postprocessor_overrides: A dictionary of overrides for the postprocessor configuration.
|
||||
dataset_stats: Dataset statistics for normalization.
|
||||
**Attributes**:
|
||||
- **preprocessor_config_filename** (`str | None`) -- The filename for the preprocessor configuration.
|
||||
- **postprocessor_config_filename** (`str | None`) -- The filename for the postprocessor
|
||||
configuration.
|
||||
- **preprocessor_overrides** (`dict[str, Any] | None`) -- A dictionary of overrides for the
|
||||
preprocessor configuration.
|
||||
- **postprocessor_overrides** (`dict[str, Any] | None`) -- A dictionary of overrides for the
|
||||
postprocessor configuration.
|
||||
- **dataset_stats** (`dict[str, dict[str, torch.Tensor]] | None`) -- Dataset statistics for
|
||||
normalization.
|
||||
"""
|
||||
|
||||
preprocessor_config_filename: str | None
|
||||
@@ -204,8 +158,7 @@ def make_pre_post_processors(
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""
|
||||
Create or load pre- and post-processor pipelines for a given policy.
|
||||
"""Create or load pre- and post-processor pipelines for a given policy.
|
||||
|
||||
This function acts as a factory. It can either load existing processor pipelines
|
||||
from a pretrained path or create new ones from scratch based on the policy
|
||||
@@ -216,6 +169,7 @@ def make_pre_post_processors(
|
||||
policy_cfg: The configuration of the policy for which to create processors.
|
||||
pretrained_path: An optional path to load pretrained processor pipelines from.
|
||||
If provided, pipelines are loaded from this path.
|
||||
pretrained_revision: The Hub revision to load `pretrained_path` from, if it's a Hub repo id.
|
||||
**kwargs: Keyword arguments for processor configuration, as defined in
|
||||
`ProcessorConfigKwargs`.
|
||||
|
||||
@@ -245,20 +199,12 @@ def make_pre_post_processors(
|
||||
),
|
||||
)
|
||||
|
||||
# The relative/absolute override keys only match if the saved base already contains those
|
||||
# steps (e.g. pi0/pi05). For bases that predate the feature (FastWAM/LingBot) they would
|
||||
# raise "Override keys ... do not match any step". Pop them here and let
|
||||
# _ensure_relative_actions() enable-or-inject the steps after loading (handles both cases).
|
||||
pre_overrides = dict(kwargs.get("preprocessor_overrides") or {})
|
||||
post_overrides = dict(kwargs.get("postprocessor_overrides") or {})
|
||||
pre_overrides.pop("relative_actions_processor", None)
|
||||
post_overrides.pop("absolute_actions_processor", None)
|
||||
preprocessor = PolicyProcessorPipeline.from_pretrained(
|
||||
pretrained_model_name_or_path=pretrained_path,
|
||||
config_filename=kwargs.get(
|
||||
"preprocessor_config_filename", f"{POLICY_PREPROCESSOR_DEFAULT_NAME}.json"
|
||||
),
|
||||
overrides=pre_overrides,
|
||||
overrides=kwargs.get("preprocessor_overrides", {}),
|
||||
to_transition=batch_to_transition,
|
||||
to_output=transition_to_batch,
|
||||
revision=pretrained_revision,
|
||||
@@ -268,12 +214,11 @@ def make_pre_post_processors(
|
||||
config_filename=kwargs.get(
|
||||
"postprocessor_config_filename", f"{POLICY_POSTPROCESSOR_DEFAULT_NAME}.json"
|
||||
),
|
||||
overrides=post_overrides,
|
||||
overrides=kwargs.get("postprocessor_overrides", {}),
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
revision=pretrained_revision,
|
||||
)
|
||||
_ensure_relative_actions(preprocessor, postprocessor, policy_cfg)
|
||||
_reconnect_relative_absolute_steps(preprocessor, postprocessor)
|
||||
if isinstance(policy_cfg, Evo1Config):
|
||||
from .evo1.processor_evo1 import reconcile_evo1_processors
|
||||
@@ -299,9 +244,9 @@ def make_policy(
|
||||
ds_meta: LeRobotDatasetMetadata | None = None,
|
||||
env_cfg: EnvConfig | None = None,
|
||||
rename_map: dict[str, str] | None = None,
|
||||
defer_weight_load: bool = False,
|
||||
) -> PreTrainedPolicy:
|
||||
"""
|
||||
Instantiate a policy model.
|
||||
"""Instantiate a policy model.
|
||||
|
||||
This factory function handles the logic of creating a policy, which requires
|
||||
determining the input and output feature shapes. These shapes can be derived
|
||||
@@ -309,22 +254,27 @@ def make_policy(
|
||||
can either initialize a new policy from scratch or load a pretrained one.
|
||||
|
||||
Args:
|
||||
cfg: The configuration for the policy to be created. If `cfg.pretrained_path` is
|
||||
set, the policy will be loaded with weights from that path.
|
||||
ds_meta: Dataset metadata used to infer feature shapes and types. Also provides
|
||||
statistics for normalization layers.
|
||||
env_cfg: Environment configuration used to infer feature shapes and types.
|
||||
One of `ds_meta` or `env_cfg` must be provided.
|
||||
rename_map: Optional mapping of dataset or environment feature keys to match
|
||||
expected policy feature names (e.g., `"left"` → `"camera1"`).
|
||||
cfg (PreTrainedConfig): The configuration for the policy to be created. If
|
||||
`cfg.pretrained_path` is set, the policy will be loaded with weights from that path.
|
||||
ds_meta (LeRobotDatasetMetadata | None, *optional*): Dataset metadata used to infer feature shapes and
|
||||
types. Also provides statistics for normalization layers.
|
||||
env_cfg (EnvConfig | None, *optional*): Environment configuration used to infer feature shapes and
|
||||
types. One of `ds_meta` or `env_cfg` must be provided.
|
||||
rename_map (dict[str, str] | None, *optional*): Optional mapping of dataset or environment feature
|
||||
keys to match expected policy feature names (e.g., `"left"` → `"camera1"`).
|
||||
defer_weight_load (bool, *optional*, defaults to `False`): Build the exact policy `from_pretrained` would build — same
|
||||
config resolution, same stats-derived buffers, same device placement and eval mode —
|
||||
but skip the safetensors weight load. Used when resuming from a DCP checkpoint, whose
|
||||
sharded weights stream in after `accelerator.prepare()` (the distributed checkpoint
|
||||
engine overwrites the random init).
|
||||
|
||||
Returns:
|
||||
An instantiated and device-placed policy model.
|
||||
PreTrainedPolicy: An instantiated and device-placed policy model.
|
||||
|
||||
Raises:
|
||||
ValueError: If both or neither of `ds_meta` and `env_cfg` are provided.
|
||||
NotImplementedError: If attempting to use an unsupported policy-backend
|
||||
combination (e.g., VQBeT with 'mps').
|
||||
NotImplementedError: If attempting to use an unsupported policy-backend combination
|
||||
(e.g., VQBeT with 'mps').
|
||||
"""
|
||||
if bool(ds_meta) == bool(env_cfg):
|
||||
raise ValueError("Either one of a dataset metadata or a sim env must be provided.")
|
||||
@@ -389,11 +339,18 @@ def make_policy(
|
||||
)
|
||||
|
||||
if cfg.pretrained_path and not cfg.use_peft:
|
||||
# Load a pretrained policy and override the config if needed (for example, if there are inference-time
|
||||
# hyperparameters that we want to vary).
|
||||
kwargs["pretrained_name_or_path"] = cfg.pretrained_path
|
||||
kwargs["revision"] = cfg.pretrained_revision
|
||||
policy = policy_cls.from_pretrained(**kwargs)
|
||||
if defer_weight_load:
|
||||
# Same construction path as from_pretrained (config already resolved from the
|
||||
# checkpoint by the caller; dataset_stats/dataset_meta kwargs identical), minus the
|
||||
# weight load — parity by construction.
|
||||
policy = policy_cls(**kwargs)
|
||||
policy.eval()
|
||||
else:
|
||||
# Load a pretrained policy and override the config if needed (for example, if there
|
||||
# are inference-time hyperparameters that we want to vary).
|
||||
kwargs["pretrained_name_or_path"] = cfg.pretrained_path
|
||||
kwargs["revision"] = cfg.pretrained_revision
|
||||
policy = policy_cls.from_pretrained(**kwargs)
|
||||
elif cfg.pretrained_path and cfg.use_peft:
|
||||
# Load a pretrained PEFT model on top of the policy. The pretrained path points to the folder/repo
|
||||
# of the adapter and the adapter's config contains the path to the base policy. So we need the
|
||||
@@ -452,6 +409,7 @@ def _get_policy_cls_from_policy_name(name: str) -> type[PreTrainedPolicy]:
|
||||
|
||||
Args:
|
||||
name: The name of the policy.
|
||||
|
||||
Returns:
|
||||
The policy class corresponding to the given name.
|
||||
"""
|
||||
@@ -507,10 +465,10 @@ def _make_processors_from_policy_config(
|
||||
dataset_stats: Dataset statistics for normalization.
|
||||
dataset_meta: Dataset metadata, forwarded only to factories that declare a
|
||||
``dataset_meta`` parameter (e.g. groot, molmoact2).
|
||||
|
||||
Returns:
|
||||
A tuple containing the input (pre-processor) and output (post-processor) pipelines.
|
||||
"""
|
||||
|
||||
policy_type = config.type
|
||||
function_name = f"make_{policy_type}_pre_post_processors"
|
||||
module_path = config.__class__.__module__.replace(
|
||||
|
||||
@@ -27,8 +27,6 @@ from lerobot.configs import (
|
||||
from lerobot.optim import AdamWConfig
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
|
||||
from ..rtc.configuration_rtc import RTCConfig
|
||||
|
||||
WAN22_MODEL_ID = "Wan-AI/Wan2.2-TI2V-5B"
|
||||
WAN22_DIFFUSERS_MODEL_ID = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FASTWAM_BASE_MODEL_ID = "lerobot/fastwam_base"
|
||||
@@ -60,6 +58,7 @@ _FASTWAM_ACTION_BASE_COMPAT_KEYS = (
|
||||
|
||||
|
||||
def default_video_dit_config(action_dim: int) -> dict[str, Any]:
|
||||
"""Return the default kwargs dict for the video-generation DiT backbone, sized for `action_dim`."""
|
||||
return {
|
||||
"patch_size": [1, 2, 2],
|
||||
"in_dim": 48,
|
||||
@@ -83,6 +82,7 @@ def default_video_dit_config(action_dim: int) -> dict[str, Any]:
|
||||
|
||||
|
||||
def default_action_dit_config(action_dim: int) -> dict[str, Any]:
|
||||
"""Return the default kwargs dict for the action-generation DiT backbone, sized for `action_dim`."""
|
||||
return {
|
||||
"action_dim": action_dim,
|
||||
"hidden_dim": 1024,
|
||||
@@ -138,7 +138,6 @@ def _validate_wan_model_id(value: str, field_name: str) -> str:
|
||||
|
||||
def is_fastwam_base_compatible_config(config: FastWAMConfig) -> bool:
|
||||
"""Return whether `fastwam_base` partial weights can initialize this config."""
|
||||
|
||||
default_video_config = default_video_dit_config(config.action_dim)
|
||||
default_action_config = default_action_dit_config(config.action_dim)
|
||||
return all(
|
||||
@@ -155,30 +154,129 @@ def is_fastwam_base_compatible_config(config: FastWAMConfig) -> bool:
|
||||
class FastWAMConfig(PreTrainedConfig):
|
||||
"""Configuration for the FastWAM LeRobot policy.
|
||||
|
||||
FastWAM adapts the Wan2.2 video-diffusion backbone into a robot policy: a video expert and an action
|
||||
expert are jointly trained (or fine-tuned) as a Mixture-of-Transformers, sharing attention over a
|
||||
predicted future video and the corresponding action chunk.
|
||||
|
||||
Args:
|
||||
action_dim (int): Number of scalar action channels per timestep.
|
||||
proprio_dim (int | None): Number of proprioception channels used as an
|
||||
extra text-context token. `None` disables proprio conditioning.
|
||||
action_horizon (int): Number of actions predicted by one policy call.
|
||||
num_video_frames (int): Raw video sampling window (in dataset frames). The
|
||||
model actually operates on `model_video_frames` frames after subsampling
|
||||
by `action_video_freq_ratio`.
|
||||
action_video_freq_ratio (int): Actions are sampled at this multiple of the
|
||||
video frame rate. Video frames are taken every `action_video_freq_ratio`-th
|
||||
raw frame, so the model sees `(num_video_frames - 1) // ratio + 1` frames
|
||||
spanning the same time window as `action_horizon` actions (ratio actions
|
||||
per video frame).
|
||||
image_size (tuple[int, int]): Concatenated image size as `(height, width)`.
|
||||
context_len (int): Maximum text embedding token length.
|
||||
video_dit_config (dict[str, Any] | None): Wan video expert config.
|
||||
action_dit_config (dict[str, Any] | None): Action expert config.
|
||||
use_gradient_checkpointing (bool): Enable activation checkpointing in both DiT
|
||||
experts (trades compute for memory; propagated into the DiT configs).
|
||||
freeze_video_expert (bool): Freeze the ~5B Wan video expert
|
||||
(`model.video_expert`) so only the action expert + proprio encoder train.
|
||||
Cuts the AdamW optimizer footprint substantially; the video expert keeps its
|
||||
pretrained weights. (If enabled, also set `loss.lambda_video=0` to skip the
|
||||
now-gradient-free video loss compute.)
|
||||
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||
Number of environment steps of observation to pass to the policy.
|
||||
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Input feature specification, keyed by feature name. `__post_init__` builds a synthetic
|
||||
single-image default at `image_size` when left unset; `set_dataset_feature_metadata` later
|
||||
replaces it with the dataset's real per-camera keys.
|
||||
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||
Output feature specification, keyed by feature name. `__post_init__` builds a default `action`
|
||||
feature of shape `(action_dim,)` when left unset.
|
||||
device (`str`, *optional*):
|
||||
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||
unavailable.
|
||||
use_amp (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub.
|
||||
repo_id (`str`, *optional*):
|
||||
Hub repository id to push the policy to.
|
||||
private (`bool`, *optional*):
|
||||
Whether the pushed Hub repository is private.
|
||||
tags (`list[str]`, *optional*):
|
||||
Tags to attach to the policy on the Hub.
|
||||
license (`str`, *optional*):
|
||||
License identifier for the policy on the Hub.
|
||||
pretrained_path (`Path`, *optional*):
|
||||
Repo id or local directory of pretrained weights saved with `save_pretrained`. Auto-populated
|
||||
from `base_model_id` when the DiT configs are `fastwam_base`-compatible; otherwise left unset
|
||||
to initialize from scratch.
|
||||
pretrained_revision (`str`, *optional*):
|
||||
Hub revision to pin when loading `pretrained_path`.
|
||||
action_dim (`int`, *optional*, defaults to 7):
|
||||
Number of scalar action channels per timestep.
|
||||
proprio_dim (`int`, *optional*, defaults to 8):
|
||||
Number of proprioception channels used as an extra text-context token. `None` disables proprio
|
||||
conditioning.
|
||||
action_horizon (`int`, *optional*, defaults to 32):
|
||||
Number of actions predicted by one policy call.
|
||||
n_action_steps (`int`, *optional*, defaults to 32):
|
||||
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||
policy. Must not exceed `action_horizon`.
|
||||
num_video_frames (`int`, *optional*, defaults to 33):
|
||||
Raw video sampling window, in dataset frames. The model actually operates on
|
||||
`model_video_frames` frames after subsampling by `action_video_freq_ratio`.
|
||||
action_video_freq_ratio (`int`, *optional*, defaults to 4):
|
||||
Actions are sampled at this multiple of the video frame rate. Video frames are taken every
|
||||
`action_video_freq_ratio`-th raw frame, so the model sees `(num_video_frames - 1) // ratio + 1`
|
||||
frames spanning the same time window as `action_horizon` actions.
|
||||
image_size (`tuple[int, int]`, *optional*, defaults to `(224, 448)`):
|
||||
Concatenated image size as `(height, width)`, shared across every camera view.
|
||||
context_len (`int`, *optional*, defaults to 128):
|
||||
Maximum text embedding token length.
|
||||
model_id (`str`, *optional*, defaults to `"Wan-AI/Wan2.2-TI2V-5B"`):
|
||||
Hub id (or local path) of the Wan2.2 video-diffusion backbone.
|
||||
tokenizer_model_id (`str`, *optional*, defaults to `"google/umt5-xxl"`):
|
||||
Hub id of the UMT5 tokenizer.
|
||||
text_encoder_model_id (`str`, *optional*, defaults to `"Wan-AI/Wan2.2-TI2V-5B-Diffusers"`):
|
||||
Hub id of the frozen UMT5 text encoder and VAE used for text/video conditioning.
|
||||
base_model_id (`str`, *optional*, defaults to `"lerobot/fastwam_base"`):
|
||||
Hub id of the FastWAM base checkpoint used to auto-populate `pretrained_path` when the DiT
|
||||
configs are compatible with it. `None` disables this auto-loading.
|
||||
tokenizer_max_len (`int`, *optional*, defaults to 128):
|
||||
Maximum token length passed to the tokenizer.
|
||||
load_text_encoder (`bool`, *optional*, defaults to `True`):
|
||||
Whether to load the frozen UMT5 text encoder. Disable when the batch always supplies
|
||||
precomputed `context`/`context_mask`.
|
||||
mot_checkpoint_mixed_attn (`bool`, *optional*, defaults to `False`):
|
||||
Whether the Mixture-of-Transformers module checkpoints its mixed video/action attention.
|
||||
torch_dtype (`str`, *optional*, defaults to `"bfloat16"`):
|
||||
Dtype the Wan backbone and action expert are built and run in.
|
||||
prompt_template (`str`, *optional*, defaults to `"A video recorded from a robot's point of view executing the following instruction: {task}"`):
|
||||
Template the raw `task` string is formatted into before text encoding.
|
||||
num_inference_steps (`int`, *optional*, defaults to 10):
|
||||
Number of denoising steps used at inference time.
|
||||
inference_seed (`int`, *optional*, defaults to 42):
|
||||
Random seed for the inference noise sampler. `None` samples fresh noise every call.
|
||||
rand_device (`str`, *optional*, defaults to `"cpu"`):
|
||||
Device the inference noise sampler draws from.
|
||||
text_cfg_scale (`float`, *optional*, defaults to 1.0):
|
||||
Classifier-free-guidance scale applied against `negative_prompt` at inference time.
|
||||
negative_prompt (`str`, *optional*, defaults to `""`):
|
||||
Negative prompt used for classifier-free guidance.
|
||||
sigma_shift (`float`, *optional*):
|
||||
Overrides the diffusion schedule's sigma shift at inference time. `None` uses the scheduler's
|
||||
own shift.
|
||||
tiled (`bool`, *optional*, defaults to `False`):
|
||||
Whether to run the Wan VAE in tiled mode to reduce memory use.
|
||||
fp32_attention (`bool`, *optional*, defaults to `True`):
|
||||
Whether the video and action DiT experts compute attention in fp32.
|
||||
use_gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||
Whether to enable activation checkpointing in both DiT experts, trading compute for memory.
|
||||
Propagated into `video_dit_config` and `action_dit_config`.
|
||||
freeze_video_expert (`bool`, *optional*, defaults to `False`):
|
||||
Whether to freeze the ~5B Wan video expert so only the action expert and proprio encoder
|
||||
train, cutting the AdamW optimizer footprint substantially. Also set `loss.lambda_video=0` to
|
||||
skip the now-gradient-free video loss compute.
|
||||
toggle_action_dimensions (`list[int]`, *optional*):
|
||||
Action dimensions the postprocessor flips between two fixed values, for LIBERO-style toggle
|
||||
actions such as the gripper. Empty disables the toggle.
|
||||
video_scheduler (`dict[str, float | int]`, *optional*):
|
||||
Train/inference shift and step-count settings for the video diffusion scheduler.
|
||||
action_scheduler (`dict[str, float | int]`, *optional*):
|
||||
Train/inference shift and step-count settings for the action diffusion scheduler.
|
||||
loss (`dict[str, float]`, *optional*):
|
||||
Per-term loss weights, keyed by `"lambda_video"` and `"lambda_action"`.
|
||||
video_dit_config (`dict[str, Any]`, *optional*):
|
||||
Wan video expert architecture config. Built from `default_video_dit_config(action_dim)` when
|
||||
left unset.
|
||||
action_dit_config (`dict[str, Any]`, *optional*):
|
||||
Action expert architecture config. Built from `default_action_dit_config(action_dim)` when
|
||||
left unset.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||
Learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 0.01):
|
||||
Weight decay for the default optimizer preset.
|
||||
"""
|
||||
|
||||
n_obs_steps: int = 1
|
||||
@@ -190,25 +288,12 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
action_video_freq_ratio: int = 4
|
||||
image_size: tuple[int, int] = (224, 448)
|
||||
context_len: int = 128
|
||||
|
||||
# Relative actions: converts absolute actions to relative (action -= state) during
|
||||
# preprocessing, and reverses it at postprocessing. Requires `proprio_dim` (OBS_STATE).
|
||||
use_relative_actions: bool = False
|
||||
# Joint names to keep absolute (not converted to relative). Empty list = all dims relative.
|
||||
relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
|
||||
# Populated at runtime from dataset metadata by make_policy (used to build the exclude mask).
|
||||
action_feature_names: list[str] | None = None
|
||||
model_id: str = WAN22_MODEL_ID
|
||||
tokenizer_model_id: str = WAN_T5_TOKENIZER_ID
|
||||
text_encoder_model_id: str = WAN22_DIFFUSERS_MODEL_ID
|
||||
base_model_id: str | None = FASTWAM_BASE_MODEL_ID
|
||||
tokenizer_max_len: int = 128
|
||||
load_text_encoder: bool = True
|
||||
# Device for the frozen ~11GB UMT5-XXL text encoder. `None` keeps it on the main
|
||||
# policy `device` (default). Set to e.g. "cpu" to keep it off the GPU and save VRAM;
|
||||
# prompts are then encoded on that device and the resulting embeddings moved to the
|
||||
# policy device. Trades GPU memory for slower (CPU) text encoding.
|
||||
text_encoder_device: str | None = None
|
||||
mot_checkpoint_mixed_attn: bool = False
|
||||
torch_dtype: str = "bfloat16"
|
||||
prompt_template: str = (
|
||||
@@ -216,11 +301,6 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
)
|
||||
num_inference_steps: int = 10
|
||||
inference_seed: int | None = 42
|
||||
# Real-Time Chunking (RTC): async chunk generation with prefix guidance so a new chunk
|
||||
# inpaints smoothly onto the still-executing tail of the previous one. `None` disables it
|
||||
# (default synchronous inference). Consumed by `RTCInferenceEngine`, which calls
|
||||
# `predict_action_chunk(..., inference_delay=, prev_chunk_left_over=)`.
|
||||
rtc_config: RTCConfig | None = None
|
||||
rand_device: str = "cpu"
|
||||
text_cfg_scale: float = 1.0
|
||||
negative_prompt: str = ""
|
||||
@@ -252,6 +332,7 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
optimizer_weight_decay: float = 1.0e-2
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the DiT/video backbone configuration."""
|
||||
super().__post_init__()
|
||||
self.image_size = tuple(self.image_size)
|
||||
self.model_id = _validate_wan_model_id(self.model_id, "model_id")
|
||||
@@ -299,14 +380,12 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
finally:
|
||||
self.pretrained_path = pretrained_path
|
||||
|
||||
@property
|
||||
def chunk_size(self) -> int:
|
||||
return self.action_horizon
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
return AdamWConfig(lr=self.optimizer_lr, weight_decay=self.optimizer_weight_decay)
|
||||
|
||||
def get_scheduler_preset(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return None
|
||||
|
||||
def set_dataset_feature_metadata(self, dataset_features: dict[str, Any]) -> None:
|
||||
@@ -341,6 +420,7 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
self.validate_features()
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||
if self.action_dim <= 0:
|
||||
raise ValueError(f"`action_dim` must be positive, got {self.action_dim}.")
|
||||
if self.action_horizon <= 0:
|
||||
@@ -401,12 +481,16 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def model_video_frames(self) -> int:
|
||||
"""Number of video frames the model actually operates on, after subsampling the
|
||||
raw `num_video_frames` window by `action_video_freq_ratio` (e.g. 33 -> 9)."""
|
||||
"""Number of video frames the model actually operates on.
|
||||
|
||||
Computed by subsampling the raw `num_video_frames` window by `action_video_freq_ratio` (e.g.
|
||||
33 -> 9).
|
||||
"""
|
||||
return (self.num_video_frames - 1) // self.action_video_freq_ratio + 1
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list[int]:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
# Load the video frames the model is supervised on: the future window subsampled by
|
||||
# action_video_freq_ratio (e.g. [0, 4, 8, ..., 32] -> 9 frames). Each video frame is
|
||||
# thus `action_video_freq_ratio` actions apart, while actions load at the full rate
|
||||
@@ -416,8 +500,10 @@ class FastWAMConfig(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list[int]:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return list(range(self.action_horizon))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -22,7 +22,6 @@ import torch
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||
from lerobot.policies.rtc.modeling_rtc import RTCProcessor
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
from lerobot.utils.import_utils import require_package
|
||||
|
||||
@@ -46,15 +45,13 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
arbitrary boolean ``[query, key]`` masks that the FlashAttention varlen API cannot express;
|
||||
installing ``flash-attn`` has no effect on the FastWAM path. (SDPA may still dispatch to
|
||||
PyTorch's own flash/mem-efficient/math kernel internally, unrelated to the ``flash-attn`` package.)
|
||||
|
||||
Args:
|
||||
config (FastWAMConfig): FastWAM policy configuration.
|
||||
dataset_stats (dict[str, dict[str, Tensor]] | None): Optional LeRobot
|
||||
dataset statistics passed by the training/evaluation stack.
|
||||
"""
|
||||
|
||||
config_class = FastWAMConfig
|
||||
name = "fastwam"
|
||||
# FSDP2 wrap units: MoTLayer is the single FSDP owner of each layer's expert blocks
|
||||
# (the blocks are re-parented onto it precisely so sharding has one boundary to hook).
|
||||
_fsdp_wrap_modules = ["MoTLayer"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -62,6 +59,17 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
dataset_stats: dict[str, dict[str, Tensor]] | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Build the FastWAM core model (video expert, action expert, and MoT router).
|
||||
|
||||
Args:
|
||||
config (`FastWAMConfig`):
|
||||
FastWAM policy configuration.
|
||||
dataset_stats (`dict[str, dict[str, Tensor]]`, *optional*):
|
||||
LeRobot dataset statistics passed by the training/evaluation stack. Accepted for
|
||||
signature compatibility with other policies but not otherwise used here.
|
||||
kwargs: Additional keyword arguments (e.g. `dataset_meta`) forwarded by `make_policy` or
|
||||
`from_pretrained`; accepted and ignored.
|
||||
"""
|
||||
# FastWAM's Wan2.2 backbone needs transformers (UMT5 text encoder/tokenizer) and
|
||||
# diffusers (Wan VAE), both behind the `fastwam` extra. Fail fast with an actionable
|
||||
# message in base installs rather than deep in Wan component construction.
|
||||
@@ -86,29 +94,8 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
for layer in mot.layers:
|
||||
if "video" in layer.blocks:
|
||||
layer.blocks["video"].requires_grad_(False)
|
||||
self.init_rtc_processor()
|
||||
self.reset()
|
||||
|
||||
def init_rtc_processor(self) -> None:
|
||||
"""Attach a Real-Time Chunking processor to the core model when configured.
|
||||
|
||||
Mirrors the PI0/PI05 pattern: the policy owns the `RTCProcessor` and hands it to
|
||||
the core `FastWAM` model, which consults it inside `infer_action`'s denoising loop.
|
||||
Must stay public and named exactly `init_rtc_processor`: the rollout loader
|
||||
(`lerobot.rollout.context`) sets `policy.config.rtc_config = cfg.inference.rtc` and
|
||||
then calls `policy.init_rtc_processor()` to (re)build the processor after load, so
|
||||
`--inference.type=rtc` alone is enough to enable guidance — no separate policy-side
|
||||
`rtc_config` needed. A private/renamed method would be silently skipped (guidance
|
||||
off), degrading RTC to unguided async chunk-swapping.
|
||||
"""
|
||||
self.rtc_processor = None
|
||||
if self.config.rtc_config is not None:
|
||||
self.rtc_processor = RTCProcessor(self.config.rtc_config)
|
||||
self.model.rtc_processor = self.rtc_processor
|
||||
|
||||
def _rtc_enabled(self) -> bool:
|
||||
return self.config.rtc_config is not None and self.config.rtc_config.enabled
|
||||
|
||||
@classmethod
|
||||
def _load_as_safetensor(cls, model, model_file: str, map_location: str, strict: bool):
|
||||
"""Shape-aware load that supports cross-embodiment fine-tuning.
|
||||
@@ -159,6 +146,12 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
return model
|
||||
|
||||
def get_optim_params(self) -> list[Tensor]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Returns a flat list of trainable tensors (DiT parameters plus the proprio encoder's, when
|
||||
present) rather than a param-group dict, so parameters frozen via `freeze_video_expert` are
|
||||
excluded.
|
||||
"""
|
||||
# Return the trainable tensors directly (a single param group). The optimizer
|
||||
# builder wraps these in a param group; returning a bare {"params": [...]} dict
|
||||
# instead would make `list(...)` yield the key string "params".
|
||||
@@ -171,25 +164,8 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
return [p for p in params if p.requires_grad]
|
||||
|
||||
def reset(self) -> None:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||
self._action_queue: deque[Tensor] = deque([], maxlen=self.config.n_action_steps)
|
||||
# Per-episode text-embedding cache (mirrors LingBot-VA's `_prompt_embeds`). The task
|
||||
# is fixed for an episode, so the ~11GB UMT5 encoder runs once on the first chunk and
|
||||
# the resulting context is reused for every subsequent chunk. Cleared here on reset so
|
||||
# a new episode's (possibly different) task is re-encoded. Proprio is still appended
|
||||
# fresh each chunk downstream, so only the text-only context is cached.
|
||||
self._cached_prompt: Any = None
|
||||
self._cached_context: Tensor | None = None
|
||||
self._cached_context_mask: Tensor | None = None
|
||||
|
||||
def _encode_prompt_cached(self, prompt: Any) -> tuple[Tensor, Tensor]:
|
||||
"""Encode `prompt` to `(context, context_mask)`, reusing the cache when the prompt is
|
||||
unchanged so UMT5 runs at most once per episode (per distinct task)."""
|
||||
if self._cached_context is None or self._cached_prompt != prompt:
|
||||
context, context_mask = self.model.encode_prompt(prompt)
|
||||
self._cached_prompt = prompt
|
||||
self._cached_context = context
|
||||
self._cached_context_mask = context_mask
|
||||
return self._cached_context, self._cached_context_mask
|
||||
|
||||
def _batch_to_training_sample(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
|
||||
"""Adapt a standard LeRobot batch to the FastWAM-native sample that
|
||||
@@ -223,73 +199,27 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
sample["proprio"] = state.unsqueeze(1) if state.ndim == 2 else state
|
||||
return sample
|
||||
|
||||
def forward(
|
||||
self, batch: dict[str, Tensor], reduction: str = "mean"
|
||||
) -> tuple[Tensor, dict[str, Any]]:
|
||||
"""Compute FastWAM training loss for a LeRobot batch.
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
Args:
|
||||
batch (dict[str, Tensor]): Batch containing FastWAM-ready keys
|
||||
(`video`, `action`, `context`, `context_mask`) or LeRobot keys
|
||||
that can be adapted (`observation.images.*`, `observation.state`,
|
||||
`action`, `action_is_pad`).
|
||||
reduction (str): "mean" returns the scalar loss (default, backward
|
||||
compatible); "none" returns per-sample losses of shape (batch_size,)
|
||||
for sample weighting (RA-BC).
|
||||
|
||||
Returns:
|
||||
tuple[Tensor, dict[str, Any]]: The loss to backprop (scalar for "mean",
|
||||
per-sample (B,) for "none"), and a dict of logging metrics (e.g.
|
||||
`loss_video`, `loss_action`) — the `(loss, output_dict)` contract the
|
||||
LeRobot training loop expects.
|
||||
Accepts either FastWAM-native batch keys (`video`, `action`, `context`, `context_mask`) or
|
||||
standard LeRobot keys (`observation.images.*`, `observation.state`, `action`, `action_is_pad`),
|
||||
which are adapted internally. The metrics dict includes per-term losses such as `loss_video` and
|
||||
`loss_action`.
|
||||
"""
|
||||
|
||||
sample = self._batch_to_training_sample(batch)
|
||||
loss, metrics = self.model.training_loss(sample, reduction=reduction)
|
||||
loss, metrics = self.model.training_loss(sample)
|
||||
return loss, dict(metrics or {})
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(
|
||||
self,
|
||||
batch: dict[str, Tensor],
|
||||
inference_delay: int | None = None,
|
||||
prev_chunk_left_over: Tensor | None = None,
|
||||
execution_horizon: int | None = None,
|
||||
**_: Any,
|
||||
) -> Tensor:
|
||||
"""Predict a chunk of actions from the current FastWAM observation.
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], **_: Any) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Args:
|
||||
batch (dict[str, Tensor]): Inference batch with `input_image` or
|
||||
image observation keys, plus `context/context_mask` or `prompt`.
|
||||
inference_delay (int | None): RTC — number of prefix steps assumed already
|
||||
executed by the time this chunk lands (from measured inference latency).
|
||||
prev_chunk_left_over (Tensor | None): RTC — the previous chunk's unexecuted
|
||||
action tail `[T_prev, action_dim]` in model space; guides denoising so the
|
||||
new chunk inpaints onto it. `None` (default) = plain synchronous inference.
|
||||
execution_horizon (int | None): RTC — override for the prefix-weight horizon;
|
||||
`None` falls back to `rtc_config.execution_horizon`.
|
||||
|
||||
Returns:
|
||||
Tensor: Action chunk with shape `[B, action_horizon, action_dim]`.
|
||||
Accepts an inference batch with `input_image` or image-observation keys, plus a `context`/
|
||||
`context_mask` pair or a `prompt`. Returns a chunk of shape `[B, action_horizon, action_dim]`.
|
||||
"""
|
||||
|
||||
self.eval()
|
||||
infer_kwargs = _batch_to_infer_kwargs(batch=batch, config=self.config)
|
||||
# Encode the task once per episode and reuse it (LingBot-VA parity): swap the raw
|
||||
# `prompt` for the cached `context`/`context_mask` so `infer_action` skips `encode_prompt`
|
||||
# and the text encoder isn't re-run every chunk. Skipped when the caller supplies its own
|
||||
# precomputed `context` (the two are mutually exclusive downstream).
|
||||
if infer_kwargs.get("context") is None and infer_kwargs.get("prompt") is not None:
|
||||
context, context_mask = self._encode_prompt_cached(infer_kwargs["prompt"])
|
||||
infer_kwargs["prompt"] = None
|
||||
infer_kwargs["context"] = context
|
||||
infer_kwargs["context_mask"] = context_mask
|
||||
# RTC guidance args flow straight to `infer_action`; they are inert unless an
|
||||
# RTCProcessor is attached, enabled, and `prev_chunk_left_over` is provided.
|
||||
infer_kwargs["inference_delay"] = inference_delay
|
||||
infer_kwargs["prev_chunk_left_over"] = prev_chunk_left_over
|
||||
infer_kwargs["execution_horizon"] = execution_horizon
|
||||
batch_size = _infer_kwargs_batch_size(infer_kwargs)
|
||||
if batch_size == 1:
|
||||
action = _action_from_model_output(self.model.infer_action(**infer_kwargs))
|
||||
@@ -309,6 +239,7 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor], **kwargs: Any) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Uses an action queue populated by `predict_action_chunk`."""
|
||||
self.eval()
|
||||
if len(self._action_queue) == 0:
|
||||
actions = self.predict_action_chunk(batch, **kwargs)[:, : self.config.n_action_steps]
|
||||
@@ -334,10 +265,9 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
mixtures={"video": video_expert, "action": action_expert},
|
||||
mot_checkpoint_mixed_attn=config.mot_checkpoint_mixed_attn,
|
||||
)
|
||||
text_encoder_device = config.text_encoder_device or device
|
||||
text_encoder = (
|
||||
load_pretrained_wan_text_encoder(
|
||||
model_id=config.text_encoder_model_id, torch_dtype=dtype, device=text_encoder_device
|
||||
model_id=config.text_encoder_model_id, torch_dtype=dtype, device=device
|
||||
)
|
||||
if config.load_text_encoder
|
||||
else None
|
||||
@@ -348,7 +278,6 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
mot=mot,
|
||||
vae=load_pretrained_wan_vae(torch_dtype=dtype, device=device),
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_device=config.text_encoder_device,
|
||||
tokenizer=build_wan_tokenizer(
|
||||
model_id=config.tokenizer_model_id, tokenizer_max_len=config.tokenizer_max_len
|
||||
),
|
||||
|
||||
@@ -21,13 +21,10 @@ import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
ActionProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RelativeActionsProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
)
|
||||
@@ -76,14 +73,13 @@ def make_fastwam_pre_post_processors(
|
||||
Args:
|
||||
config (FastWAMConfig): Policy configuration controlling device and
|
||||
normalization feature metadata.
|
||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None): Optional
|
||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): Optional
|
||||
LeRobot dataset statistics used by normalization processors.
|
||||
|
||||
Returns:
|
||||
tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]: Input and
|
||||
output processor pipelines discoverable by LeRobot.
|
||||
"""
|
||||
|
||||
# NOTE: no visual normalization here. VISUAL is IDENTITY (see configuration_fastwam.normalization_mapping)
|
||||
# — images pass through in [0, 1] and the model maps them to the Wan VAE's [-1, 1] at the encode
|
||||
# boundary. This is deliberate: `lerobot_train.py` overrides the normalizer stats with
|
||||
@@ -101,25 +97,14 @@ def make_fastwam_pre_post_processors(
|
||||
|
||||
steps = make_default_policy_processor_steps(config, normalization_stats, normalizer_device=config.device)
|
||||
|
||||
# Shared relative-action step (OpenPI order: raw -> relative -> normalize -> model ->
|
||||
# unnormalize -> absolute). The SAME instance is passed to AbsoluteActionsProcessorStep
|
||||
# below so its cached raw state (set during preprocessing) flows to postprocessing.
|
||||
relative_step = RelativeActionsProcessorStep(
|
||||
enabled=config.use_relative_actions,
|
||||
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
)
|
||||
|
||||
input_steps: list[ProcessorStep] = [
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
steps.to_device,
|
||||
relative_step,
|
||||
steps.normalize,
|
||||
]
|
||||
output_steps: list[ProcessorStep] = [
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step),
|
||||
]
|
||||
if config.toggle_action_dimensions:
|
||||
output_steps.append(
|
||||
|
||||
@@ -839,7 +839,6 @@ class FastWAM(torch.nn.Module):
|
||||
text_dim: int | None = None,
|
||||
proprio_dim: int | None = None,
|
||||
device: str = "cpu",
|
||||
text_encoder_device: str | torch.device | None = None,
|
||||
torch_dtype: torch.dtype = torch.float32,
|
||||
video_train_shift: float = 5.0,
|
||||
video_infer_shift: float = 5.0,
|
||||
@@ -908,27 +907,12 @@ class FastWAM(torch.nn.Module):
|
||||
self.train_scheduler = self.train_video_scheduler
|
||||
self.infer_scheduler = self.infer_video_scheduler
|
||||
|
||||
# Optional Real-Time Chunking processor (set by the policy wrapper). When present and
|
||||
# enabled it guides the action denoising loop in `infer_action` so a freshly generated
|
||||
# chunk inpaints onto the previous chunk's unexecuted tail. Plain attribute (not an
|
||||
# nn.Module) — carries no parameters and stays out of state_dict / device moves.
|
||||
self.rtc_processor = None
|
||||
self.device = torch.device(device)
|
||||
# When pinned (e.g. "cpu"), the frozen text encoder stays on this device instead
|
||||
# of following the model onto the GPU — `_apply` skips it and `encode_prompt` runs
|
||||
# it here, moving embeddings back to `self.device`. `None` = follow `self.device`.
|
||||
self._text_encoder_device = (
|
||||
torch.device(text_encoder_device) if text_encoder_device is not None else None
|
||||
)
|
||||
self.torch_dtype = torch_dtype
|
||||
self.loss_lambda_video = float(loss_lambda_video)
|
||||
self.loss_lambda_action = float(loss_lambda_action)
|
||||
|
||||
self.to(self.device)
|
||||
# `self.to` above (via `_apply`) skips a pinned text encoder; make sure it actually
|
||||
# sits on the pinned device (it was loaded there, but this is a cheap safety net).
|
||||
if self.text_encoder is not None and self._text_encoder_device is not None:
|
||||
self.text_encoder._apply(lambda t: t.to(self._text_encoder_device))
|
||||
|
||||
@classmethod
|
||||
def from_wan22_pretrained(
|
||||
@@ -1019,8 +1003,7 @@ class FastWAM(torch.nn.Module):
|
||||
# while staying out of `state_dict()` / `parameters()`.
|
||||
super()._apply(fn, *args, **kwargs)
|
||||
self.vae._apply(fn)
|
||||
# A pinned text encoder (e.g. on CPU) must NOT follow device moves — leave it put.
|
||||
if self.text_encoder is not None and self._text_encoder_device is None:
|
||||
if self.text_encoder is not None:
|
||||
self.text_encoder._apply(fn)
|
||||
return self
|
||||
|
||||
@@ -1041,12 +1024,9 @@ class FastWAM(torch.nn.Module):
|
||||
"Prompt encoding requires loaded text encoder/tokenizer. "
|
||||
"Set `load_text_encoder=true` or provide precomputed `context/context_mask`."
|
||||
)
|
||||
# Run the encoder on its own device (may be pinned to CPU to save VRAM), then
|
||||
# move the resulting embeddings/mask to the model device for the DiT.
|
||||
te_device = self._text_encoder_device or self.device
|
||||
ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True)
|
||||
ids = ids.to(te_device)
|
||||
mask = mask.to(te_device, dtype=torch.bool)
|
||||
ids = ids.to(self.device)
|
||||
mask = mask.to(self.device, dtype=torch.bool)
|
||||
prompt_emb = self.text_encoder(ids, mask)
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
for i, v in enumerate(seq_lens):
|
||||
@@ -1054,7 +1034,7 @@ class FastWAM(torch.nn.Module):
|
||||
# Match FastWAM/Wan2.2 context semantics: padding embeddings are zeroed,
|
||||
# while cross-attention still sees a fixed-length context.
|
||||
mask = torch.ones_like(mask)
|
||||
return prompt_emb.to(device=self.device), mask.to(device=self.device)
|
||||
return prompt_emb.to(device=self.device), mask
|
||||
|
||||
def _append_proprio_to_context(
|
||||
self,
|
||||
@@ -1379,9 +1359,7 @@ class FastWAM(torch.nn.Module):
|
||||
pred_action = self.action_expert.post_dit(tokens_out["action"], action_pre)
|
||||
return pred_video, pred_action
|
||||
|
||||
def _compute_training_video_loss(
|
||||
self, inputs, pred_video, target_video, timestep_video, reduction: str = "mean"
|
||||
):
|
||||
def _compute_training_video_loss(self, inputs, pred_video, target_video, timestep_video):
|
||||
include_initial_video_step = inputs["first_frame_latents"] is None
|
||||
if inputs["first_frame_latents"] is not None:
|
||||
pred_video = pred_video[:, :, 1:]
|
||||
@@ -1396,13 +1374,9 @@ class FastWAM(torch.nn.Module):
|
||||
loss_video_per_sample.device,
|
||||
dtype=loss_video_per_sample.dtype,
|
||||
)
|
||||
weighted = loss_video_per_sample * video_weight
|
||||
# reduction="none" returns the per-sample vector (B,) for sample weighting (RA-BC).
|
||||
return weighted if reduction == "none" else weighted.mean()
|
||||
return (loss_video_per_sample * video_weight).mean()
|
||||
|
||||
def _compute_training_action_loss(
|
||||
self, inputs, pred_action, target_action, timestep_action, reduction: str = "mean"
|
||||
):
|
||||
def _compute_training_action_loss(self, inputs, pred_action, target_action, timestep_action):
|
||||
action_loss_token = functional.mse_loss(
|
||||
pred_action.float(), target_action.float(), reduction="none"
|
||||
).mean(dim=2)
|
||||
@@ -1419,11 +1393,9 @@ class FastWAM(torch.nn.Module):
|
||||
action_loss_per_sample.device,
|
||||
dtype=action_loss_per_sample.dtype,
|
||||
)
|
||||
weighted = action_loss_per_sample * action_weight
|
||||
# reduction="none" returns the per-sample vector (B,) for sample weighting (RA-BC).
|
||||
return weighted if reduction == "none" else weighted.mean()
|
||||
return (action_loss_per_sample * action_weight).mean()
|
||||
|
||||
def training_loss(self, sample, tiled: bool = False, reduction: str = "mean"):
|
||||
def training_loss(self, sample, tiled: bool = False):
|
||||
inputs = self.build_inputs(sample, tiled=tiled)
|
||||
targets = self._sample_training_targets(inputs)
|
||||
pred_video, pred_action = self._run_training_mot(inputs=inputs, targets=targets)
|
||||
@@ -1432,20 +1404,17 @@ class FastWAM(torch.nn.Module):
|
||||
pred_video=pred_video,
|
||||
target_video=targets["target_video"],
|
||||
timestep_video=targets["timestep_video"],
|
||||
reduction=reduction,
|
||||
)
|
||||
loss_action = self._compute_training_action_loss(
|
||||
inputs=inputs,
|
||||
pred_action=pred_action,
|
||||
target_action=targets["target_action"],
|
||||
timestep_action=targets["timestep_action"],
|
||||
reduction=reduction,
|
||||
)
|
||||
# With reduction="none" both terms are (B,), so loss_total is the per-sample loss (B,).
|
||||
loss_total = self.loss_lambda_video * loss_video + self.loss_lambda_action * loss_action
|
||||
loss_dict = {
|
||||
"loss_video": self.loss_lambda_video * float(loss_video.detach().mean().item()),
|
||||
"loss_action": self.loss_lambda_action * float(loss_action.detach().mean().item()),
|
||||
"loss_video": self.loss_lambda_video * float(loss_video.detach().item()),
|
||||
"loss_action": self.loss_lambda_action * float(loss_action.detach().item()),
|
||||
}
|
||||
return loss_total, loss_dict
|
||||
|
||||
@@ -1830,9 +1799,6 @@ class FastWAM(torch.nn.Module):
|
||||
seed: int | None = None,
|
||||
rand_device: str = "cpu",
|
||||
tiled: bool = False,
|
||||
inference_delay: int | None = None,
|
||||
prev_chunk_left_over: torch.Tensor | None = None,
|
||||
execution_horizon: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
self.eval()
|
||||
if str(getattr(self.video_expert, "video_attention_mask_mode", "")) != "first_frame_causal":
|
||||
@@ -1885,43 +1851,18 @@ class FastWAM(torch.nn.Module):
|
||||
dtype=latents_action.dtype,
|
||||
shift_override=sigma_shift,
|
||||
)
|
||||
rtc_active = (
|
||||
self.rtc_processor is not None
|
||||
and getattr(self.rtc_processor.rtc_config, "enabled", False)
|
||||
and prev_chunk_left_over is not None
|
||||
)
|
||||
num_train_timesteps = float(self.infer_action_scheduler.num_train_timesteps)
|
||||
for step_t_action, step_delta_action in zip(infer_timesteps_action, infer_deltas_action, strict=True):
|
||||
timestep_action = step_t_action.unsqueeze(0).to(dtype=latents_action.dtype, device=self.device)
|
||||
|
||||
def denoise(x_t, ts=timestep_action):
|
||||
return self._predict_action_noise_with_cache(
|
||||
latents_action=x_t,
|
||||
timestep_action=ts,
|
||||
context=context,
|
||||
context_mask=context_mask,
|
||||
video_kv_cache=video_kv_cache,
|
||||
attention_mask=attention_mask,
|
||||
video_seq_len=video_seq_len,
|
||||
)
|
||||
|
||||
if rtc_active:
|
||||
# `time` is the flow-matching noise level sigma in [0, 1]: FastWAM's model
|
||||
# predicts velocity v = noise - clean, so the clean-action estimate is
|
||||
# x1 = x_t - sigma * v — exactly RTC's `x1_t = x_t - time * v_t`.
|
||||
sigma = float(step_t_action.item()) / num_train_timesteps
|
||||
pred_action = self.rtc_processor.denoise_step(
|
||||
x_t=latents_action,
|
||||
prev_chunk_left_over=prev_chunk_left_over.to(
|
||||
device=latents_action.device, dtype=latents_action.dtype
|
||||
),
|
||||
inference_delay=inference_delay or 0,
|
||||
time=sigma,
|
||||
original_denoise_step_partial=denoise,
|
||||
execution_horizon=execution_horizon,
|
||||
)
|
||||
else:
|
||||
pred_action = denoise(latents_action)
|
||||
pred_action = self._predict_action_noise_with_cache(
|
||||
latents_action=latents_action,
|
||||
timestep_action=timestep_action,
|
||||
context=context,
|
||||
context_mask=context_mask,
|
||||
video_kv_cache=video_kv_cache,
|
||||
attention_mask=attention_mask,
|
||||
video_seq_len=video_seq_len,
|
||||
)
|
||||
|
||||
latents_action = self.infer_action_scheduler.step(pred_action, step_delta_action, latents_action)
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ def is_image_feature(key: str) -> bool:
|
||||
"""Check if a feature key represents an image feature.
|
||||
|
||||
Args:
|
||||
key: The feature key to check
|
||||
key (`str`): The feature key to check.
|
||||
|
||||
Returns:
|
||||
True if the key represents an image feature, False otherwise
|
||||
@@ -54,6 +54,8 @@ class ConcurrencyConfig:
|
||||
|
||||
@dataclass
|
||||
class ActorLearnerConfig:
|
||||
"""Actor-learner distributed architecture settings (network address, weight-push frequency)."""
|
||||
|
||||
learner_host: str = "127.0.0.1"
|
||||
learner_port: int = 50051
|
||||
policy_parameters_push_frequency: int = 4
|
||||
@@ -62,6 +64,8 @@ class ActorLearnerConfig:
|
||||
|
||||
@dataclass
|
||||
class CriticNetworkConfig:
|
||||
"""MLP architecture settings for the critic network(s)."""
|
||||
|
||||
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
||||
activate_final: bool = True
|
||||
final_activation: str | None = None
|
||||
@@ -69,12 +73,16 @@ class CriticNetworkConfig:
|
||||
|
||||
@dataclass
|
||||
class ActorNetworkConfig:
|
||||
"""MLP architecture settings for the actor network."""
|
||||
|
||||
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
||||
activate_final: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class PolicyConfig:
|
||||
"""Gaussian-policy output-head settings (tanh squashing, std clamping)."""
|
||||
|
||||
use_tanh_squash: bool = True
|
||||
std_min: float = 1e-5
|
||||
std_max: float = 10.0
|
||||
@@ -94,9 +102,95 @@ class GaussianActorConfig(PreTrainedConfig):
|
||||
logic live on the algorithm side (see ``lerobot.rl.algorithms.sac``).
|
||||
|
||||
CLI: ``--policy.type=gaussian_actor``.
|
||||
|
||||
Args:
|
||||
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||
Number of environment steps of observation to pass to the policy (the current step and
|
||||
additional steps going back). This policy predicts a single action from a single step, so
|
||||
this is not expected to be changed from 1.
|
||||
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||
automatically from the dataset when not explicitly provided.
|
||||
device (`str`, *optional*, defaults to `"cpu"`):
|
||||
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`.
|
||||
use_amp (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`):
|
||||
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||
repo_id (`str | None`, *optional*):
|
||||
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||
private (`bool | None`, *optional*):
|
||||
Whether to create/push the Hub repository as private.
|
||||
tags (`list[str] | None`, *optional*):
|
||||
Tags to attach to the policy's Hub model card.
|
||||
license (`str | None`, *optional*):
|
||||
License identifier to add to the policy's Hub model card.
|
||||
pretrained_path (`Path | None`, *optional*):
|
||||
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the
|
||||
policy is initialized from scratch.
|
||||
pretrained_revision (`str | None`, *optional*):
|
||||
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||
it. Defaults to mean/std normalization for visual features and min/max normalization for
|
||||
state, environment, and action features.
|
||||
dataset_stats (`dict[str, dict[str, list[float]]] | None`, *optional*):
|
||||
Statistics used to normalize image, state, and action features. Defaults to placeholder
|
||||
values; normally overridden with statistics computed from the actual training dataset.
|
||||
storage_device (`str`, *optional*, defaults to `"cpu"`):
|
||||
Device on which a copy of the model's parameters is kept for transport between the actor and
|
||||
learner processes in the actor-learner architecture.
|
||||
vision_encoder_name (`str | None`, *optional*):
|
||||
Name of a pretrained vision encoder to use for image observations, e.g.
|
||||
`"lerobot/resnet10"` for the HIL-SERL ResNet10 encoder. `None` (the default) uses a
|
||||
lightweight from-scratch CNN encoder instead.
|
||||
freeze_vision_encoder (`bool`, *optional*, defaults to `True`):
|
||||
Whether to freeze the vision encoder's parameters during training.
|
||||
image_encoder_hidden_dim (`int`, *optional*, defaults to 32):
|
||||
Hidden dimension size for the from-scratch image encoder (unused when `vision_encoder_name`
|
||||
is set).
|
||||
shared_encoder (`bool`, *optional*, defaults to `True`):
|
||||
Whether the actor and critic(s) share the same observation encoder instance.
|
||||
num_discrete_actions (`int | None`, *optional*):
|
||||
Number of discrete actions appended to the continuous action output, e.g. for a gripper
|
||||
open/close action. `None` disables the discrete critic and action head.
|
||||
image_embedding_pooling_dim (`int`, *optional*, defaults to 8):
|
||||
Number of learned spatial pooling features per image, used by the image encoder's spatial
|
||||
embedding layer.
|
||||
state_encoder_hidden_dim (`int`, *optional*, defaults to 256):
|
||||
Hidden dimension size for the state encoder.
|
||||
latent_dim (`int`, *optional*, defaults to 256):
|
||||
Dimension of the observation encoder's output latent space.
|
||||
online_steps (`int`, *optional*, defaults to 1000000):
|
||||
Number of steps to run during online training.
|
||||
online_buffer_capacity (`int`, *optional*, defaults to 100000):
|
||||
Capacity of the online replay buffer.
|
||||
offline_buffer_capacity (`int`, *optional*, defaults to 100000):
|
||||
Capacity of the offline replay buffer.
|
||||
async_prefetch (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use asynchronous prefetching for the replay buffers.
|
||||
online_step_before_learning (`int`, *optional*, defaults to 100):
|
||||
Number of steps to collect before online learning starts.
|
||||
actor_learner_config (`ActorLearnerConfig`, *optional*):
|
||||
Transport configuration (host, port, push frequency, queue timeout) for the actor-learner
|
||||
architecture.
|
||||
concurrency (`ConcurrencyConfig`, *optional*):
|
||||
Concurrency configuration (threads or processes) for the actor and learner.
|
||||
actor_network_kwargs (`ActorNetworkConfig`, *optional*):
|
||||
Architecture configuration (hidden dimensions, final activation) for the actor network.
|
||||
policy_kwargs (`PolicyConfig`, *optional*):
|
||||
Configuration for the Gaussian policy head (tanh squashing, std bounds, final-layer init
|
||||
scale).
|
||||
discrete_critic_network_kwargs (`CriticNetworkConfig`, *optional*):
|
||||
Architecture configuration (hidden dimensions, final activation) for the discrete critic
|
||||
network.
|
||||
"""
|
||||
|
||||
# Mapping of feature types to normalization modes
|
||||
normalization_mapping: dict[str, NormalizationMode] = field(
|
||||
default_factory=lambda: {
|
||||
"VISUAL": NormalizationMode.MEAN_STD,
|
||||
@@ -106,7 +200,6 @@ class GaussianActorConfig(PreTrainedConfig):
|
||||
}
|
||||
)
|
||||
|
||||
# Statistics for normalizing different types of inputs
|
||||
dataset_stats: dict[str, dict[str, list[float]]] | None = field(
|
||||
default_factory=lambda: {
|
||||
OBS_IMAGE: {
|
||||
@@ -125,60 +218,42 @@ class GaussianActorConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
# Architecture specifics
|
||||
# Device to run the model on (e.g., "cuda", "cpu")
|
||||
device: str = "cpu"
|
||||
# Device to store the model on
|
||||
storage_device: str = "cpu"
|
||||
# Name of the vision encoder model (Set to "lerobot/resnet10" for hil serl resnet10)
|
||||
vision_encoder_name: str | None = None
|
||||
# Whether to freeze the vision encoder during training
|
||||
freeze_vision_encoder: bool = True
|
||||
# Hidden dimension size for the image encoder
|
||||
image_encoder_hidden_dim: int = 32
|
||||
# Whether to use a shared encoder for actor and critic
|
||||
shared_encoder: bool = True
|
||||
# Number of discrete actions, eg for gripper actions
|
||||
num_discrete_actions: int | None = None
|
||||
# Dimension of the image embedding pooling
|
||||
image_embedding_pooling_dim: int = 8
|
||||
|
||||
# Encoder architecture
|
||||
# Hidden dimension size for the state encoder
|
||||
state_encoder_hidden_dim: int = 256
|
||||
# Dimension of the latent space
|
||||
latent_dim: int = 256
|
||||
|
||||
# Online training (TODO(Khalil): relocate to TrainRLServerPipelineConfig)
|
||||
# Number of steps for online training
|
||||
online_steps: int = 1000000
|
||||
# Capacity of the online replay buffer
|
||||
online_buffer_capacity: int = 100000
|
||||
# Capacity of the offline replay buffer
|
||||
offline_buffer_capacity: int = 100000
|
||||
# Whether to use asynchronous prefetching for the buffers
|
||||
async_prefetch: bool = False
|
||||
# Number of steps before learning starts
|
||||
online_step_before_learning: int = 100
|
||||
|
||||
# Actor-learner transport (TODO(Khalil): relocate to TrainRLServerPipelineConfig).
|
||||
# Configuration for actor-learner architecture
|
||||
actor_learner_config: ActorLearnerConfig = field(default_factory=ActorLearnerConfig)
|
||||
# Configuration for concurrency settings (you can use threads or processes for the actor and learner)
|
||||
concurrency: ConcurrencyConfig = field(default_factory=ConcurrencyConfig)
|
||||
|
||||
# Network architecture
|
||||
# Configuration for the actor network architecture
|
||||
actor_network_kwargs: ActorNetworkConfig = field(default_factory=ActorNetworkConfig)
|
||||
# Configuration for the policy parameters (Gaussian head)
|
||||
policy_kwargs: PolicyConfig = field(default_factory=PolicyConfig)
|
||||
# Configuration for the discrete critic network
|
||||
discrete_critic_network_kwargs: CriticNetworkConfig = field(default_factory=CriticNetworkConfig)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates actor/critic network and learner configuration."""
|
||||
super().__post_init__()
|
||||
# Any validation specific to GaussianActor configuration
|
||||
|
||||
def get_optimizer_preset(self) -> MultiAdamConfig:
|
||||
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||
# Default learning rate used to satisfy the abstract ``get_optimizer_preset()``
|
||||
# contract from ``PreTrainedConfig``. The actual optimizers used during RL
|
||||
# training are built by ``SACAlgorithm.make_optimizers_and_scheduler()`` from
|
||||
@@ -195,9 +270,11 @@ class GaussianActorConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||
return None
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||
has_image = any(is_image_feature(key) for key in self.input_features)
|
||||
has_state = OBS_STATE in self.input_features
|
||||
|
||||
@@ -211,16 +288,20 @@ class GaussianActorConfig(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def image_features(self) -> list[str]:
|
||||
"""The names of the input features that are images."""
|
||||
return [key for key in self.input_features if is_image_feature(key)]
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list:
|
||||
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list:
|
||||
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||
return None # SAC typically predicts one action at a time
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||
return None
|
||||
|
||||
@@ -35,6 +35,14 @@ DISCRETE_DIMENSION_INDEX = -1 # Gripper is always the last dimension
|
||||
class GaussianActorPolicy(
|
||||
PreTrainedPolicy,
|
||||
):
|
||||
"""Tanh-squashed diagonal Gaussian actor policy for SAC and related maximum-entropy continuous-control
|
||||
algorithms.
|
||||
|
||||
This policy only implements the actor (and its observation encoder) plus an optional discrete-action
|
||||
critic head; the Q-critics, temperature, and Bellman-update logic live on the algorithm side (see
|
||||
`lerobot.rl.algorithms.sac`).
|
||||
"""
|
||||
|
||||
config_class = GaussianActorConfig
|
||||
name = "gaussian_actor"
|
||||
|
||||
@@ -42,6 +50,11 @@ class GaussianActorPolicy(
|
||||
self,
|
||||
config: GaussianActorConfig | None = None,
|
||||
):
|
||||
"""Build the observation encoder(s), the Gaussian actor network, and the optional discrete critic.
|
||||
|
||||
Args:
|
||||
config (GaussianActorConfig): The policy configuration.
|
||||
"""
|
||||
super().__init__(config)
|
||||
config.validate_features()
|
||||
self.config = config
|
||||
@@ -53,6 +66,12 @@ class GaussianActorPolicy(
|
||||
self._init_discrete_critic()
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Returns only the `"actor"` parameter group, excluding the shared encoder's parameters when
|
||||
`shared_encoder` is enabled. The critic, encoder, and temperature parameters are optimized
|
||||
separately by the SAC algorithm.
|
||||
"""
|
||||
optim_params = {
|
||||
"actor": [
|
||||
p
|
||||
@@ -63,20 +82,30 @@ class GaussianActorPolicy(
|
||||
return optim_params
|
||||
|
||||
def reset(self):
|
||||
"""Reset the policy"""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. This policy holds no episode-scoped state,
|
||||
so this is a no-op.
|
||||
"""
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Predict a chunk of actions given environment observations."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Not supported: this policy predicts a single action per call rather than a chunk of actions, and
|
||||
calling this always raises `NotImplementedError`.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"GaussianActorPolicy does not support action chunking. It returns single actions!"
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Select action for inference/evaluation"""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
Samples one action directly from the actor network, re-using cached image features from the
|
||||
shared encoder when available, and appends an argmax discrete action (e.g. a gripper command)
|
||||
when `num_discrete_actions` is set.
|
||||
"""
|
||||
observations_features = None
|
||||
if self.shared_encoder and self.actor.encoder.has_images:
|
||||
observations_features = self.actor.encoder.get_cached_image_features(batch)
|
||||
@@ -96,15 +125,19 @@ class GaussianActorPolicy(
|
||||
return actions
|
||||
|
||||
def forward(self, batch: dict[str, Tensor | dict[str, Tensor]]) -> dict[str, Tensor]:
|
||||
"""Actor forward pass: sample actions and return log-probabilities.
|
||||
"""Actor forward pass: sample actions and return their log-probabilities.
|
||||
|
||||
Deviates from the base contract: rather than returning a training loss, this returns the actor's
|
||||
sampled actions, log-probabilities, and means directly. Loss computation and the Bellman update
|
||||
live on the algorithm side (see `lerobot.rl.algorithms.sac`).
|
||||
|
||||
Args:
|
||||
batch: A flat observation dict, or a training dict containing
|
||||
``"state"`` (observations) and optionally ``"observation_feature"``
|
||||
batch (dict[str, Tensor | dict[str, Tensor]]): A flat observation dict, or a training dict
|
||||
containing `"state"` (observations) and optionally `"observation_feature"`
|
||||
(pre-computed encoder features).
|
||||
|
||||
Returns:
|
||||
Dict with ``"action"``, ``"log_prob"``, and ``"action_mean"`` tensors.
|
||||
dict[str, Tensor]: Dict with `"action"`, `"log_prob"`, and `"action_mean"` tensors.
|
||||
"""
|
||||
observations = batch.get("state", batch)
|
||||
observation_features = batch.get("observation_feature") if isinstance(batch, dict) else None
|
||||
@@ -311,10 +344,10 @@ class MLP(nn.Module):
|
||||
Arguments:
|
||||
input_dim (int): Size of input feature dimension.
|
||||
hidden_dims (list[int]): Sizes for each hidden layer.
|
||||
activations (Callable or str): Activation to apply between layers.
|
||||
activate_final (bool): Whether to apply activation at the final layer.
|
||||
dropout_rate (Optional[float]): Dropout probability applied before normalization and activation.
|
||||
final_activation (Optional[Callable or str]): Activation for the final layer when `activate_final` is True.
|
||||
activations (Callable or str, *optional*, defaults to `SiLU()`): Activation to apply between layers.
|
||||
activate_final (bool, *optional*, defaults to `False`): Whether to apply activation at the final layer.
|
||||
dropout_rate (Optional[float], *optional*): Dropout probability applied before normalization and activation.
|
||||
final_activation (Optional[Callable or str], *optional*): Activation for the final layer when `activate_final` is True.
|
||||
|
||||
For each layer, `in_dim` is updated to the previous `out_dim`. All constructed modules are
|
||||
stored in `self.net` as an `nn.Sequential` container.
|
||||
@@ -562,8 +595,7 @@ def orthogonal_init():
|
||||
|
||||
class SpatialLearnedEmbeddings(nn.Module):
|
||||
def __init__(self, height, width, channel, num_features=8):
|
||||
"""
|
||||
PyTorch implementation of learned spatial embeddings
|
||||
"""PyTorch implementation of learned spatial embeddings
|
||||
|
||||
Args:
|
||||
height: Spatial height of input features
|
||||
@@ -582,8 +614,7 @@ class SpatialLearnedEmbeddings(nn.Module):
|
||||
nn.init.kaiming_normal_(self.kernel, mode="fan_in", nonlinearity="linear")
|
||||
|
||||
def forward(self, features):
|
||||
"""
|
||||
Forward pass for spatial embedding
|
||||
"""Forward pass for spatial embedding
|
||||
|
||||
Args:
|
||||
features: Input tensor of shape [B, C, H, W] where B is batch size,
|
||||
@@ -591,7 +622,6 @@ class SpatialLearnedEmbeddings(nn.Module):
|
||||
Returns:
|
||||
Output tensor of shape [B, C*F] where F is the number of features
|
||||
"""
|
||||
|
||||
features_expanded = features.unsqueeze(-1) # [B, C, H, W, 1]
|
||||
kernel_expanded = self.kernel.unsqueeze(0) # [1, C, H, W, F]
|
||||
|
||||
|
||||
@@ -35,8 +35,7 @@ def make_gaussian_actor_pre_post_processors(
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""
|
||||
Constructs pre-processor and post-processor pipelines for the Gaussian actor policy.
|
||||
"""Constructs pre-processor and post-processor pipelines for the Gaussian actor policy.
|
||||
|
||||
The pre-processing pipeline prepares input data for the model by:
|
||||
1. Renaming features to match pretrained configurations.
|
||||
@@ -49,8 +48,8 @@ def make_gaussian_actor_pre_post_processors(
|
||||
2. Unnormalizing the output features to their original scale.
|
||||
|
||||
Args:
|
||||
config: The configuration object for the tanh-Gaussian policy.
|
||||
dataset_stats: A dictionary of statistics for normalization.
|
||||
config (`GaussianActorConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
|
||||
@@ -74,6 +74,11 @@ _GROOT_ACTION_DECODE_TRANSFORM_ALIASES = {
|
||||
|
||||
|
||||
def normalize_groot_model_version(model_version: str) -> str:
|
||||
"""Resolve `model_version` to a canonical GR00T version string.
|
||||
|
||||
Raises:
|
||||
ValueError: If `model_version` isn't a recognized alias.
|
||||
"""
|
||||
normalized = _GROOT_MODEL_VERSION_ALIASES.get(model_version.lower())
|
||||
if normalized is None:
|
||||
supported = GROOT_N1_7
|
||||
@@ -85,6 +90,11 @@ def normalize_groot_model_version(model_version: str) -> str:
|
||||
|
||||
|
||||
def normalize_groot_action_decode_transform(transform: str | None) -> str | None:
|
||||
"""Resolve `transform` to a canonical action-decode-transform name, or `None`.
|
||||
|
||||
Raises:
|
||||
ValueError: If `transform` isn't a recognized alias.
|
||||
"""
|
||||
if transform is None:
|
||||
return None
|
||||
normalized = _GROOT_ACTION_DECODE_TRANSFORM_ALIASES.get(transform.lower())
|
||||
@@ -100,6 +110,7 @@ def normalize_groot_action_decode_transform(transform: str | None) -> str | None
|
||||
|
||||
|
||||
def infer_groot_model_version(model_path: str | None) -> str | None:
|
||||
"""Infer the GR00T model version (`GROOT_N1_7`) from a checkpoint path, or `None` if undetermined."""
|
||||
if not model_path:
|
||||
return None
|
||||
model_path_lower = model_path.lower()
|
||||
@@ -117,6 +128,7 @@ def infer_groot_model_version(model_path: str | None) -> str | None:
|
||||
|
||||
|
||||
def is_raw_groot_n1_7_checkpoint(model_path: str | Path | None) -> bool:
|
||||
"""Return `True` if `model_path` looks like an un-migrated, raw upstream GR00T N1.7 checkpoint."""
|
||||
if model_path is None:
|
||||
return False
|
||||
|
||||
@@ -133,6 +145,7 @@ def is_raw_groot_n1_7_checkpoint(model_path: str | Path | None) -> bool:
|
||||
|
||||
|
||||
def infer_groot_n1_7_embodiment_tag(model_path: str | Path | None) -> str | None:
|
||||
"""Infer the embodiment tag from a raw GR00T N1.7 checkpoint's `processor_config.json`, if resolvable."""
|
||||
if model_path is None:
|
||||
return None
|
||||
|
||||
@@ -152,6 +165,13 @@ def infer_groot_n1_7_embodiment_tag(model_path: str | Path | None) -> str | None
|
||||
def infer_groot_n1_7_action_horizon(
|
||||
model_path: str | Path | None, embodiment_tag: str | None = None
|
||||
) -> int | None:
|
||||
"""Infer the action horizon from a raw GR00T N1.7 checkpoint's `processor_config.json`, if resolvable.
|
||||
|
||||
Args:
|
||||
model_path (`str | pathlib.Path | None`): Path to the checkpoint directory.
|
||||
embodiment_tag (`str | None`, *optional*): The embodiment tag to look up. Inferred via
|
||||
`infer_groot_n1_7_embodiment_tag` when `None`.
|
||||
"""
|
||||
if model_path is None:
|
||||
return None
|
||||
|
||||
@@ -185,6 +205,13 @@ def infer_groot_n1_7_action_horizon(
|
||||
def infer_groot_n1_7_action_execution_horizon(
|
||||
model_path: str | Path | None, embodiment_tag: str | None = None
|
||||
) -> int | None:
|
||||
"""Infer the action execution horizon (<= action horizon) for a raw GR00T N1.7 checkpoint.
|
||||
|
||||
Args:
|
||||
model_path (`str | pathlib.Path | None`): Path to the checkpoint directory.
|
||||
embodiment_tag (`str | None`, *optional*): The embodiment tag to look up. Inferred via
|
||||
`infer_groot_n1_7_embodiment_tag` when `None`.
|
||||
"""
|
||||
action_horizon = infer_groot_n1_7_action_horizon(model_path, embodiment_tag)
|
||||
if action_horizon is None:
|
||||
return None
|
||||
@@ -241,7 +268,127 @@ def _infer_groot_model_version_from_config(config: dict) -> str | None:
|
||||
@PreTrainedConfig.register_subclass("groot")
|
||||
@dataclass
|
||||
class GrootConfig(PreTrainedConfig):
|
||||
"""Configuration for Groot policy wrapper."""
|
||||
"""Configuration for the GR00T N1.7 policy wrapper.
|
||||
|
||||
Wraps NVIDIA's Isaac-GR00T N1.7 model (a Qwen3-VL/Cosmos-Reason2 backbone plus a flow-matching
|
||||
action head) for fine-tuning and inference through LeRobot. GR00T N1.5 checkpoints and configs are
|
||||
no longer supported; loading one raises with `GROOT_N1_5_REMOVAL_GUIDANCE`.
|
||||
|
||||
Args:
|
||||
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||
pass to the policy (the current step plus this many additional steps looking back).
|
||||
input_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from input feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||
output_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from output feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||
chunk_size (`int`, *optional*, defaults to 40): The size of the action prediction chunk decoded
|
||||
per call to `predict_action_chunk`.
|
||||
n_action_steps (`int`, *optional*, defaults to 40): The number of actions from a predicted
|
||||
chunk that are actually queued for execution. Must not exceed `chunk_size`.
|
||||
max_state_dim (`int`, *optional*, defaults to 132): Maximum observation-state dimension expected
|
||||
by the pretrained GR00T model; shorter states are zero-padded.
|
||||
max_action_dim (`int`, *optional*, defaults to 132): Maximum action dimension expected by the
|
||||
pretrained GR00T model; shorter actions are zero-padded.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||
normalization mode. Always `IDENTITY` for every feature: GR00T normalizes state/action
|
||||
internally in its own processor steps and the Qwen3-VL image processor handles image
|
||||
normalization, so this mapping is not consulted by `make_groot_pre_post_processors`.
|
||||
base_model_path (`str | None`, *optional*): Path or Hub id of the base GR00T N1.7 model whose
|
||||
backbone weights and checkpoint sidecars (`statistics.json`, `processor_config.json`, ...)
|
||||
are loaded. Distinct from the inherited `pretrained_path`, which points at a saved LeRobot
|
||||
checkpoint directory. Defaults to `GROOT_N1_7_BASE_MODEL` when left unset.
|
||||
action_decode_transform (`str | None`, *optional*, defaults to `"auto"`): Named action transform
|
||||
applied after raw N1.7 checkpoint decoding and before `env.step()`. `"auto"` resolves to the
|
||||
embodiment default (`"libero"` for the `libero_sim` embodiment, otherwise no transform);
|
||||
pass `"none"` to explicitly disable it.
|
||||
embodiment_tag (`str`, *optional*, defaults to `"new_embodiment"`): Embodiment tag to use for
|
||||
training, e.g. `"new_embodiment"` or `"gr1"`.
|
||||
tune_llm (`bool`, *optional*, defaults to `False`): Whether to fine-tune the LLM backbone.
|
||||
tune_visual (`bool`, *optional*, defaults to `False`): Whether to fine-tune the vision tower.
|
||||
tune_projector (`bool`, *optional*, defaults to `True`): Whether to fine-tune the projector.
|
||||
tune_diffusion_model (`bool`, *optional*, defaults to `True`): Whether to fine-tune the
|
||||
flow-matching action head.
|
||||
tune_vlln (`bool`, *optional*, defaults to `True`): Whether to fine-tune the VL LayerNorm and VL
|
||||
self-attention projector in the action head.
|
||||
tune_top_llm_layers (`int`, *optional*, defaults to 0): Number of top LLM backbone layers to
|
||||
fine-tune (0 means none). Lets you adapt just the final language layers without unfreezing
|
||||
the whole backbone; independent of `tune_llm`, which tunes the entire LLM.
|
||||
num_inference_timesteps (`int | None`, *optional*): Number of flow-matching denoising steps used
|
||||
to decode an action chunk at inference time. `None` keeps the checkpoint value (GR00T N1.7
|
||||
default: 4).
|
||||
rtc_ramp_rate (`float | None`, *optional*): Real-Time Chunking overlap-blend ramp rate, used
|
||||
when the RTC engine supplies a previous-chunk prefix. `None` keeps the checkpoint value
|
||||
(GR00T N1.7 default: 6.0).
|
||||
use_flash_attention (`bool`, *optional*, defaults to `False`): Whether to request the
|
||||
flash-attention-2 kernel for the Qwen3-VL backbone. Set to `True` only after installing a
|
||||
flash-attn build matching your torch/CUDA environment; otherwise the backbone falls back to
|
||||
SDPA, which is numerically equivalent.
|
||||
use_relative_actions (`bool`, *optional*, defaults to `False`): Whether to enable GR00T-style
|
||||
state-relative action chunks (the action chunk is expressed relative to the current
|
||||
observation state).
|
||||
relative_exclude_joints (`list[str]`, *optional*): Action dimensions that stay absolute when
|
||||
`use_relative_actions` is set; matched as a case-insensitive substring against the dataset's
|
||||
action feature names. With the empty default every dimension is treated as relative,
|
||||
including the gripper; set e.g. `["gripper"]` to keep the gripper absolute.
|
||||
optimizer_lr (`float`, *optional*, defaults to 0.0001): Learning rate for the AdamW optimizer.
|
||||
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`): AdamW betas, as
|
||||
used by the Isaac-GR00T N1.7 fine-tuning recipe.
|
||||
optimizer_eps (`float`, *optional*, defaults to 1e-08): AdamW epsilon.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 1e-05): AdamW weight decay.
|
||||
warmup_ratio (`float`, *optional*, defaults to 0.05): Fraction of `max_steps` used as cosine
|
||||
scheduler warmup.
|
||||
use_bf16 (`bool`, *optional*, defaults to `True`): Whether to run the GR00T forward/inference
|
||||
passes under BF16 autocast.
|
||||
model_params_fp32 (`bool`, *optional*, defaults to `True`): Whether to keep model parameters in
|
||||
FP32 while computing under BF16 autocast, matching the native N1.7 fine-tuning recipe.
|
||||
image_size (`tuple[int, int]`, *optional*, defaults to `(256, 256)`): Legacy field kept only so
|
||||
that a GR00T N1.5-era `image_size=(224, 224)` config is detected and remapped to the N1.7
|
||||
default in `__post_init__`; image sizing is otherwise handled by the backbone's image
|
||||
processor.
|
||||
tokenizer_assets_repo (`str | None`, *optional*): Deprecated GR00T N1.5 field. Must stay `None`;
|
||||
a non-`None` value is treated as an N1.5 checkpoint/config and rejected in `__post_init__`.
|
||||
lora_rank (`int`, *optional*, defaults to 0): Deprecated, never-wired LoRA field kept only so
|
||||
older saved configs still parse.
|
||||
lora_alpha (`int`, *optional*, defaults to 16): Deprecated, never-wired LoRA field kept only so
|
||||
older saved configs still parse.
|
||||
lora_dropout (`float`, *optional*, defaults to 0.1): Deprecated, never-wired LoRA field kept only
|
||||
so older saved configs still parse.
|
||||
lora_full_model (`bool`, *optional*, defaults to `False`): Deprecated, never-wired LoRA field
|
||||
kept only so older saved configs still parse.
|
||||
video_backend (`str`, *optional*, defaults to `"decord"`): Deprecated Isaac-GR00T runner field;
|
||||
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
balance_dataset_weights (`bool`, *optional*, defaults to `True`): Deprecated Isaac-GR00T runner
|
||||
field; unused by the LeRobot N1.7 implementation, kept only so older saved configs still
|
||||
parse.
|
||||
balance_trajectory_weights (`bool`, *optional*, defaults to `True`): Deprecated Isaac-GR00T
|
||||
runner field; unused by the LeRobot N1.7 implementation, kept only so older saved configs
|
||||
still parse.
|
||||
dataset_paths (`list[str] | None`, *optional*): Deprecated Isaac-GR00T runner field; unused by
|
||||
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
output_dir (`str`, *optional*, defaults to `"./tmp/gr00t"`): Deprecated Isaac-GR00T runner field;
|
||||
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
save_steps (`int`, *optional*, defaults to 1000): Deprecated Isaac-GR00T runner field; unused by
|
||||
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
max_steps (`int`, *optional*, defaults to 10000): Total training steps; used together with
|
||||
`warmup_ratio` to derive the cosine scheduler's warmup step count in
|
||||
`get_scheduler_preset`.
|
||||
batch_size (`int`, *optional*, defaults to 32): Deprecated Isaac-GR00T runner field; unused by
|
||||
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
dataloader_num_workers (`int`, *optional*, defaults to 8): Deprecated Isaac-GR00T runner field;
|
||||
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
report_to (`str`, *optional*, defaults to `"wandb"`): Deprecated Isaac-GR00T runner field; unused
|
||||
by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
resume (`bool`, *optional*, defaults to `False`): Deprecated Isaac-GR00T runner field; unused by
|
||||
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||
"""
|
||||
|
||||
# Basic policy settings
|
||||
n_obs_steps: int = 1
|
||||
@@ -372,6 +519,12 @@ class GrootConfig(PreTrainedConfig):
|
||||
resume: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
"""Reject legacy GR00T N1.5 configs, normalize fields, and remap N1.5-era defaults.
|
||||
|
||||
Raises:
|
||||
ValueError: If `tokenizer_assets_repo` is set (an N1.5-only field), if `base_model_path`
|
||||
resolves to a GR00T N1.5 checkpoint, or if `n_action_steps` exceeds `chunk_size`.
|
||||
"""
|
||||
if self.tokenizer_assets_repo is not None:
|
||||
raise ValueError(
|
||||
"Config sets 'tokenizer_assets_repo', which only existed for GR00T N1.5; this looks "
|
||||
|
||||
@@ -14,8 +14,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""
|
||||
Groot Policy Wrapper for LeRobot Integration
|
||||
"""Groot Policy Wrapper for LeRobot Integration
|
||||
|
||||
Minimal integration that delegates to Isaac-GR00T N1.7 components where
|
||||
possible without porting their code. Dataset loading and training
|
||||
@@ -69,10 +68,17 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
config_class = GrootConfig
|
||||
|
||||
def supports_rtc(self) -> bool:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. GR00T N1.7 implements RTC."""
|
||||
return True
|
||||
|
||||
def __init__(self, config: GrootConfig, **kwargs):
|
||||
"""Initialize Groot policy wrapper."""
|
||||
"""Build the underlying GR00T N1.7 model from `config` and reset the action queue.
|
||||
|
||||
Args:
|
||||
config (GrootConfig): Policy configuration; also validated/completed via
|
||||
`config.validate_features()`.
|
||||
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||
"""
|
||||
require_package("transformers", extra="groot")
|
||||
super().__init__(config)
|
||||
config.validate_features()
|
||||
@@ -149,7 +155,7 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
]
|
||||
|
||||
def reset(self):
|
||||
"""Reset policy state when environment resets."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue."""
|
||||
self._action_queue = deque([], maxlen=self._action_queue_steps)
|
||||
|
||||
@classmethod
|
||||
@@ -168,27 +174,40 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
strict: bool = True,
|
||||
**kwargs,
|
||||
) -> T:
|
||||
"""Load Groot policy from pretrained model.
|
||||
"""Load a Groot policy from either a raw N1.7 checkpoint or a fine-tuned LeRobot checkpoint.
|
||||
|
||||
Handles two cases:
|
||||
1. Base GR00T N1.7 models - loads the raw model
|
||||
2. Fine-tuned LeRobot checkpoints - loads config and weights from safetensors
|
||||
|
||||
Args:
|
||||
pretrained_name_or_path: Path to the GR00T model or fine-tuned checkpoint
|
||||
config: Optional GrootConfig. If None, loads from checkpoint or creates default
|
||||
force_download: Force download even if cached
|
||||
resume_download: Resume interrupted download
|
||||
proxies: Proxy settings
|
||||
token: HuggingFace authentication token
|
||||
cache_dir: Cache directory path
|
||||
local_files_only: Only use local files
|
||||
revision: Specific model revision
|
||||
strict: Strict state dict loading
|
||||
**kwargs: Additional arguments (passed to config)
|
||||
pretrained_name_or_path (str | Path): Hub id or local path to the GR00T model or the
|
||||
fine-tuned checkpoint.
|
||||
config (GrootConfig | None, *optional*): Config to use. If `None`, one is loaded from the
|
||||
checkpoint (fine-tuned case) or created with defaults (base-model case).
|
||||
force_download (bool, *optional*, defaults to `False`): Whether to force (re-)downloading
|
||||
the files, overriding the existing cache.
|
||||
resume_download (bool | None, *optional*): Deprecated; ignored by the underlying Hub client.
|
||||
proxies (dict | None, *optional*): A dictionary of proxy servers to use by protocol or
|
||||
endpoint.
|
||||
token (str | bool | None, *optional*): The token to use as HTTP bearer authorization for
|
||||
remote files.
|
||||
cache_dir (str | Path | None, *optional*): Path to the folder where cached files are stored.
|
||||
local_files_only (bool, *optional*, defaults to `False`): If `True`, avoid downloading the
|
||||
file and use the local cache only.
|
||||
revision (str | None, *optional*): Revision on the Hub: a branch name, git tag, or commit id.
|
||||
strict (bool, *optional*, defaults to `True`): Whether to require an exact match between the
|
||||
checkpoint's and the instantiated model's parameter keys.
|
||||
kwargs: For the fine-tuned-checkpoint case, forwarded to
|
||||
[`~policies.pretrained.PreTrainedPolicy.from_pretrained`]. For the base-model case,
|
||||
applied as config field overrides.
|
||||
|
||||
Returns:
|
||||
Initialized GrootPolicy instance with loaded model
|
||||
T: The loaded `GrootPolicy` instance, in eval mode.
|
||||
|
||||
Raises:
|
||||
ValueError: If `config.base_model_path` (or `pretrained_name_or_path`) resolves to an
|
||||
unsupported GR00T model version.
|
||||
"""
|
||||
requested_version = infer_groot_model_version(str(pretrained_name_or_path)) or GROOT_N1_7
|
||||
logger.info(
|
||||
@@ -285,7 +304,11 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
return policy
|
||||
|
||||
def get_optim_params(self): # type: ignore[override]
|
||||
"""Isaac-GR00T excludes biases and normalization parameters from weight decay."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Splits parameters into weight-decay and no-weight-decay groups, matching the Isaac-GR00T
|
||||
recipe of excluding biases and normalization parameters from weight decay.
|
||||
"""
|
||||
return self._build_weight_decay_parameter_groups(self)
|
||||
|
||||
def _resolve_action_queue_steps(self) -> int:
|
||||
@@ -307,7 +330,6 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
|
||||
def _resolve_prediction_horizon(self, actions: Tensor) -> int:
|
||||
"""Return the policy-facing action horizon for a native GR00T prediction."""
|
||||
|
||||
horizons = [actions.shape[1]]
|
||||
checkpoint_action_horizon = infer_groot_n1_7_action_horizon(
|
||||
self.config.base_model_path,
|
||||
@@ -444,9 +466,10 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
return inputs, options
|
||||
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||
"""Training forward pass.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
Delegates to Isaac-GR00T model.forward when inputs are compatible.
|
||||
Delegates to the underlying Isaac-GR00T model's `forward`, run under BF16 autocast when
|
||||
`config.use_bf16` is set.
|
||||
"""
|
||||
groot_inputs = self._filter_groot_inputs(batch, include_action=True)
|
||||
|
||||
@@ -472,12 +495,11 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: object) -> Tensor:
|
||||
"""Predict a chunk of actions for inference by delegating to Isaac-GR00T.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
Returns a tensor of shape (B, n_action_steps, action_dim).
|
||||
|
||||
For N1.7, LeRobot's RTC leftovers are converted into the native GR00T
|
||||
action-overlap options before calling the underlying model.
|
||||
Delegates to the underlying Isaac-GR00T model's `get_action`, returning a tensor of shape
|
||||
`(B, n_action_steps, action_dim)`. LeRobot's RTC leftovers, if any, are converted into the
|
||||
native GR00T action-overlap options before calling the model.
|
||||
"""
|
||||
self.eval()
|
||||
|
||||
@@ -513,7 +535,15 @@ class GrootPolicy(PreTrainedPolicy):
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Select single action from action queue."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
Uses an action queue populated by `predict_action_chunk`.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If `config.use_relative_actions` is set, since cached relative-chunk
|
||||
actions can be decoded against newer observation states; use `predict_action_chunk`
|
||||
directly instead.
|
||||
"""
|
||||
if getattr(self.config, "use_relative_actions", False):
|
||||
raise NotImplementedError(
|
||||
"GrootPolicy.select_action does not support relative-action policies because cached "
|
||||
|
||||
@@ -165,7 +165,6 @@ def _load_n1_7_checkpoint_processor_assets(config: GrootConfig) -> _GrootN17Chec
|
||||
Returns ``None`` for non-raw N1.7 checkpoints so the generic GR00T pipeline
|
||||
can keep using caller-provided dataset stats and config values.
|
||||
"""
|
||||
|
||||
if not is_raw_groot_n1_7_checkpoint(config.base_model_path):
|
||||
return None
|
||||
|
||||
@@ -273,7 +272,6 @@ def _load_n1_7_checkpoint_stats(
|
||||
joints. LeRobot normalizers operate over a single vector, so this function
|
||||
preserves checkpoint group order while flattening each selected statistic.
|
||||
"""
|
||||
|
||||
if raw_stats is None:
|
||||
all_stats = read_json(checkpoint_path / "statistics.json")
|
||||
raw_stats = all_stats.get(embodiment_tag)
|
||||
@@ -381,7 +379,6 @@ _GROOT_ABSENT_STANDARD_OVERRIDE_KEYS = frozenset(
|
||||
|
||||
def _drop_groot_absent_standard_overrides(overrides: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""Strip standard override keys that a GR00T pipeline has no step for."""
|
||||
|
||||
if not overrides:
|
||||
return overrides
|
||||
|
||||
@@ -414,7 +411,6 @@ def _apply_groot_step_overrides(
|
||||
silently (standard normalization keys GR00T has no step for are removed
|
||||
beforehand by ``_drop_groot_absent_standard_overrides``).
|
||||
"""
|
||||
|
||||
if not overrides:
|
||||
return
|
||||
|
||||
@@ -487,7 +483,6 @@ def make_groot_pre_post_processors_from_pretrained(
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""Load Groot processors for a raw N1.7 checkpoint or a serialized LeRobot pipeline."""
|
||||
|
||||
# Drop the standard normalizer/unnormalizer override keys lerobot-train emits unconditionally:
|
||||
# GR00T has no such steps, so they would make both the raw-checkpoint and serialized override
|
||||
# paths raise. This must happen before either branch below.
|
||||
@@ -584,7 +579,6 @@ def _reconnect_groot_n1_7_pack_decode_steps(
|
||||
The pack step holds the per-instance raw-state cache that relative-action
|
||||
decoding reads its reference state from; the link itself is not serialized.
|
||||
"""
|
||||
|
||||
pack_step = next(
|
||||
(step for step in preprocessor.steps if isinstance(step, GrootN17PackInputsStep)),
|
||||
None,
|
||||
@@ -1155,13 +1149,13 @@ def make_groot_pre_post_processors(
|
||||
This mirrors SO100-style preprocessing and keeps scales consistent with GR00T.
|
||||
|
||||
Args:
|
||||
config: Groot configuration containing data_config, embodiment_tag, etc.
|
||||
dataset_stats: Optional per-key min/max statistics for normalization before padding.
|
||||
config (`GrootConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||
dataset_meta (`typing.Any | None`, *optional*): Dataset metadata, forwarded to factories that need more than just `dataset_stats`.
|
||||
|
||||
Returns:
|
||||
Tuple of (preprocessor, postprocessor) pipelines
|
||||
"""
|
||||
|
||||
dataset_meta = dataset_meta or getattr(config, "_runtime_dataset_meta", None)
|
||||
checkpoint_assets = _load_n1_7_checkpoint_processor_assets(config)
|
||||
checkpoint_stats = checkpoint_assets.stats if checkpoint_assets is not None else None
|
||||
@@ -1354,7 +1348,6 @@ def _to_uint8_np_bthwc(img_t: torch.Tensor) -> np.ndarray:
|
||||
|
||||
def _align_video_horizon(video: np.ndarray, horizon: int | None) -> np.ndarray:
|
||||
"""Match the checkpoint video horizon by truncating or left-padding frames."""
|
||||
|
||||
if horizon is None or horizon <= 0:
|
||||
return video
|
||||
current = video.shape[1]
|
||||
@@ -2010,7 +2003,6 @@ class GrootN17PackInputsStep(ProcessorStep):
|
||||
|
||||
def get_cached_raw_state(self) -> dict[str, np.ndarray] | None:
|
||||
"""Return the latest unnormalized state split by checkpoint modality key."""
|
||||
|
||||
return self._last_raw_state
|
||||
|
||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||
@@ -2225,7 +2217,6 @@ def _n1_7_decode_stats_for_action(
|
||||
use_percentiles: bool,
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Select the min/max arrays needed to decode one checkpoint action group."""
|
||||
|
||||
is_relative = use_relative_action and config_value(action_config.get("rep")) == "relative"
|
||||
modality = "relative_action" if is_relative else "action"
|
||||
stats = raw_stats.get(modality, {}).get(key, {})
|
||||
@@ -2524,8 +2515,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
return features
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
"""
|
||||
Returns a serializable dictionary of the processor's configuration.
|
||||
"""Returns a serializable dictionary of the processor's configuration.
|
||||
|
||||
Excludes 'stats' since they are saved separately via state_dict().
|
||||
"""
|
||||
@@ -2538,8 +2528,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
}
|
||||
|
||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||
"""
|
||||
Returns normalization statistics as a flat state dictionary.
|
||||
"""Returns normalization statistics as a flat state dictionary.
|
||||
|
||||
This enables saving stats to safetensors files, similar to normalizer_processor.
|
||||
"""
|
||||
@@ -2554,8 +2543,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
return flat
|
||||
|
||||
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||
"""
|
||||
Loads normalization statistics from a flat state dictionary.
|
||||
"""Loads normalization statistics from a flat state dictionary.
|
||||
|
||||
This enables loading stats from safetensors files during from_pretrained.
|
||||
"""
|
||||
|
||||
@@ -28,18 +28,110 @@ from dataclasses import dataclass, field
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature
|
||||
from lerobot.optim.optimizers import AdamWConfig
|
||||
from lerobot.optim.schedulers import (
|
||||
ConstantWithWarmupSchedulerConfig,
|
||||
CosineAnnealingWithWarmupSchedulerConfig,
|
||||
LRSchedulerConfig,
|
||||
)
|
||||
from lerobot.optim.schedulers import ConstantWithWarmupSchedulerConfig, LRSchedulerConfig
|
||||
from lerobot.utils.constants import ACTION
|
||||
|
||||
|
||||
@PreTrainedConfig.register_subclass("lingbot_va")
|
||||
@dataclass
|
||||
class LingBotVAConfig(PreTrainedConfig):
|
||||
"""Configuration for the native LingBot-VA policy integration in LeRobot."""
|
||||
"""Configuration for the native LingBot-VA policy integration in LeRobot.
|
||||
|
||||
Defaults match the upstream LIBERO configuration (`wan_va/configs/va_libero_cfg.py`) and the
|
||||
`transformer/config.json` of the released checkpoints.
|
||||
|
||||
Args:
|
||||
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||
pass to the policy (the current step plus this many additional steps looking back).
|
||||
input_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from input feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||
output_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from output feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||
patch_size (`tuple[int, int, int]`, *optional*, defaults to `(1, 2, 2)`): Wan transformer's
|
||||
spatiotemporal patch size (time, height, width).
|
||||
num_attention_heads (`int`, *optional*, defaults to 24): Number of attention heads in the Wan
|
||||
transformer.
|
||||
attention_head_dim (`int`, *optional*, defaults to 128): Dimension per attention head.
|
||||
in_channels (`int`, *optional*, defaults to 48): Number of input channels to the transformer
|
||||
(VAE latent channels).
|
||||
out_channels (`int`, *optional*, defaults to 48): Number of output channels from the
|
||||
transformer.
|
||||
action_dim (`int`, *optional*, defaults to 30): Dimension of the action stream fed to and
|
||||
predicted by the transformer.
|
||||
text_dim (`int`, *optional*, defaults to 4096): Dimension of the UMT5 text embeddings.
|
||||
freq_dim (`int`, *optional*, defaults to 256): Dimension of the sinusoidal timestep embedding.
|
||||
ffn_dim (`int`, *optional*, defaults to 14336): Hidden dimension of the transformer's
|
||||
feed-forward blocks.
|
||||
num_layers (`int`, *optional*, defaults to 30): Number of transformer layers.
|
||||
cross_attn_norm (`bool`, *optional*, defaults to `True`): Whether to normalize the
|
||||
cross-attention inputs.
|
||||
eps (`float`, *optional*, defaults to 1e-06): Epsilon used in the transformer's normalization
|
||||
layers.
|
||||
rope_max_seq_len (`int`, *optional*, defaults to 1024): Maximum sequence length for the
|
||||
transformer's rotary position embeddings.
|
||||
attn_mode (`str`, *optional*, defaults to `"torch"`): Attention backend. `"torch"` (SDPA) or
|
||||
`"flashattn"` for inference; `"flex"` for training only, and only on a recent torch.
|
||||
wan_pretrained_path (`str`, *optional*, defaults to `"robbyant/lingbot-va-base"`): Hub id or
|
||||
local directory holding the frozen VAE, UMT5 text encoder, and tokenizer sub-folders
|
||||
(diffusers layout, ~20 GB). Lazily loaded and not bundled in the checkpoint.
|
||||
dtype (`str`, *optional*, defaults to `"bfloat16"`): Transformer/VAE/text-encoder dtype:
|
||||
`"bfloat16"`, `"float16"`, or `"float32"`.
|
||||
text_encoder_device (`str`, *optional*, defaults to `"cpu"`): Device for the frozen UMT5-XXL
|
||||
text encoder, which runs once per episode. `"cpu"` frees ~11 GB of VRAM.
|
||||
obs_cam_keys (`list[str]`, *optional*): Observation camera keys, in concatenation order (order
|
||||
matters: latents are concatenated on width). Defaults to the LIBERO camera keys.
|
||||
image_hflip (`bool`, *optional*, defaults to `False`): Whether to undo the LIBERO env
|
||||
processor's extra horizontal flip, to match the model's training orientation.
|
||||
camera_layout (`str`, *optional*, defaults to `"width_concat"`): Camera latent layout:
|
||||
`"width_concat"` (cameras concatenated on width; LIBERO) or `"robotwin_tshape"` (full-res
|
||||
head plus half-res wrists in a "T"; RoboTwin).
|
||||
height (`int`, *optional*, defaults to 128): Observation image height fed to the VAE.
|
||||
width (`int`, *optional*, defaults to 128): Observation image width fed to the VAE.
|
||||
action_per_frame (`int`, *optional*, defaults to 4): Number of single-step actions decoded per
|
||||
predicted video frame.
|
||||
frame_chunk_size (`int`, *optional*, defaults to 4): Number of video frames predicted per
|
||||
autoregressive chunk.
|
||||
attn_window (`int`, *optional*, defaults to 30): Attention window size, in frames, for the
|
||||
causal streaming KV cache.
|
||||
num_inference_steps (`int`, *optional*, defaults to 20): Number of denoising steps for the
|
||||
video-latent flow-matching scheduler.
|
||||
video_exec_step (`int`, *optional*, defaults to -1): Which decoded video frame index to treat
|
||||
as "executed" for KV-cache feedback. `-1` uses the last frame.
|
||||
action_num_inference_steps (`int`, *optional*, defaults to 50): Number of denoising steps for
|
||||
the action flow-matching scheduler.
|
||||
guidance_scale (`float`, *optional*, defaults to 5.0): Classifier-free guidance scale for the
|
||||
video-latent stream.
|
||||
action_guidance_scale (`float`, *optional*, defaults to 1.0): Classifier-free guidance scale
|
||||
for the action stream.
|
||||
snr_shift (`float`, *optional*, defaults to 5.0): Flow-matching noise-schedule shift for the
|
||||
video-latent stream.
|
||||
action_snr_shift (`float`, *optional*, defaults to 0.05): Flow-matching noise-schedule shift
|
||||
for the action stream.
|
||||
max_sequence_length (`int`, *optional*, defaults to 512): Maximum UMT5 prompt length.
|
||||
used_action_channel_ids (`list[int]`, *optional*): Subset of the 30-d action space used by the
|
||||
benchmark; defaults to the first 7 channels (LIBERO's 7-DoF action). The action
|
||||
(un)normalization quantiles live in the checkpoint's `policy_postprocessor.json`, not here.
|
||||
save_predicted_video (`bool`, *optional*, defaults to `False`): Whether to VAE-decode predicted
|
||||
video latents into `self.last_predicted_frames`, opt-in for saving MP4s.
|
||||
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||
normalization mode. Always `IDENTITY`: images are scaled and VAE-encoded, and actions are
|
||||
quantile-(un)normalized, inside the policy or a dedicated processor step.
|
||||
optimizer_lr (`float`, *optional*, defaults to 1e-05): AdamW learning rate.
|
||||
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`): AdamW betas.
|
||||
optimizer_eps (`float`, *optional*, defaults to 1e-08): AdamW epsilon.
|
||||
optimizer_weight_decay (`float`, *optional*, defaults to 0.0001): AdamW weight decay.
|
||||
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0): Gradient clipping norm.
|
||||
scheduler_warmup_steps (`int`, *optional*, defaults to 1000): Number of linear-warmup steps
|
||||
before the constant learning-rate phase.
|
||||
"""
|
||||
|
||||
# Wan transformer architecture
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
@@ -96,15 +188,6 @@ class LingBotVAConfig(PreTrainedConfig):
|
||||
# (un)normalization quantiles live in the checkpoint's ``policy_postprocessor.json``, not here.
|
||||
used_action_channel_ids: list[int] = field(default_factory=lambda: list(range(7)))
|
||||
|
||||
# Relative actions: converts absolute actions to relative (action -= state) during
|
||||
# preprocessing, and reverses it at postprocessing. Requires the dataset to provide
|
||||
# observation.state whose leading dims align 1:1 with the used action channels.
|
||||
use_relative_actions: bool = False
|
||||
# Joint names to keep absolute (not converted to relative). Empty list = all dims relative.
|
||||
relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
|
||||
# Populated at runtime from dataset metadata by make_policy (used to build the exclude mask).
|
||||
action_feature_names: list[str] | None = None
|
||||
|
||||
# Opt-in: VAE-decode predicted video latents to ``self.last_predicted_frames`` for saving MP4s.
|
||||
save_predicted_video: bool = False
|
||||
|
||||
@@ -125,19 +208,13 @@ class LingBotVAConfig(PreTrainedConfig):
|
||||
optimizer_weight_decay: float = 1e-4
|
||||
optimizer_grad_clip_norm: float = 1.0
|
||||
scheduler_warmup_steps: int = 1000
|
||||
# Scheduler after warmup. "constant_with_warmup" (upstream default: warmup then flat peak LR)
|
||||
# or "cosine_annealing_with_warmup" (warmup then cosine anneal peak->0 over the remaining steps).
|
||||
# Cosine tightens the loss tail and often nudges final loss down; it does NOT reduce the
|
||||
# flow-matching estimator's step-to-step noise (that's metric variance, LR-independent).
|
||||
scheduler_type: str = "constant_with_warmup"
|
||||
# Probability of corrupting the action stream's conditioning (clean/context) tokens with
|
||||
# flow-matching noise during training, mirroring the video stream's noisy_cond_prob=0.5.
|
||||
# Upstream train.py hardcodes 0.0 for actions (never corrupted) with no exposed knob; this is
|
||||
# an experimental deviation to make the model more tolerant of imperfect action history
|
||||
# (e.g. clamp-induced drift between predicted and executed actions during rollout).
|
||||
action_noisy_cond_prob: float = 0.0
|
||||
|
||||
def __post_init__(self):
|
||||
"""Validate `attn_mode`.
|
||||
|
||||
Raises:
|
||||
ValueError: If `attn_mode` is not one of `"torch"`, `"flashattn"`, or `"flex"`.
|
||||
"""
|
||||
super().__post_init__()
|
||||
if self.attn_mode not in ("torch", "flashattn", "flex"):
|
||||
raise ValueError(f"attn_mode must be one of 'torch', 'flashattn', 'flex'; got {self.attn_mode!r}")
|
||||
@@ -153,6 +230,11 @@ class LingBotVAConfig(PreTrainedConfig):
|
||||
return self.chunk_size
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""Validate and set up input/output features for LingBot-VA.
|
||||
|
||||
Raises:
|
||||
ValueError: If no visual input feature is present in `input_features`.
|
||||
"""
|
||||
image_features = [key for key, feat in self.input_features.items() if feat.type == FeatureType.VISUAL]
|
||||
if not image_features:
|
||||
raise ValueError(
|
||||
@@ -165,6 +247,7 @@ class LingBotVAConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
"""Return the AdamW optimizer configuration built from the `optimizer_*` fields."""
|
||||
return AdamWConfig(
|
||||
lr=self.optimizer_lr,
|
||||
betas=self.optimizer_betas,
|
||||
@@ -174,37 +257,23 @@ class LingBotVAConfig(PreTrainedConfig):
|
||||
)
|
||||
|
||||
def get_scheduler_preset(self) -> LRSchedulerConfig | None:
|
||||
# Default (upstream): linear warmup then constant LR (warmup_constant_lambda).
|
||||
# Optionally cosine-anneal peak->0 over the remaining steps via scheduler_type.
|
||||
if self.scheduler_type == "cosine_annealing_with_warmup":
|
||||
return CosineAnnealingWithWarmupSchedulerConfig(num_warmup_steps=self.scheduler_warmup_steps)
|
||||
"""Return the linear-warmup-then-constant scheduler configuration, matching upstream's `warmup_constant_lambda`."""
|
||||
# Upstream uses a linear warmup followed by a constant LR (warmup_constant_lambda).
|
||||
return ConstantWithWarmupSchedulerConfig(num_warmup_steps=self.scheduler_warmup_steps)
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list[int]:
|
||||
"""Observation frame deltas for the training clip, sized to what the VAE actually reads.
|
||||
|
||||
``diffusers``' ``AutoencoderKLWan._encode`` runs ``iter_ = 1 + (n - 1) // 4`` passes over
|
||||
``x[:, :, :1]`` then ``x[:, :, 1 + 4*(i-1) : 1 + 4*i]``, so it only ever consumes the first
|
||||
``4 * (iter_ - 1) + 1`` frames of an ``n``-frame clip. Asking for ``frame_chunk_size * 4``
|
||||
frames (the previous formula) therefore decoded 3 frames per sample that never reached the
|
||||
encoder: at ``frame_chunk_size=2`` the deltas were ``[0, 4, ..., 28]`` and only
|
||||
``[0, 4, 8, 12, 16]`` were used -- verified by ablation, scrambling the tail left the latents
|
||||
bit-identical.
|
||||
|
||||
Requesting exactly ``4 * (frame_chunk_size - 1) + 1`` frames yields the same
|
||||
``frame_chunk_size`` latent frames with every loaded frame used, and drops the wasted video
|
||||
decode. The stride is unchanged, so the frames that do reach the model are the same ones.
|
||||
"""
|
||||
"""Return the keyframe-sampling indices used to build the observed-frame history."""
|
||||
temporal_downsample = 4
|
||||
stride = max(1, self.action_per_frame // temporal_downsample)
|
||||
num_frames = temporal_downsample * (self.frame_chunk_size - 1) + 1
|
||||
return [i * stride for i in range(num_frames)]
|
||||
return list(range(0, self.frame_chunk_size * temporal_downsample * stride, stride))
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list[int]:
|
||||
"""Return indices for delta actions."""
|
||||
return list(range(self.chunk_size))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
"""Return indices for delta rewards (None for LingBot-VA)."""
|
||||
return None
|
||||
|
||||
@@ -38,7 +38,7 @@ import torch.nn.functional as F # noqa: N812
|
||||
from einops import rearrange
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.policies.pretrained import PreTrainedPolicy, unpack_action_output
|
||||
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||
from lerobot.utils.constants import ACTION
|
||||
from lerobot.utils.import_utils import require_package
|
||||
|
||||
@@ -66,6 +66,17 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
name = "lingbot_va"
|
||||
|
||||
def __init__(self, config: LingBotVAConfig, **kwargs):
|
||||
"""Build the trainable Wan dual-stream transformer and reset per-episode streaming state.
|
||||
|
||||
The VAE, UMT5 text encoder, and tokenizer are frozen and lazily loaded from
|
||||
`config.wan_pretrained_path` on first use; only the transformer is saved in the LeRobot
|
||||
checkpoint.
|
||||
|
||||
Args:
|
||||
config (LingBotVAConfig): Policy configuration; also validated/completed via
|
||||
`config.validate_features()`.
|
||||
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||
"""
|
||||
require_package("diffusers", extra="lingbot_va")
|
||||
require_package("transformers", extra="lingbot_va")
|
||||
super().__init__(config)
|
||||
@@ -99,6 +110,8 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
# from ``config.wan_pretrained_path`` the first time inference runs.
|
||||
self._frozen: dict = {}
|
||||
|
||||
self.last_predicted_frames: Tensor | None = None
|
||||
self.last_predicted_latents: Tensor | None = None
|
||||
self.reset()
|
||||
|
||||
# Frozen-module lazy loading (VAE + UMT5 + tokenizer)
|
||||
@@ -144,12 +157,18 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
|
||||
# PreTrainedPolicy API
|
||||
def get_optim_params(self) -> dict:
|
||||
# Only the transformer is trainable; the VAE / text encoder stay frozen (kept outside the
|
||||
# nn.Module registry). With PEFT/LoRA this naturally returns just the adapter params.
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||
|
||||
Only the transformer is trainable; the VAE and text encoder stay frozen (kept outside the
|
||||
`nn.Module` registry). With PEFT/LoRA this naturally returns just the adapter params.
|
||||
"""
|
||||
return [p for p in self.transformer.parameters() if p.requires_grad]
|
||||
|
||||
def reset(self):
|
||||
"""Reset all per-episode streaming state (KV cache, queues, frame counter)."""
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||
|
||||
Resets all per-episode streaming state (KV cache, queues, frame counter).
|
||||
"""
|
||||
cfg = self.config
|
||||
self._action_queue: deque = deque(maxlen=cfg.n_action_steps)
|
||||
self._obs_buffer: list = [] # raw keyframe obs (one per env substep) observed this chunk
|
||||
@@ -168,6 +187,8 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
self._prompt: str | None = None
|
||||
self._prompt_embeds = None
|
||||
self._negative_prompt_embeds = None
|
||||
self.last_predicted_frames = None
|
||||
self.last_predicted_latents = None
|
||||
self._use_cfg = (cfg.guidance_scale > 1) or (cfg.action_guidance_scale > 1)
|
||||
# Two independent flow-matching schedulers (video latent + action streams).
|
||||
self._scheduler = FlowMatchScheduler(shift=cfg.snr_shift, sigma_min=0.0, extra_one_step=True)
|
||||
@@ -253,12 +274,8 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
"grid_id": grid_id,
|
||||
}
|
||||
|
||||
def _flow_matching_loss(self, input_dict, pred, reduction: str = "mean"):
|
||||
"""Dual-stream flow-matching loss (port of upstream ``Trainer.compute_loss``).
|
||||
|
||||
``reduction="mean"`` returns scalar (latent_loss, action_loss); ``"none"`` returns
|
||||
per-sample vectors of shape ``(B,)`` each (averaged over latent frames) for RA-BC.
|
||||
"""
|
||||
def _flow_matching_loss(self, input_dict, pred):
|
||||
"""Dual-stream flow-matching loss (port of upstream ``Trainer.compute_loss``)."""
|
||||
latent_pred, action_pred = pred
|
||||
ld, ad = input_dict["latent_dict"], input_dict["action_dict"]
|
||||
action_pred = rearrange(action_pred, "b (f n) c -> b c f n 1", f=ad["targets"].shape[-3])
|
||||
@@ -278,8 +295,7 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
latent_loss = (
|
||||
(latent_loss * lw[:, None, :, None, None]).permute(0, 2, 3, 4, 1).flatten(0, 1).flatten(1)
|
||||
)
|
||||
# per (batch*frame) mean over spatial/channel -> (B*F,)
|
||||
latent_loss = latent_loss.sum(dim=1) / (torch.ones_like(latent_loss).sum(dim=1) + 1e-6)
|
||||
latent_loss = (latent_loss.sum(dim=1) / (torch.ones_like(latent_loss).sum(dim=1) + 1e-6)).mean()
|
||||
|
||||
amask = ad["actions_mask"].float()
|
||||
action_loss = F.mse_loss(action_pred.float(), ad["targets"].float().detach(), reduction="none")
|
||||
@@ -287,14 +303,10 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
(action_loss * aw[:, None, :, None, None] * amask).permute(0, 2, 3, 4, 1).flatten(0, 1).flatten(1)
|
||||
)
|
||||
amask_f = amask.permute(0, 2, 3, 4, 1).flatten(0, 1).flatten(1)
|
||||
action_loss = action_loss.sum(dim=1) / (amask_f.sum(dim=1) + 1e-6)
|
||||
action_loss = (action_loss.sum(dim=1) / (amask_f.sum(dim=1) + 1e-6)).mean()
|
||||
return latent_loss, action_loss
|
||||
|
||||
if reduction == "none":
|
||||
# (B*F,) -> (B, F) -> (B,): per-sample losses for RA-BC weighting.
|
||||
return latent_loss.reshape(bn, fn).mean(dim=1), action_loss.reshape(bn, fn).mean(dim=1)
|
||||
return latent_loss.mean(), action_loss.mean()
|
||||
|
||||
def training_loss_from_streams(self, latents, actions, actions_mask, text_emb, reduction: str = "mean"):
|
||||
def training_loss_from_streams(self, latents, actions, actions_mask, text_emb):
|
||||
"""Core dual-stream training loss given prepared latents / actions / text embeddings.
|
||||
|
||||
``latents``: ``[B, in_channels, F, h, w]`` (normalized video latents).
|
||||
@@ -311,11 +323,7 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
latents, self._train_sched_latent, action_mask=None, action_mode=False, noisy_cond_prob=0.5
|
||||
)
|
||||
action_dict = self._add_noise_stream(
|
||||
actions,
|
||||
self._train_sched_action,
|
||||
action_mask=actions_mask,
|
||||
action_mode=True,
|
||||
noisy_cond_prob=self.config.action_noisy_cond_prob,
|
||||
actions, self._train_sched_action, action_mask=actions_mask, action_mode=True, noisy_cond_prob=0.0
|
||||
)
|
||||
latent_dict["text_emb"] = text_emb
|
||||
action_dict["text_emb"] = text_emb
|
||||
@@ -327,24 +335,20 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
"window_size": int(torch.randint(4, 65, (1,)).item()),
|
||||
}
|
||||
pred = self.transformer(input_dict, train_mode=True)
|
||||
latent_loss, action_loss = self._flow_matching_loss(input_dict, pred, reduction)
|
||||
# reduction="none": latent_loss/action_loss are (B,) -> loss is per-sample (B,).
|
||||
latent_loss, action_loss = self._flow_matching_loss(input_dict, pred)
|
||||
loss = latent_loss + action_loss
|
||||
return loss, {"latent_loss": latent_loss.mean().item(), "action_loss": action_loss.mean().item()}
|
||||
return loss, {"latent_loss": latent_loss.item(), "action_loss": action_loss.item()}
|
||||
|
||||
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict | None]:
|
||||
"""Training forward: dual-stream flow-matching loss.
|
||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||
|
||||
Builds the (video-latent, action, text) training streams from a LeRobot batch
|
||||
(VAE-encoding the camera frames and UMT5-encoding the task), then runs the flow-matching
|
||||
dual-stream loss. Requires the policy to be built with ``attn_mode='flex'``.
|
||||
|
||||
``reduction="mean"`` returns the scalar loss (default); ``"none"`` returns per-sample
|
||||
losses of shape ``(B,)`` for sample weighting (RA-BC).
|
||||
dual-stream loss. Requires the policy to be built with `attn_mode='flex'`.
|
||||
"""
|
||||
self._ensure_frozen_modules()
|
||||
latents, actions, actions_mask, text_emb = self._build_training_streams(batch)
|
||||
return self.training_loss_from_streams(latents, actions, actions_mask, text_emb, reduction=reduction)
|
||||
return self.training_loss_from_streams(latents, actions, actions_mask, text_emb)
|
||||
|
||||
@torch.no_grad()
|
||||
def _build_training_streams(self, batch):
|
||||
@@ -413,31 +417,24 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
return torch.cat(per_cam, dim=-1).to(self.config.device)
|
||||
|
||||
@torch.no_grad()
|
||||
def select_action(
|
||||
self, batch: dict[str, Tensor], return_intermediate_predictions: bool = False, **kwargs
|
||||
) -> Tensor | tuple[Tensor, dict[str, Tensor]]:
|
||||
"""Return one action, refilling the chunk (and feeding back observed keyframes) as needed.
|
||||
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||
|
||||
Mirrors the upstream LIBERO client loop (``evaluation/libero/client.py``): the first obs is
|
||||
the conditioning frame; every observation produced afterwards is buffered as a keyframe and,
|
||||
once the chunk's actions are exhausted, the buffered frames + executed actions are fed back
|
||||
into the KV cache before the next chunk is predicted.
|
||||
|
||||
When ``return_intermediate_predictions=True`` returns ``(action, predictions)``. Predictions
|
||||
are produced only on the ticks that predict a fresh chunk (first tick and each chunk refill);
|
||||
on the intermediate ticks that just pop a cached action, ``predictions`` is an empty dict.
|
||||
Uses an action queue populated by `predict_action_chunk`, refilling it (and feeding back
|
||||
observed keyframes) as needed. Mirrors the upstream LIBERO client loop
|
||||
(`evaluation/libero/client.py`): the first observation is the conditioning frame; every
|
||||
observation produced afterwards is buffered as a keyframe and, once the chunk's actions are
|
||||
exhausted, the buffered frames plus executed actions are fed back into the KV cache before the
|
||||
next chunk is predicted.
|
||||
"""
|
||||
self.eval()
|
||||
self._ensure_frozen_modules()
|
||||
self._maybe_init_prompt(batch)
|
||||
|
||||
predictions: dict[str, Tensor] = {}
|
||||
if not self._started:
|
||||
# First call: this observation conditions the first chunk (it is *not* a keyframe).
|
||||
self._started = True
|
||||
actions, predictions = unpack_action_output(
|
||||
self.predict_action_chunk(batch, return_intermediate_predictions=return_intermediate_predictions)
|
||||
) # [B, chunk_size, n_used]
|
||||
actions = self.predict_action_chunk(batch) # [B, chunk_size, n_used]
|
||||
self._action_queue.extend(actions.transpose(0, 1)) # [chunk_size, B, n_used]
|
||||
self._obs_buffer = []
|
||||
self._exec_step = 0
|
||||
@@ -449,30 +446,20 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
if len(self._action_queue) == 0:
|
||||
# All actions for the current chunk have been executed; feed the observed
|
||||
# keyframes + executed actions back and predict the next chunk.
|
||||
actions, predictions = unpack_action_output(
|
||||
self.predict_action_chunk(
|
||||
None, return_intermediate_predictions=return_intermediate_predictions
|
||||
)
|
||||
)
|
||||
actions = self.predict_action_chunk(None)
|
||||
self._action_queue.extend(actions.transpose(0, 1))
|
||||
self._exec_step = 0
|
||||
|
||||
self._prev_j = self._exec_step % self.config.action_per_frame
|
||||
self._exec_step += 1
|
||||
action = self._action_queue.popleft()
|
||||
if return_intermediate_predictions:
|
||||
return action, predictions
|
||||
return action
|
||||
return self._action_queue.popleft()
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(
|
||||
self, batch: dict[str, Tensor], return_intermediate_predictions: bool = False, **kwargs
|
||||
) -> Tensor | tuple[Tensor, dict[str, Tensor]]:
|
||||
"""Run one autoregressive chunk and return actions ``[B, chunk_size, n_used]`` (normalized).
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||
|
||||
When ``return_intermediate_predictions=True`` returns ``(actions, predictions)`` where
|
||||
``predictions`` holds this chunk's VAE-decoded imagined video under ``"images.predicted"``
|
||||
(``[T, H, W, 3]`` uint8 on CPU).
|
||||
Runs one autoregressive chunk and returns actions of shape `[B, chunk_size, n_used]`
|
||||
(normalized).
|
||||
"""
|
||||
self.eval()
|
||||
self._ensure_frozen_modules()
|
||||
@@ -495,6 +482,12 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
# actions: [B, action_dim, F, action_per_frame, 1] (model-normalized). Keep for KV feedback.
|
||||
self._executed_actions = actions
|
||||
|
||||
if self.config.save_predicted_video:
|
||||
# Match upstream LingBot-VA visualization: collect chunk latents and decode the
|
||||
# concatenated latent sequence once after the rollout finishes.
|
||||
self.last_predicted_frames = None
|
||||
self.last_predicted_latents = latents.detach().to("cpu")
|
||||
|
||||
# On the first chunk, frame 0 is the conditioning frame (already "known"): the upstream
|
||||
# LIBERO client skips it (start_idx=1), so we drop the first frame's actions here.
|
||||
used = self.config.used_action_channel_ids
|
||||
@@ -503,15 +496,7 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
a = a[:, :, 1:] # drop frame 0 -> (F-1) frames of actions
|
||||
a = a.squeeze(-1).flatten(2) # [B, n_used, n_steps]
|
||||
a = a.transpose(1, 2).contiguous() # [B, n_steps, n_used]
|
||||
a = a.to(torch.float32)
|
||||
|
||||
if return_intermediate_predictions:
|
||||
# Decode this chunk's imagined video for visualization / eval. Per-chunk decode (the VAE
|
||||
# has no streaming decoder) may differ slightly at chunk boundaries from a single decode
|
||||
# over the whole concatenated latent sequence; acceptable for monitoring/inspection.
|
||||
frames = self._decode_predicted_video(latents) # [T, H, W, 3] uint8, CPU
|
||||
return a, {"images.predicted": frames}
|
||||
return a
|
||||
return a.to(torch.float32)
|
||||
|
||||
# Prompt / text encoding
|
||||
def _maybe_init_prompt(self, batch):
|
||||
@@ -872,6 +857,11 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
||||
return actions, latents
|
||||
|
||||
# Predicted-video decoding (opt-in)
|
||||
@torch.no_grad()
|
||||
def decode_predicted_latents(self, latents) -> Tensor:
|
||||
"""Decode a concatenated predicted-latent sequence into ``[T, H, W, 3]`` uint8 frames."""
|
||||
return self._decode_predicted_video(latents)
|
||||
|
||||
@torch.no_grad()
|
||||
def _decode_predicted_video(self, latents) -> Tensor:
|
||||
"""VAE-decode predicted latents into a uint8 frame stack ``[T, H, W, 3]`` on CPU."""
|
||||
|
||||
@@ -25,11 +25,9 @@ import torch
|
||||
|
||||
from lerobot.configs.types import FeatureType, NormalizationMode
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
RelativeActionsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
@@ -49,33 +47,20 @@ def make_lingbot_va_pre_post_processors(
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
# Shared relative-action step (OpenPI order: raw -> relative -> normalize -> model ->
|
||||
# unnormalize -> absolute). The SAME instance is passed to AbsoluteActionsProcessorStep
|
||||
# below so its cached raw state (set during preprocessing) flows to postprocessing.
|
||||
relative_step = RelativeActionsProcessorStep(
|
||||
enabled=config.use_relative_actions,
|
||||
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
)
|
||||
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
relative_step,
|
||||
steps.normalize,
|
||||
steps.to_device,
|
||||
]
|
||||
|
||||
# Unnormalize actions back to physical units. Config-driven norm_map (was hardcoded QUANTILES)
|
||||
# so it stays symmetric with the preprocessor's NormalizerProcessorStep — required for
|
||||
# use_relative_actions with ACTION=IDENTITY (and unchanged for QUANTILES runs).
|
||||
# Unnormalize actions from [-1, 1] to physical units (QUANTILES) using q01/q99 restored from the checkpoint.
|
||||
output_steps: list[ProcessorStep] = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features,
|
||||
norm_map=config.normalization_mapping,
|
||||
norm_map={FeatureType.ACTION: NormalizationMode.QUANTILES},
|
||||
stats=dataset_stats,
|
||||
),
|
||||
AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step),
|
||||
steps.to_cpu,
|
||||
]
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user