mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| fc573d46da |
@@ -24,24 +24,19 @@ on:
|
||||
required: false
|
||||
type: string
|
||||
|
||||
# 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.
|
||||
# Triggers the workflow on push events to main for the docs folder
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**"
|
||||
- "src/**"
|
||||
|
||||
# Same for pull requests, so a docstring change gets a preview build and a broken `[[autodoc]]` path
|
||||
# fails the PR rather than main.
|
||||
# Triggers the workflow on pull request events targeting main for the docs folder
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**"
|
||||
- "src/**"
|
||||
|
||||
release:
|
||||
types: [published]
|
||||
@@ -64,21 +59,12 @@ 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 }}
|
||||
@@ -97,6 +83,4 @@ jobs:
|
||||
commit_sha: ${{ github.event.pull_request.head.sha }}
|
||||
pr_number: ${{ github.event.number }}
|
||||
package: lerobot
|
||||
# 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]"
|
||||
additional_args: --not_python_module
|
||||
|
||||
@@ -56,41 +56,3 @@ 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
|
||||
|
||||
+2
-11
@@ -67,11 +67,7 @@ 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).
|
||||
#
|
||||
# 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)$
|
||||
exclude: ^src/lerobot/templates/.*\.md$
|
||||
|
||||
##### Security #####
|
||||
- repo: https://github.com/gitleaks/gitleaks
|
||||
@@ -108,13 +104,8 @@ 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: ["--config=pyproject.toml"]
|
||||
# pass_filenames: false
|
||||
# args: ["-vv", "--config=pyproject.toml"]
|
||||
|
||||
@@ -50,10 +50,6 @@ 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,29 +184,3 @@ 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
|
||||
|
||||
-60
@@ -1,60 +0,0 @@
|
||||
# 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]
|
||||
@@ -191,30 +191,6 @@
|
||||
- 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
|
||||
- local: api/rl
|
||||
title: Reinforcement Learning
|
||||
title: "API Reference"
|
||||
|
||||
@@ -33,7 +33,7 @@ LeRobot provides processor steps for converting between joint and EE spaces usin
|
||||
```python
|
||||
from lerobot.model.kinematics import RobotKinematics
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
|
||||
@@ -44,7 +44,7 @@ kinematics = RobotKinematics(
|
||||
)
|
||||
|
||||
# Joints → EE (for observations: "where is my gripper?")
|
||||
fk_step = ForwardKinematicsJointsToEE(kinematics=kinematics, motor_names=[...])
|
||||
fk_step = ForwardKinematicsJointsToEEObservation(kinematics=kinematics, motor_names=[...])
|
||||
|
||||
# EE → Joints (for actions: "move my gripper here")
|
||||
ik_step = InverseKinematicsEEToJoints(kinematics=kinematics, motor_names=[...])
|
||||
@@ -197,7 +197,7 @@ Here is how the different processors compose. Each arrow is a processor step, an
|
||||
```
|
||||
┌─────────────────────────────────────────┐
|
||||
Action Space │ Joint Space ←──IK──→ EE Space │
|
||||
│ ForwardKinematicsJointsToEE │
|
||||
│ ForwardKinematicsJointsToEEAction │
|
||||
│ InverseKinematicsEEToJoints │
|
||||
└─────────────────────────────────────────┘
|
||||
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
# 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
|
||||
@@ -1,27 +0,0 @@
|
||||
# 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
|
||||
@@ -1,23 +0,0 @@
|
||||
# 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
|
||||
@@ -1,19 +0,0 @@
|
||||
# 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
|
||||
@@ -1,23 +0,0 @@
|
||||
# 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
|
||||
@@ -1,20 +0,0 @@
|
||||
# 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
|
||||
|
||||
[[autodoc]] lerobot.policies.pretrained.PreTrainedPolicy
|
||||
|
||||
## PreTrainedConfig
|
||||
|
||||
[[autodoc]] lerobot.configs.PreTrainedConfig
|
||||
|
||||
## make_policy
|
||||
|
||||
[[autodoc]] lerobot.policies.factory.make_policy
|
||||
@@ -1,20 +0,0 @@
|
||||
# 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
|
||||
@@ -1,87 +0,0 @@
|
||||
# Reinforcement Learning
|
||||
|
||||
`lerobot.rl` is the distributed actor/learner reinforcement-learning stack behind
|
||||
[Train a Robot with RL](../hilserl) (HIL-SERL) and [Train RL in Simulation](../hilserl_sim). Algorithms,
|
||||
the replay buffer, data sources, and the trainer are gRPC-free and usable standalone; the actor/learner
|
||||
entry points (`actor`, `learner`, `learner_service`) additionally require `pip install 'lerobot[hilserl]'`.
|
||||
|
||||
## TrainRLServerPipelineConfig
|
||||
|
||||
Top-level configuration for both the `lerobot-actor` and `lerobot-learner` CLIs.
|
||||
|
||||
[[autodoc]] lerobot.rl.train_rl.TrainRLServerPipelineConfig
|
||||
|
||||
## RLAlgorithm
|
||||
|
||||
Abstract base every RL algorithm subclasses.
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.base.RLAlgorithm
|
||||
|
||||
## RLAlgorithmConfig
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.configs.RLAlgorithmConfig
|
||||
|
||||
## TrainingStats
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.configs.TrainingStats
|
||||
|
||||
## make_algorithm
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.factory.make_algorithm
|
||||
|
||||
## make_algorithm_config
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.factory.make_algorithm_config
|
||||
|
||||
## get_algorithm_class
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.factory.get_algorithm_class
|
||||
|
||||
## SAC
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.sac.sac_algorithm.SACAlgorithm
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.rl.algorithms.sac.configuration_sac.SACAlgorithmConfig
|
||||
|
||||
## ReplayBuffer
|
||||
|
||||
In-memory replay buffer of transitions, sampled in batches for off-policy training.
|
||||
|
||||
[[autodoc]] lerobot.rl.buffer.ReplayBuffer
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.rl.buffer.BatchTransition
|
||||
|
||||
## DataMixer
|
||||
|
||||
Abstract interface for combining online and offline data sources into training batches.
|
||||
|
||||
[[autodoc]] lerobot.rl.data_sources.DataMixer
|
||||
- all
|
||||
|
||||
[[autodoc]] lerobot.rl.data_sources.OnlineOfflineMixer
|
||||
- all
|
||||
|
||||
## RLTrainer
|
||||
|
||||
Unified training-step orchestrator: holds the algorithm, a `DataMixer`, and an optional preprocessor.
|
||||
|
||||
[[autodoc]] lerobot.rl.trainer.RLTrainer
|
||||
- all
|
||||
|
||||
## Actor / learner CLIs
|
||||
|
||||
The distributed actor and learner processes communicate over gRPC; see [Train a Robot with
|
||||
RL](../hilserl) for the full workflow.
|
||||
|
||||
[[autodoc]] lerobot.rl.actor.actor_cli
|
||||
|
||||
[[autodoc]] lerobot.rl.learner.train_cli
|
||||
|
||||
[[autodoc]] lerobot.rl.learner_service.LearnerService
|
||||
- all
|
||||
|
||||
## eval_policy
|
||||
|
||||
[[autodoc]] lerobot.rl.eval_policy.eval_policy
|
||||
@@ -1,147 +0,0 @@
|
||||
# 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
|
||||
@@ -1,30 +0,0 @@
|
||||
# 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,16 +161,6 @@ 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).
|
||||
@@ -310,7 +300,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 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.
|
||||
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.
|
||||
|
||||
Mirror an existing policy that's structurally similar to yours; the diff is small.
|
||||
|
||||
@@ -354,7 +344,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). 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.
|
||||
**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.
|
||||
|
||||
**Report results in your policy's MDX**, with the exact `lerobot-eval` command and hardware so anyone can re-run:
|
||||
|
||||
|
||||
@@ -145,7 +145,7 @@ The environment processor (`env_processor`) handles incoming observations and en
|
||||
1. **VanillaObservationProcessorStep**: Converts raw robot observations into standardized format
|
||||
2. **JointVelocityProcessorStep** (optional): Adds joint velocity information to observations
|
||||
3. **MotorCurrentProcessorStep** (optional): Adds motor current readings to observations
|
||||
4. **ForwardKinematicsJointsToEE** (optional): Computes end-effector pose from joint positions
|
||||
4. **ForwardKinematicsJointsToEEObservation** (optional): Computes end-effector pose from joint positions
|
||||
5. **ImageCropResizeProcessorStep** (optional): Crops and resizes camera images
|
||||
6. **TimeLimitProcessorStep** (optional): Enforces episode time limits
|
||||
7. **GripperPenaltyProcessorStep** (optional): Applies penalties for inappropriate gripper usage
|
||||
@@ -413,7 +413,7 @@ We support using a gamepad or a keyboard or the leader arm of the robot.
|
||||
|
||||
HIL-Serl learns actions in the end-effector space of the robot. Therefore, the teleoperation will control the end-effector's x,y,z displacements.
|
||||
|
||||
The end-effector transformation is applied by the processor pipeline (`InverseKinematicsRLStep`, `EEBoundsAndSafety`, `EEReferenceAndDelta`, `GripperVelocityToJoint`) configured under `env.processor.inverse_kinematics` (`InverseKinematicsConfig`) and `env.processor.gripper` / `env.processor.max_gripper_pos`. The defaults related to the end-effector space are:
|
||||
The end-effector transformation is applied by the processor pipeline (`EEReferenceAndDelta`, `EEBoundsAndSafety`, `GripperVelocityToJoint`, `InverseKinematicsEEToJoints`, `AddIKSolutionStep`) configured under `env.processor.inverse_kinematics` (`InverseKinematicsConfig`) and `env.processor.gripper` / `env.processor.max_gripper_pos`. The defaults related to the end-effector space are:
|
||||
|
||||
<!-- prettier-ignore-start -->
|
||||
```python
|
||||
|
||||
@@ -108,7 +108,6 @@ 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)"
|
||||
@@ -128,9 +127,7 @@ ask_vqa_top:
|
||||
}
|
||||
```
|
||||
|
||||
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.
|
||||
Add one such sub-recipe per camera the dataset records.
|
||||
|
||||
## Layer 3 — training format
|
||||
|
||||
@@ -144,20 +141,7 @@ 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` 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.
|
||||
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.
|
||||
|
||||
+129
-125
@@ -1,29 +1,28 @@
|
||||
# Multi-GPU Training
|
||||
|
||||
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` |
|
||||
This guide shows you how to train policies on multiple GPUs using [Hugging Face Accelerate](https://huggingface.co/docs/accelerate).
|
||||
|
||||
## Installation
|
||||
|
||||
`accelerate` is included in the `training` extra:
|
||||
`accelerate` is included in the `training` extra. Install it with:
|
||||
|
||||
```bash
|
||||
pip install 'lerobot[training]'
|
||||
```
|
||||
|
||||
## Launching
|
||||
## Training with Multiple GPUs
|
||||
|
||||
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.
|
||||
You can launch training in two ways:
|
||||
|
||||
With `torchrun`:
|
||||
### Option 1: Without config (specify parameters directly)
|
||||
|
||||
You can specify all parameters directly in the command without running `accelerate config`:
|
||||
|
||||
```bash
|
||||
torchrun --nproc-per-node=2 $(which lerobot-train) \
|
||||
accelerate launch \
|
||||
--multi_gpu \
|
||||
--num_processes=2 \
|
||||
$(which lerobot-train) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=act \
|
||||
--policy.repo_id=${HF_USER}/my_trained_policy \
|
||||
@@ -32,145 +31,150 @@ torchrun --nproc-per-node=2 $(which lerobot-train) \
|
||||
--wandb.enable=true
|
||||
```
|
||||
|
||||
With `accelerate launch` (as a plain launcher):
|
||||
**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:
|
||||
|
||||
```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) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=act \
|
||||
--policy.repo_id=${HF_USER}/my_trained_policy \
|
||||
--output_dir=outputs/train/act_multi_gpu \
|
||||
--job_name=act_multi_gpu \
|
||||
--wandb.enable=true
|
||||
```
|
||||
|
||||
## How It Works
|
||||
|
||||
When you launch training with accelerate:
|
||||
|
||||
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
|
||||
|
||||
## Learning Rate and Training Steps Scaling
|
||||
|
||||
**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) \
|
||||
--dataset.repo_id=${HF_USER}/my_dataset \
|
||||
--policy.type=act \
|
||||
--policy.repo_id=${HF_USER}/my_trained_policy \
|
||||
--output_dir=outputs/train/act_multi_gpu \
|
||||
--job_name=act_multi_gpu \
|
||||
--wandb.enable=true
|
||||
--optimizer.lr=2e-4 \
|
||||
--dataset.repo_id=lerobot/pusht \
|
||||
--policy.type=act
|
||||
```
|
||||
|
||||
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`).
|
||||
**Training Steps Scaling:**
|
||||
|
||||
> [!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.
|
||||
|
||||
## Batch semantics, learning rate, and steps
|
||||
|
||||
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:
|
||||
|
||||
```
|
||||
effective_batch_size = batch_size × dp_world_size × gradient_accumulation_steps
|
||||
```
|
||||
|
||||
The training banner prints this factorization at startup. `--steps` counts loop steps (micro-batches per worker), not optimizer updates.
|
||||
|
||||
Gradient accumulation is a first-class flag:
|
||||
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:
|
||||
|
||||
```bash
|
||||
torchrun --nproc-per-node=2 $(which lerobot-train) \
|
||||
--batch_size=8 --accelerator.gradient_accumulation.steps=4 ...
|
||||
# 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
|
||||
```
|
||||
|
||||
**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`.
|
||||
## Training Large Models with FSDP
|
||||
|
||||
## Sharded training (FSDP)
|
||||
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.
|
||||
|
||||
If a model is too large to train with DDP, shard it with FSDP2:
|
||||
An example on how to launch LeRobot training with FSDP across 4 GPUs (1 machine):
|
||||
|
||||
```bash
|
||||
torchrun --nproc-per-node=4 $(which lerobot-train) \
|
||||
accelerate launch --config_file fsdp.yaml --num_processes=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
|
||||
```
|
||||
|
||||
`--parallelism.dp_shard=-1` shards over however many processes the launcher started.
|
||||
A minimal `fsdp.yaml` (FSDP1; shards params/grads/optimizer — ZeRO-3-equivalent):
|
||||
|
||||
### 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
|
||||
```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
|
||||
```
|
||||
|
||||
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).
|
||||
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()`.
|
||||
|
||||
Other sharding settings:
|
||||
### FSDP checkpoints
|
||||
|
||||
- `--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.
|
||||
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:
|
||||
|
||||
### 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**.
|
||||
- **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.
|
||||
|
||||
## Notes
|
||||
|
||||
- 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).
|
||||
- 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.
|
||||
|
||||
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).
|
||||
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).
|
||||
|
||||
@@ -187,7 +187,7 @@ We use different IK initial guesses in the kinematic steps. As initial guess eit
|
||||
- EEBoundsAndSafety: clamps the EE pose to a workspace and rate‑limits jumps for safety. Also declares `action.ee.*` features.
|
||||
- InverseKinematicsEEToJoints: turns an EE pose into joint positions with IK. `initial_guess_current_joints=True` is recommended for closed‑loop control; set `False` for open‑loop replay for stability.
|
||||
- GripperVelocityToJoint: integrates a velocity‑like gripper input into an absolute gripper position using the current measured state.
|
||||
- ForwardKinematicsJointsToEE: computes `observation.state.ee.*` from observed joints for logging and training on EE state.
|
||||
- ForwardKinematicsJointsToEEObservation: computes `observation.state.ee.*` from observed joints for logging and training on EE state.
|
||||
|
||||
### Troubleshooting
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ robot_ee_to_joints_processor = RobotProcessorPipeline[RobotAction, RobotAction](
|
||||
|
||||
robot_joints_to_ee_pose = RobotProcessorPipeline[RobotObservation, RobotObservation]( # robot obs -> dataset obs
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys()))
|
||||
ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys()))
|
||||
],
|
||||
to_transition=observation_to_transition,
|
||||
to_output=transition_to_observation,
|
||||
|
||||
@@ -40,15 +40,3 @@ 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.
|
||||
|
||||
@@ -1,287 +0,0 @@
|
||||
# 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.
|
||||
@@ -36,7 +36,7 @@ from lerobot.processor import (
|
||||
)
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.utils.constants import ACTION, OBS_STR
|
||||
@@ -95,7 +95,7 @@ def main():
|
||||
# Build pipeline to convert joints observation to EE observation
|
||||
robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEObservation(
|
||||
kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys())
|
||||
)
|
||||
],
|
||||
|
||||
@@ -29,7 +29,7 @@ from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
EEBoundsAndSafety,
|
||||
EEReferenceAndDelta,
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
GripperVelocityToJoint,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
@@ -111,7 +111,7 @@ def main():
|
||||
# Build pipeline to convert joint observation to EE observation (FK).
|
||||
robot_joints_to_ee_pose = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEObservation(
|
||||
kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys())
|
||||
)
|
||||
],
|
||||
|
||||
@@ -38,7 +38,7 @@ from lerobot.processor import (
|
||||
)
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.rollout import BaseStrategyConfig, RolloutConfig, build_rollout_context
|
||||
@@ -75,7 +75,7 @@ def main():
|
||||
)
|
||||
|
||||
robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=motor_names)],
|
||||
steps=[ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=motor_names)],
|
||||
to_transition=observation_to_transition,
|
||||
to_output=transition_to_observation,
|
||||
)
|
||||
|
||||
@@ -36,7 +36,7 @@ from lerobot.processor import (
|
||||
)
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.utils.constants import ACTION, OBS_STR
|
||||
@@ -95,7 +95,7 @@ def main():
|
||||
# Build pipeline to convert joints observation to EE observation
|
||||
robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEObservation(
|
||||
kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys())
|
||||
)
|
||||
],
|
||||
|
||||
@@ -29,7 +29,8 @@ from lerobot.processor import (
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
EEBoundsAndSafety,
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEAction,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.scripts.lerobot_record import record_loop
|
||||
@@ -78,7 +79,7 @@ def main():
|
||||
# Build pipeline to convert follower joints to EE observation.
|
||||
follower_joints_to_ee = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEObservation(
|
||||
kinematics=follower_kinematics_solver, motor_names=list(follower.bus.motors.keys())
|
||||
),
|
||||
],
|
||||
@@ -89,7 +90,7 @@ def main():
|
||||
# Build pipeline to convert leader joints to EE action.
|
||||
leader_joints_to_ee = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEAction(
|
||||
kinematics=leader_kinematics_solver, motor_names=list(leader.bus.motors.keys())
|
||||
),
|
||||
],
|
||||
|
||||
@@ -36,7 +36,7 @@ from lerobot.processor import (
|
||||
)
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.rollout import BaseStrategyConfig, RolloutConfig, build_rollout_context
|
||||
@@ -78,7 +78,7 @@ def main():
|
||||
|
||||
# Joint-space observation → EE-space observation (consumed by the policy).
|
||||
robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation](
|
||||
steps=[ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=motor_names)],
|
||||
steps=[ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=motor_names)],
|
||||
to_transition=observation_to_transition,
|
||||
to_output=transition_to_observation,
|
||||
)
|
||||
|
||||
@@ -27,7 +27,7 @@ from lerobot.processor import (
|
||||
from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
EEBoundsAndSafety,
|
||||
ForwardKinematicsJointsToEE,
|
||||
ForwardKinematicsJointsToEEAction,
|
||||
InverseKinematicsEEToJoints,
|
||||
)
|
||||
from lerobot.teleoperators.so_leader import SO100Leader, SO100LeaderConfig
|
||||
@@ -65,7 +65,7 @@ def main():
|
||||
# Build pipeline to convert teleop joints to EE action
|
||||
leader_to_ee = RobotProcessorPipeline[RobotAction, RobotAction](
|
||||
steps=[
|
||||
ForwardKinematicsJointsToEE(
|
||||
ForwardKinematicsJointsToEEAction(
|
||||
kinematics=leader_kinematics_solver, motor_names=list(leader.bus.motors.keys())
|
||||
),
|
||||
],
|
||||
|
||||
+17
-80
@@ -346,7 +346,6 @@ 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"
|
||||
@@ -401,63 +400,19 @@ 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", "D" #, "A", "S", "RUF"
|
||||
"E", "W", "F", "I", "B", "C4", "T20", "N", "UP", "SIM" #, "A", "S", "D", "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", "D104"]
|
||||
"__init__.py" = ["F401", "F403", "E402"]
|
||||
# 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/**" = ["D"]
|
||||
"src/lerobot/processor/**" = ["D"]
|
||||
"src/lerobot/rewards/**" = ["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"]
|
||||
@@ -501,34 +456,25 @@ default.extend-ignore-identifiers-re = [
|
||||
"seperated_timestep",
|
||||
]
|
||||
|
||||
# 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 = 55.5
|
||||
output-format = "term-missing"
|
||||
color = true
|
||||
paths = ["src/lerobot"]
|
||||
exclude = ["src/lerobot/policies/molmoact2/molmoact2_hf_model"]
|
||||
# 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"]
|
||||
|
||||
# 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
|
||||
@@ -575,15 +521,6 @@ 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,7 +14,8 @@
|
||||
# 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,20 +40,17 @@ class OpenCVCameraConfig(CameraConfig):
|
||||
OpenCVCameraConfig(0, 30, 1280, 720, fourcc="YUYV") # With YUYV format
|
||||
```
|
||||
|
||||
**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.
|
||||
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.
|
||||
|
||||
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** (`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.
|
||||
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.
|
||||
|
||||
Note:
|
||||
- Only 3-channel color output (RGB/BGR) is currently supported.
|
||||
|
||||
@@ -36,28 +36,27 @@ 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** (`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).
|
||||
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).
|
||||
|
||||
Note:
|
||||
- Either name or serial_number must be specified.
|
||||
|
||||
+173
-603
@@ -13,41 +13,16 @@
|
||||
# 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.
|
||||
"""Training-output persistence: checkpoints, two-phase resume, and hub publishing.
|
||||
|
||||
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 huggingface_hub import HfApi, 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,
|
||||
@@ -65,39 +40,14 @@ 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.
|
||||
|
||||
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>`.
|
||||
"""
|
||||
"""Returns the checkpoint sub-directory corresponding to the step number."""
|
||||
step_identifier = get_step_identifier(step, total_steps)
|
||||
return output_dir / CHECKPOINTS_DIR / step_identifier
|
||||
|
||||
@@ -113,15 +63,37 @@ 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 update_last_checkpoint(checkpoint_dir: Path) -> None:
|
||||
"""Point the `last` symlink in the checkpoints directory at the given checkpoint.
|
||||
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)
|
||||
|
||||
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.
|
||||
|
||||
Args:
|
||||
checkpoint_dir (Path): The checkpoint step directory the `last` link should target.
|
||||
"""
|
||||
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:
|
||||
last_checkpoint_dir = checkpoint_dir.parent / LAST_CHECKPOINT_LINK
|
||||
if last_checkpoint_dir.is_symlink():
|
||||
last_checkpoint_dir.unlink()
|
||||
@@ -129,68 +101,6 @@ def update_last_checkpoint(checkpoint_dir: Path) -> None:
|
||||
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,
|
||||
@@ -200,301 +110,192 @@ def save_checkpoint(
|
||||
scheduler: LRScheduler | None = None,
|
||||
preprocessor: PolicyProcessorPipeline | None = None,
|
||||
postprocessor: PolicyProcessorPipeline | None = None,
|
||||
accelerator: "Accelerator | None" = None,
|
||||
num_processes: int | None = None,
|
||||
batch_size: int | None = None,
|
||||
model_state_dict: dict | None = None,
|
||||
optim_state_dict: dict | 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 (checkpoint_format ∈ {safetensors, safetensors_dcp}, or any non-sharded run)
|
||||
│ ├── pytorch_model_fsdp_0/ # DCP model shards (checkpoint_format ∈ {dcp, safetensors_dcp})
|
||||
│ ├── model.safetensors # policy weights
|
||||
│ ├── train_config.json # train config
|
||||
│ ├── 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
|
||||
│ ├── processor.json # processor config (if preprocessor provided)
|
||||
│ └── step_*.safetensors # processor state files (if any)
|
||||
└── training_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)
|
||||
├── optimizer_param_groups.json # optimizer param groups
|
||||
├── optimizer_state.safetensors # optimizer state
|
||||
├── rng_state.safetensors # rng states
|
||||
├── 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.
|
||||
├── scheduler_state.json # scheduler state
|
||||
└── training_step.json # training step
|
||||
|
||||
Args:
|
||||
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.
|
||||
step (int): The training step at that checkpoint.
|
||||
policy (PreTrainedPolicy): The policy to save.
|
||||
optimizer (Optimizer): The optimizer to save the state from.
|
||||
optimizer (Optimizer | None, optional): The optimizer to save the state from. Defaults to None.
|
||||
scheduler (LRScheduler | None, optional): The scheduler to save the state from. Defaults to None.
|
||||
preprocessor (PolicyProcessorPipeline | None, optional): The preprocessor/pipeline to save.
|
||||
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.
|
||||
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
|
||||
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 ----------------------------------
|
||||
policy.save_pretrained(pretrained_dir, state_dict=model_state_dict)
|
||||
cfg.save_pretrained(pretrained_dir)
|
||||
if cfg.peft is not None:
|
||||
# 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)
|
||||
|
||||
# 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)
|
||||
save_training_state(
|
||||
checkpoint_dir, step, cfg, optimizer, scheduler, accelerator, sharded=sharded, model=policy_to_save
|
||||
checkpoint_dir,
|
||||
step,
|
||||
optimizer,
|
||||
scheduler,
|
||||
num_processes=num_processes,
|
||||
batch_size=batch_size,
|
||||
optim_state_dict=optim_state_dict,
|
||||
)
|
||||
if accelerator is not None:
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
|
||||
def save_training_state(
|
||||
checkpoint_dir: Path,
|
||||
step: int,
|
||||
cfg: TrainPipelineConfig,
|
||||
optimizer: Optimizer | dict[str, Optimizer] | None = None,
|
||||
train_step: int,
|
||||
optimizer: Optimizer | None = None,
|
||||
scheduler: LRScheduler | None = None,
|
||||
accelerator: "Accelerator | None" = None,
|
||||
*,
|
||||
sharded: bool = False,
|
||||
model: PreTrainedPolicy | None = None,
|
||||
num_processes: int | None = None,
|
||||
batch_size: int | None = None,
|
||||
optim_state_dict: dict | None = None,
|
||||
) -> None:
|
||||
"""Write training_state/. Collective under sharding: call on every rank.
|
||||
"""
|
||||
Saves the training step, optimizer state, scheduler state, and rng state.
|
||||
|
||||
Args:
|
||||
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.
|
||||
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.
|
||||
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.
|
||||
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.
|
||||
"""
|
||||
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)
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# 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`.
|
||||
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.
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The resumed training config; `cfg.checkpoint_path` locates
|
||||
the checkpoint to restore from.
|
||||
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).
|
||||
|
||||
Raises:
|
||||
NotADirectoryError: If 'checkpoint_dir' doesn't contain a 'training_state' dir
|
||||
|
||||
Returns:
|
||||
int: The training step recorded in the checkpoint (micro-batch counter).
|
||||
|
||||
Raises:
|
||||
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.
|
||||
tuple[int, Optimizer, LRScheduler | None]: training step, optimizer and scheduler with their
|
||||
state_dict loaded.
|
||||
"""
|
||||
training_state_dir = cfg.checkpoint_path / TRAINING_STATE_DIR
|
||||
training_state_dir = checkpoint_dir / 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)
|
||||
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)
|
||||
|
||||
step = load_training_step(training_state_dir)
|
||||
if load_optimizer:
|
||||
optimizer = load_optimizer_state(optimizer, training_state_dir)
|
||||
if scheduler is not None:
|
||||
load_scheduler_state(scheduler, training_state_dir)
|
||||
scheduler = load_scheduler_state(scheduler, training_state_dir)
|
||||
|
||||
return step, optimizer, scheduler
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
# Hub: checkpoint push (resume artifact) and publishing (distribution artifact)
|
||||
# ---------------------------------------------------------------------------------------------
|
||||
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)
|
||||
|
||||
|
||||
def push_checkpoint_to_hub(
|
||||
@@ -510,16 +311,6 @@ 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)
|
||||
@@ -547,16 +338,6 @@ 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:
|
||||
@@ -573,214 +354,3 @@ 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
|
||||
|
||||
@@ -1,273 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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,7 +14,6 @@
|
||||
# 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
|
||||
@@ -22,8 +21,6 @@ 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:
|
||||
@@ -39,8 +36,6 @@ class DatasetConfig:
|
||||
# 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
|
||||
@@ -80,14 +75,6 @@ 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
|
||||
|
||||
@@ -1,190 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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,7 +23,6 @@ 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)",
|
||||
@@ -41,7 +40,6 @@ 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
|
||||
@@ -80,7 +78,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:
|
||||
@@ -101,16 +99,13 @@ 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`` and ``route`` are only meaningful inside a blend;
|
||||
``route: vqa`` gives sparse VQA annotations priority over normal weighted
|
||||
selection.
|
||||
sub-recipes). ``weight`` is only meaningful inside a blend.
|
||||
"""
|
||||
|
||||
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."""
|
||||
@@ -118,10 +113,6 @@ 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()
|
||||
@@ -156,9 +147,8 @@ class TrainingRecipe:
|
||||
return cls.from_dict(data)
|
||||
|
||||
def _validate_message_recipe(self) -> 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.")
|
||||
"""Ensure every templated binding is known and at least one turn is a target."""
|
||||
assert self.messages is not None
|
||||
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
|
||||
|
||||
for turn in self.messages:
|
||||
@@ -166,19 +156,12 @@ class TrainingRecipe:
|
||||
if missing:
|
||||
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
|
||||
|
||||
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)."
|
||||
)
|
||||
if not any(turn.target for turn in self.messages):
|
||||
raise ValueError("Message recipes must contain at least one target turn.")
|
||||
|
||||
def _validate_blend_recipe(self) -> None:
|
||||
"""Ensure each blend component is a non-empty, weighted message recipe."""
|
||||
if self.blend is None:
|
||||
raise ValueError("Cannot validate a blend recipe without blend components.")
|
||||
assert self.blend is not None
|
||||
if not self.blend:
|
||||
raise ValueError("Blend recipes must contain at least one component.")
|
||||
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
# 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}
|
||||
@@ -1,13 +0,0 @@
|
||||
# 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}
|
||||
@@ -1,30 +0,0 @@
|
||||
# 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}
|
||||
@@ -1,72 +0,0 @@
|
||||
# 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,7 +18,6 @@ import multiprocessing
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
@@ -27,8 +26,6 @@ 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
|
||||
@@ -42,34 +39,6 @@ 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 = (
|
||||
@@ -152,16 +121,9 @@ 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
|
||||
@@ -329,60 +291,6 @@ 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, resolve_episode_indices
|
||||
from lerobot.datasets import LeRobotDataset
|
||||
|
||||
|
||||
# Pydantic Models for SARM Subtask Annotation
|
||||
@@ -1049,10 +1049,7 @@ def main():
|
||||
torch_dtype = {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}[args.dtype]
|
||||
|
||||
# Determine 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))
|
||||
)
|
||||
episode_indices = args.episodes or 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, resolve_episode_indices
|
||||
from .utils import DEFAULT_EPISODES_PATH, create_lerobot_dataset_card
|
||||
from .video_utils import VideoEncodingManager
|
||||
|
||||
# NOTE: Low-level I/O functions (cast_stats_to_numpy, get_parquet_file_size_in_mb, etc.)
|
||||
@@ -97,7 +97,6 @@ __all__ = [
|
||||
"reencode_dataset",
|
||||
"remove_feature",
|
||||
"resolve_delta_timestamps",
|
||||
"resolve_episode_indices",
|
||||
"safe_stop_image_writer",
|
||||
"split_dataset",
|
||||
"write_stats",
|
||||
|
||||
@@ -39,7 +39,6 @@ from .io_utils import (
|
||||
hf_transform_to_torch,
|
||||
load_nested_dataset,
|
||||
)
|
||||
from .utils import resolve_episode_indices
|
||||
from .video_utils import decode_video_frames
|
||||
|
||||
|
||||
@@ -84,7 +83,7 @@ class DatasetReader:
|
||||
"""
|
||||
self._meta = meta
|
||||
self.root = root
|
||||
self.episodes = resolve_episode_indices(episodes, meta.total_episodes)
|
||||
self.episodes = episodes
|
||||
self._tolerance_s = tolerance_s
|
||||
self._video_backend = video_backend
|
||||
if image_transforms is not None and not callable(image_transforms):
|
||||
@@ -164,34 +163,10 @@ 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,7 +29,6 @@ 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(
|
||||
@@ -91,9 +90,6 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
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(
|
||||
@@ -102,7 +98,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
dataset = LeRobotDataset(
|
||||
cfg.dataset.repo_id,
|
||||
root=cfg.dataset.root,
|
||||
episodes=episodes,
|
||||
episodes=cfg.dataset.episodes,
|
||||
delta_timestamps=delta_timestamps,
|
||||
image_transforms=image_transforms,
|
||||
revision=cfg.dataset.revision,
|
||||
@@ -115,7 +111,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
dataset = StreamingLeRobotDataset(
|
||||
cfg.dataset.repo_id,
|
||||
root=cfg.dataset.root,
|
||||
episodes=episodes,
|
||||
episodes=cfg.dataset.episodes,
|
||||
delta_timestamps=delta_timestamps,
|
||||
image_transforms=image_transforms,
|
||||
revision=cfg.dataset.revision,
|
||||
|
||||
@@ -162,32 +162,14 @@ def render_sample(
|
||||
task: str | None = None,
|
||||
dataset_ctx: Any | None = None,
|
||||
) -> RenderedMessages | None:
|
||||
"""Render recipe-defined messages and supervision for one dataset sample.
|
||||
"""Render the chat-style messages for a single dataset sample.
|
||||
|
||||
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.
|
||||
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.
|
||||
"""
|
||||
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,
|
||||
@@ -201,58 +183,6 @@ 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:
|
||||
@@ -271,8 +201,7 @@ def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
||||
cumulative += component.weight or 0.0
|
||||
if draw < cumulative:
|
||||
return component
|
||||
if last_component is None:
|
||||
raise ValueError("Blend recipes must contain at least one component.")
|
||||
assert last_component is not None
|
||||
return last_component
|
||||
|
||||
|
||||
@@ -392,8 +321,7 @@ def _render_message_recipe(
|
||||
bindings: dict[str, LanguageRow | str | None],
|
||||
) -> RenderedMessages | None:
|
||||
"""Expand ``recipe.messages`` into rendered chat messages using ``bindings``."""
|
||||
if recipe.messages is None:
|
||||
raise ValueError("Cannot render a blend recipe as a message recipe.")
|
||||
assert recipe.messages is not None
|
||||
messages: list[dict[str, Any]] = []
|
||||
streams: list[str | None] = []
|
||||
target_indices: list[int] = []
|
||||
@@ -418,9 +346,7 @@ def _render_message_recipe(
|
||||
if turn.target:
|
||||
target_indices.append(message_idx)
|
||||
|
||||
# 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:
|
||||
if not target_indices:
|
||||
return None
|
||||
|
||||
rendered = {
|
||||
@@ -477,12 +403,14 @@ def _validate_rendered(rendered: RenderedMessages) -> None:
|
||||
|
||||
if len(streams) != len(messages):
|
||||
raise ValueError("message_streams must be aligned with messages.")
|
||||
# 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.")
|
||||
if not target_indices:
|
||||
raise ValueError("Rendered samples must contain at least one target 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,7 +18,6 @@ import dataclasses
|
||||
import importlib.resources
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
@@ -99,47 +98,6 @@ 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"
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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",
|
||||
]
|
||||
@@ -1,195 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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"
|
||||
@@ -1,147 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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)
|
||||
@@ -1,112 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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)."
|
||||
)
|
||||
@@ -1,94 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""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)
|
||||
@@ -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 lerobot.common.train_utils.publish_trained_model — the two must stay
|
||||
# exact log line emitted by PreTrainedPolicy.push_model_to_hub — the two must stay
|
||||
# in sync. If it ever stops matching we just fall back to stage-based completion
|
||||
# (~30s slower), so the contract is an optimization, not a correctness requirement.
|
||||
success_marker = f"Model pushed to https://huggingface.co/{repo_id}"
|
||||
|
||||
@@ -314,16 +314,11 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
To find the port, you can run our utility script:
|
||||
```bash
|
||||
lerobot-find-port.py
|
||||
```
|
||||
|
||||
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.
|
||||
>>> 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:
|
||||
@@ -600,7 +595,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
ID, and finally programs the bus' default baud-rate.
|
||||
|
||||
Args:
|
||||
motor (str): Key of the motor in `motors`.
|
||||
motor (str): Key of the motor in :pyattr:`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.
|
||||
@@ -671,7 +666,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
"""Enable torque on selected motors.
|
||||
|
||||
Args:
|
||||
motors (int | str | list[str] | None, optional): Same semantics as [`~motors.motors_bus.MotorsBus.disable_torque`].
|
||||
motors (int | str | list[str] | None, optional): Same semantics as :pymeth:`disable_torque`.
|
||||
Defaults to `None`.
|
||||
num_retry (int, optional): Number of additional retry attempts on communication failure.
|
||||
Defaults to 0.
|
||||
@@ -684,12 +679,10 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
This helper is useful to temporarily disable torque when configuring motors.
|
||||
|
||||
Example:
|
||||
```python
|
||||
>>> with bus.torque_disabled(): # doctest: +SKIP
|
||||
Examples:
|
||||
>>> with bus.torque_disabled():
|
||||
... # Safe operations here
|
||||
... pass
|
||||
```
|
||||
"""
|
||||
self.disable_torque(motors)
|
||||
try:
|
||||
@@ -702,7 +695,7 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
Args:
|
||||
timeout_ms (int | None, optional): Timeout in *milliseconds*. If `None` (default) the method falls
|
||||
back to `default_timeout`.
|
||||
back to :pyattr:`default_timeout`.
|
||||
"""
|
||||
timeout_ms = timeout_ms if timeout_ms is not None else self.default_timeout
|
||||
self.port_handler.setPacketTimeoutMillis(timeout_ms)
|
||||
@@ -753,8 +746,8 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
|
||||
Args:
|
||||
calibration_dict (dict[str, MotorCalibration]): Calibration obtained from
|
||||
[`~motors.motors_bus.MotorsBus.read_calibration`] or crafted by the user.
|
||||
cache (bool, optional): Save the calibration to `calibration`. Defaults to True.
|
||||
:pymeth:`read_calibration` or crafted by the user.
|
||||
cache (bool, optional): Save the calibration to :pyattr:`calibration`. Defaults to True.
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -762,7 +755,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 `calibration` is cleared.
|
||||
The in-memory :pyattr:`calibration` is cleared.
|
||||
|
||||
Args:
|
||||
motors (NameOrID | Sequence[NameOrID] | None, optional): Selection of motors. `None` (default)
|
||||
@@ -1076,9 +1069,9 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
) -> None:
|
||||
"""Write a value to a single motor's register.
|
||||
|
||||
Contrary to [`~motors.motors_bus.MotorsBus.sync_write`], this expects a response status packet emitted by the motor, which
|
||||
Contrary to :pymeth:`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 [`~motors.motors_bus.MotorsBus.sync_write`] but it is more reliable. It should typically be used when configuring
|
||||
slower than :pymeth:`sync_write` but it is more reliable. It should typically be used when configuring
|
||||
motors.
|
||||
|
||||
Args:
|
||||
@@ -1235,8 +1228,8 @@ class SerialMotorsBus(MotorsBusBase):
|
||||
) -> None:
|
||||
"""Write the same register on multiple motors.
|
||||
|
||||
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
|
||||
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
|
||||
frequency matters and losing some packets is acceptable (e.g. teleoperation loops).
|
||||
|
||||
Args:
|
||||
|
||||
@@ -20,6 +20,7 @@ from .optimizers import (
|
||||
SGDConfig as SGDConfig,
|
||||
XVLAAdamWConfig as XVLAAdamWConfig,
|
||||
load_optimizer_state,
|
||||
load_optimizer_state_dict,
|
||||
save_optimizer_state,
|
||||
)
|
||||
from .schedulers import (
|
||||
@@ -50,6 +51,7 @@ __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, write_json
|
||||
from lerobot.utils.io_utils import deserialize_json_into_object, load_json, write_json
|
||||
from lerobot.utils.utils import flatten_dict, unflatten_dict
|
||||
|
||||
# Type alias for parameters accepted by optimizer build() methods.
|
||||
@@ -52,11 +52,6 @@ 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"
|
||||
@@ -250,10 +245,6 @@ 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.
|
||||
|
||||
@@ -292,27 +283,35 @@ 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 (non-sharded runs; sharded runs use the DCP channel).
|
||||
"""Save optimizer state to disk.
|
||||
|
||||
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)
|
||||
_save_single_optimizer_state(optimizer, save_dir, optim_state_dict=optim_state_dict)
|
||||
|
||||
|
||||
def _save_single_optimizer_state(optimizer: torch.optim.Optimizer, save_dir: Path) -> None:
|
||||
def _save_single_optimizer_state(
|
||||
optimizer: torch.optim.Optimizer, save_dir: Path, optim_state_dict: dict | None = None
|
||||
) -> None:
|
||||
"""Save a single optimizer's state to disk."""
|
||||
state = optimizer.state_dict()
|
||||
state = dict(optim_state_dict) if optim_state_dict is not None else optimizer.state_dict()
|
||||
param_groups = state.pop("param_groups")
|
||||
flat_state = flatten_dict(state)
|
||||
save_file(flat_state, save_dir / OPTIMIZER_STATE)
|
||||
@@ -366,3 +365,19 @@ 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),
|
||||
}
|
||||
|
||||
@@ -47,8 +47,6 @@ class ACTPolicy(PreTrainedPolicy):
|
||||
|
||||
config_class = ACTConfig
|
||||
name = "act"
|
||||
# FSDP2 wrap units: one unit per transformer layer of both stacks.
|
||||
_fsdp_wrap_modules = ["ACTEncoderLayer", "ACTDecoderLayer"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -131,16 +131,12 @@ class ProcessorConfigKwargs(TypedDict, total=False):
|
||||
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** (`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.
|
||||
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.
|
||||
"""
|
||||
|
||||
preprocessor_config_filename: str | None
|
||||
@@ -246,7 +242,6 @@ 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.
|
||||
@@ -257,27 +252,22 @@ def make_policy(
|
||||
can either initialize a new policy from scratch or load a pretrained one.
|
||||
|
||||
Args:
|
||||
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): Dataset metadata used to infer feature shapes and
|
||||
types. Also provides statistics for normalization layers.
|
||||
env_cfg (EnvConfig | None): 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 mapping of dataset or environment feature
|
||||
keys to match expected policy feature names (e.g., `"left"` → `"camera1"`).
|
||||
defer_weight_load (bool): 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).
|
||||
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"`).
|
||||
|
||||
Returns:
|
||||
PreTrainedPolicy: An instantiated and device-placed policy model.
|
||||
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.")
|
||||
@@ -342,18 +332,11 @@ def make_policy(
|
||||
)
|
||||
|
||||
if cfg.pretrained_path and not cfg.use_peft:
|
||||
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)
|
||||
# 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
|
||||
|
||||
@@ -54,9 +54,6 @@ class FastWAMPolicy(PreTrainedPolicy):
|
||||
|
||||
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,
|
||||
|
||||
@@ -56,6 +56,7 @@ from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyActionProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
@@ -2297,7 +2298,7 @@ def _apply_n1_7_action_decode_transform(
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="groot_n1_7_action_decode_v1")
|
||||
class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
class GrootN17ActionDecodeStep(PolicyActionProcessorStep):
|
||||
"""Decode the full 132-D N1.7 model action back to environment actions.
|
||||
|
||||
N1.7 predicts checkpoint-order action groups. This step unnormalizes each
|
||||
@@ -2318,6 +2319,8 @@ class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
and chunk index alongside each queued action through the postprocessor.
|
||||
"""
|
||||
|
||||
skip_if_missing = True
|
||||
|
||||
env_action_dim: int = 0
|
||||
raw_stats: dict[str, Any] | None = None
|
||||
modality_config: dict[str, Any] | None = None
|
||||
@@ -2326,20 +2329,17 @@ class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
action_decode_transform: str | None = None
|
||||
pack_step: GrootN17PackInputsStep | None = field(default=None, repr=False)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if not isinstance(action, torch.Tensor):
|
||||
return transition
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
if self.raw_stats is None or self.modality_config is None:
|
||||
return transition
|
||||
return action
|
||||
|
||||
action_config = self.modality_config.get("action", {})
|
||||
if not isinstance(action_config, dict):
|
||||
return transition
|
||||
return action
|
||||
action_keys = action_config.get("modality_keys", [])
|
||||
action_configs = action_config.get("action_configs", [])
|
||||
if not isinstance(action_keys, list) or not isinstance(action_configs, list):
|
||||
return transition
|
||||
return action
|
||||
|
||||
action_np = action.detach().cpu().float().numpy()
|
||||
if self.use_relative_action and action_np.ndim != 3:
|
||||
@@ -2420,7 +2420,7 @@ class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
raise ValueError(f"Unsupported relative N1.7 action config for '{key}': {cfg}")
|
||||
|
||||
if not decoded_groups:
|
||||
return transition
|
||||
return action
|
||||
|
||||
decoded = np.concatenate(
|
||||
[decoded_groups[key] for key in action_keys if isinstance(key, str) and key in decoded_groups],
|
||||
@@ -2436,11 +2436,7 @@ class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
)
|
||||
if squeeze_horizon:
|
||||
decoded = decoded[:, 0]
|
||||
new_transition = transition.copy()
|
||||
new_transition[TransitionKey.ACTION] = torch.as_tensor(
|
||||
decoded, dtype=action.dtype, device=action.device
|
||||
)
|
||||
return new_transition
|
||||
return torch.as_tensor(decoded, dtype=action.dtype, device=action.device)
|
||||
|
||||
def transform_features(self, features):
|
||||
return features
|
||||
@@ -2461,7 +2457,9 @@ class GrootN17ActionDecodeStep(ProcessorStep):
|
||||
# silently load into it (v1 is stubbed below with the removal guidance).
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="groot_action_unpack_unnormalize_v2")
|
||||
class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
class GrootActionUnpackUnnormalizeStep(PolicyActionProcessorStep):
|
||||
skip_if_missing = True
|
||||
|
||||
env_action_dim: int = 0
|
||||
# Apply inverse of min-max normalization if it was used in preprocessor
|
||||
normalize_min_max: bool = True
|
||||
@@ -2470,12 +2468,8 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
libero_gripper_action: bool = False
|
||||
libero_gripper_binarize: bool = True
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
# Expect model outputs to be in TransitionKey.ACTION as (B, T, D_model)
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if not isinstance(action, torch.Tensor):
|
||||
return transition
|
||||
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
# Model outputs arrive as (B, T, D_model).
|
||||
# Slice to env dimension while preserving an optional action horizon.
|
||||
# Sync rollout postprocesses selected actions as (B, D); RTC postprocesses
|
||||
# chunks as (B, T, D), matching Isaac-GR00T's decode_action contract.
|
||||
@@ -2517,8 +2511,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
||||
action = action.clone()
|
||||
action[..., -1] = gripper
|
||||
|
||||
transition[TransitionKey.ACTION] = action
|
||||
return transition
|
||||
return action
|
||||
|
||||
def transform_features(self, features):
|
||||
return features
|
||||
|
||||
@@ -41,7 +41,9 @@ from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyActionProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
@@ -1007,7 +1009,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep):
|
||||
|
||||
@ProcessorStepRegistry.register(name="molmoact2_state_frame_transform")
|
||||
@dataclass
|
||||
class MolmoAct2StateFrameTransformStep(ProcessorStep):
|
||||
class MolmoAct2StateFrameTransformStep(ObservationProcessorStep):
|
||||
"""Convert robot state from arm frame to model frame before normalization.
|
||||
|
||||
Required for zero-shot deployment of MolmoAct2-SO100_101 on SO-100/101
|
||||
@@ -1023,25 +1025,21 @@ class MolmoAct2StateFrameTransformStep(ProcessorStep):
|
||||
See: https://huggingface.co/docs/lerobot/backwardcomp
|
||||
"""
|
||||
|
||||
skip_if_missing = True
|
||||
|
||||
joint_signs: list[float] | None = None
|
||||
joint_offsets: list[float] | None = None
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
if self.joint_signs is None or self.joint_offsets is None:
|
||||
return transition
|
||||
observation = transition.get(TransitionKey.OBSERVATION)
|
||||
if not isinstance(observation, dict) or OBS_STATE not in observation:
|
||||
return transition
|
||||
transition = transition.copy()
|
||||
observation = observation.copy()
|
||||
def observation(self, observation: dict[str, Any]) -> dict[str, Any]:
|
||||
if self.joint_signs is None or self.joint_offsets is None or OBS_STATE not in observation:
|
||||
return observation
|
||||
state = torch.as_tensor(observation[OBS_STATE], dtype=torch.float32).clone()
|
||||
n = len(self.joint_signs)
|
||||
signs = torch.tensor(self.joint_signs, dtype=torch.float32, device=state.device)
|
||||
offsets = torch.tensor(self.joint_offsets, dtype=torch.float32, device=state.device)
|
||||
state[..., :n] = signs * state[..., :n] + offsets
|
||||
observation[OBS_STATE] = state
|
||||
transition[TransitionKey.OBSERVATION] = observation
|
||||
return transition
|
||||
return observation
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
@@ -1054,7 +1052,7 @@ class MolmoAct2StateFrameTransformStep(ProcessorStep):
|
||||
|
||||
@ProcessorStepRegistry.register(name="molmoact2_action_frame_transform")
|
||||
@dataclass
|
||||
class MolmoAct2ActionFrameTransformStep(ProcessorStep):
|
||||
class MolmoAct2ActionFrameTransformStep(PolicyActionProcessorStep):
|
||||
"""Convert model action from model frame back to arm frame after unnormalization.
|
||||
|
||||
Inverse of MolmoAct2StateFrameTransformStep. Required for zero-shot
|
||||
@@ -1065,23 +1063,20 @@ class MolmoAct2ActionFrameTransformStep(ProcessorStep):
|
||||
See: https://huggingface.co/docs/lerobot/backwardcomp
|
||||
"""
|
||||
|
||||
skip_if_missing = True
|
||||
|
||||
joint_signs: list[float] | None = None
|
||||
joint_offsets: list[float] | None = None
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
if self.joint_signs is None or self.joint_offsets is None:
|
||||
return transition
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
return transition
|
||||
transition = transition.copy()
|
||||
return action
|
||||
action = torch.as_tensor(action, dtype=torch.float32).clone()
|
||||
n = len(self.joint_signs)
|
||||
signs = torch.tensor(self.joint_signs, dtype=torch.float32, device=action.device)
|
||||
offsets = torch.tensor(self.joint_offsets, dtype=torch.float32, device=action.device)
|
||||
action[..., :n] = signs * (action[..., :n] - offsets)
|
||||
transition[TransitionKey.ACTION] = action
|
||||
return transition
|
||||
return action
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
@@ -1094,13 +1089,11 @@ class MolmoAct2ActionFrameTransformStep(ProcessorStep):
|
||||
|
||||
@ProcessorStepRegistry.register(name="molmoact2_clamp_action")
|
||||
@dataclass
|
||||
class MolmoAct2ClampActionProcessorStep(ProcessorStep):
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
transition = transition.copy()
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is not None:
|
||||
transition[TransitionKey.ACTION] = torch.as_tensor(action).clamp(-1.0, 1.0)
|
||||
return transition
|
||||
class MolmoAct2ClampActionProcessorStep(PolicyActionProcessorStep):
|
||||
skip_if_missing = True
|
||||
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
return action.clamp(-1.0, 1.0)
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
@@ -22,9 +21,10 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import TransitionKey
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
@@ -41,7 +41,7 @@ from .configuration_pi05 import PI05Config
|
||||
|
||||
@ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step")
|
||||
@dataclass
|
||||
class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
||||
class Pi05PrepareStateTokenizerProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""
|
||||
Processor step to prepare the state and tokenize the language input.
|
||||
"""
|
||||
@@ -49,19 +49,14 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
||||
max_state_dim: int = 32
|
||||
task_key: str = "task"
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
transition = transition.copy()
|
||||
|
||||
state = transition.get(TransitionKey.OBSERVATION, {}).get(OBS_STATE)
|
||||
def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]:
|
||||
state = (self.transition.get(TransitionKey.OBSERVATION) or {}).get(OBS_STATE)
|
||||
if state is None:
|
||||
raise ValueError("State is required for PI05")
|
||||
tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key)
|
||||
tasks = complementary_data.get(self.task_key)
|
||||
if tasks is None:
|
||||
raise ValueError("No task found in complementary data")
|
||||
|
||||
# TODO: check if this necessary
|
||||
state = deepcopy(state)
|
||||
|
||||
# State should already be normalized to [-1, 1] by the NormalizerProcessorStep that runs before this step
|
||||
# Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`)
|
||||
state_np = state.cpu().numpy()
|
||||
@@ -74,10 +69,11 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
||||
full_prompt = f"Task: {cleaned_text}, State: {state_str};\nAction: "
|
||||
full_prompts.append(full_prompt)
|
||||
|
||||
transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts
|
||||
# Normalize state to [-1, 1] range if needed (assuming it's already normalized by normalizer processor step!!)
|
||||
# Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`)
|
||||
return transition
|
||||
complementary_data[self.task_key] = full_prompts
|
||||
return complementary_data
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"task_key": self.task_key, "max_state_dim": self.max_state_dim}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
@@ -22,10 +21,11 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import TransitionKey
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
ActionTokenizerProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
@@ -42,7 +42,7 @@ from .configuration_pi0_fast import PI0FastConfig
|
||||
|
||||
@ProcessorStepRegistry.register(name="pi0_fast_prepare_state_tokenizer_processor_step")
|
||||
@dataclass
|
||||
class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
||||
class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""
|
||||
Processor step to prepare the state and tokenize the language input.
|
||||
"""
|
||||
@@ -50,19 +50,14 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
||||
max_state_dim: int = 32
|
||||
task_key: str = "task"
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
transition = transition.copy()
|
||||
|
||||
state = transition.get(TransitionKey.OBSERVATION, {}).get(OBS_STATE)
|
||||
def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]:
|
||||
state = (self.transition.get(TransitionKey.OBSERVATION) or {}).get(OBS_STATE)
|
||||
if state is None:
|
||||
raise ValueError("State is required for PI0Fast")
|
||||
tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key)
|
||||
tasks = complementary_data.get(self.task_key)
|
||||
if tasks is None:
|
||||
raise ValueError("No task found in complementary data")
|
||||
|
||||
# TODO: check if this necessary
|
||||
state = deepcopy(state)
|
||||
|
||||
# State should already be normalized to [-1, 1] by the NormalizerProcessorStep that runs before this step
|
||||
# Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`)
|
||||
state_np = state.cpu().numpy()
|
||||
@@ -75,10 +70,11 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
||||
full_prompt = f"Task: {cleaned_text}, State: {state_str};\n"
|
||||
full_prompts.append(full_prompt)
|
||||
|
||||
transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts
|
||||
# Normalize state to [-1, 1] range if needed (assuming it's already normalized by normalizer processor step!!)
|
||||
# Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`)
|
||||
return transition
|
||||
complementary_data[self.task_key] = full_prompts
|
||||
return complementary_data
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"task_key": self.task_key, "max_state_dim": self.max_state_dim}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
|
||||
@@ -18,17 +18,20 @@ import builtins
|
||||
import dataclasses
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
from importlib.resources import files
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, TypeVar, Unpack
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, TypedDict, TypeVar, Unpack
|
||||
|
||||
from huggingface_hub import hf_hub_download, save_torch_state_dict
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download, save_torch_state_dict
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
from safetensors.torch import load_model as load_model_as_safetensor
|
||||
from safetensors.torch import load_model as load_model_as_safetensor, save_model as save_model_as_safetensor
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.__version__ import __version__
|
||||
from lerobot.configs import PreTrainedConfig
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||
from lerobot.utils.hub import HubMixin
|
||||
from lerobot.utils.import_utils import _peft_available, require_package
|
||||
@@ -43,14 +46,56 @@ else:
|
||||
get_peft_model = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
||||
|
||||
T = TypeVar("T", bound="PreTrainedPolicy")
|
||||
|
||||
# Pinned far above any policy's total size so save_torch_state_dict always emits exactly one
|
||||
# `model.safetensors` (no shards, no index) — a constant, not a computed byte count.
|
||||
_SINGLE_FILE_SHARD_SIZE = "1TB"
|
||||
|
||||
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.
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
class ActionSelectKwargs(TypedDict, total=False):
|
||||
@@ -65,22 +110,6 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
config_class: None
|
||||
name: None
|
||||
|
||||
# --- declarative parallelism/acceleration surface ----------------------------------------
|
||||
# Module CLASS names forming the FSDP2 wrap units (and, once wired, the activation-
|
||||
# checkpointing units). Resolved onto the accelerate plugin right before
|
||||
# `accelerator.prepare()` by `lerobot.distributed.set_fsdp_wrap_modules`; sharded training
|
||||
# with no wrap source anywhere fails loudly instead of silently wrapping only the root.
|
||||
_fsdp_wrap_modules: ClassVar[list[str] | None] = None
|
||||
# Non-`forward` entry points that must trigger FSDP2 unshard/reshard hooks when called on a
|
||||
# sharded policy (registered post-prepare via `torch.distributed.fsdp
|
||||
# .register_fsdp_forward_method`); calling them unregistered crashes on mixed Tensor/DTensor.
|
||||
_fsdp_forward_methods: ClassVar[tuple[str, ...]] = ("select_action", "predict_action_chunk")
|
||||
# Capability gate for the (future) activation-checkpointing wiring.
|
||||
supports_gradient_checkpointing: ClassVar[bool] = False
|
||||
# Declarative context-parallel plan (diffusers `ContextParallelModelPlan` semantics:
|
||||
# module FQN -> sequence split/gather spec). Reserved for the CP engine round.
|
||||
_cp_plan: ClassVar[dict[str, Any] | None] = None
|
||||
|
||||
def __init__(self, config: PreTrainedConfig, *inputs, **kwargs):
|
||||
super().__init__()
|
||||
if not isinstance(config, PreTrainedConfig):
|
||||
@@ -98,33 +127,43 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
if not getattr(cls, "name", None):
|
||||
raise TypeError(f"Class {cls.__name__} must define 'name'")
|
||||
|
||||
def _save_pretrained(self, save_directory: Path) -> None:
|
||||
"""Serialize this policy's parameters (and config) into `save_directory`.
|
||||
def save_pretrained(
|
||||
self,
|
||||
save_directory: str | Path,
|
||||
*,
|
||||
state_dict: dict[str, Tensor] | None = None,
|
||||
repo_id: str | None = None,
|
||||
push_to_hub: bool = False,
|
||||
card_kwargs: dict | None = None,
|
||||
**push_to_hub_kwargs,
|
||||
) -> str | None:
|
||||
"""Save the policy to a directory (and optionally push to the Hub).
|
||||
|
||||
Sharding is handled internally: under FSDP2 the full state dict is gathered through a
|
||||
COLLECTIVE, so when the policy is sharded this method (via `save_pretrained`) must be
|
||||
called on EVERY rank — a rank-0-gated call deadlocks. File writes happen on the main
|
||||
process only, in all layouts (single, DDP, sharded).
|
||||
|
||||
Args:
|
||||
save_directory (Path): Target directory for the policy config (`config.json`) and the
|
||||
safetensors weight file(s).
|
||||
Overrides `HubMixin.save_pretrained` to add a `state_dict` argument (mirroring
|
||||
`transformers.PreTrainedModel.save_pretrained`). Under FSDP, `self.state_dict()` would
|
||||
return sharded tensors, so the caller gathers the full state dict via a cross-rank
|
||||
collective and passes it here for `_save_pretrained` to write directly.
|
||||
"""
|
||||
# Lazy imports: the persistence layer pulls in lerobot.distributed only when saving.
|
||||
from lerobot.distributed.checkpoint import full_model_state_dict, is_sharded_module
|
||||
from lerobot.distributed.utils import is_main_process
|
||||
save_directory = Path(save_directory)
|
||||
save_directory.mkdir(parents=True, exist_ok=True)
|
||||
self._save_pretrained(save_directory, state_dict=state_dict)
|
||||
if push_to_hub:
|
||||
if repo_id is None:
|
||||
repo_id = save_directory.name
|
||||
return self.push_to_hub(repo_id=repo_id, card_kwargs=card_kwargs, **push_to_hub_kwargs)
|
||||
return None
|
||||
|
||||
model_to_save = self.module if hasattr(self, "module") else self
|
||||
if is_sharded_module(model_to_save):
|
||||
logging.info("Gathering the full state dict from all ranks (sharded policy).")
|
||||
state_dict = full_model_state_dict(model_to_save) # collective when sharded; {} off-main
|
||||
if not state_dict or not is_main_process():
|
||||
# Sharded: the gather materializes on the main rank only (emptiness check).
|
||||
# Non-sharded multi-rank (DDP): every rank holds a full dict — the explicit rank
|
||||
# gate prevents N ranks racing on the same files. Single process: never taken.
|
||||
return
|
||||
def _save_pretrained(self, save_directory: Path, state_dict: dict[str, Tensor] | None = None) -> None:
|
||||
self.config._save_pretrained(save_directory)
|
||||
save_torch_state_dict(state_dict, str(save_directory), max_shard_size=_SINGLE_FILE_SHARD_SIZE)
|
||||
model_to_save = self.module if hasattr(self, "module") else self
|
||||
if state_dict is None:
|
||||
save_model_as_safetensor(model_to_save, str(save_directory / SAFETENSORS_SINGLE_FILE))
|
||||
return
|
||||
# A pre-gathered (e.g. FSDP full) state dict was supplied: write it directly.
|
||||
# `save_torch_state_dict` discards shared-tensor duplicates just like `save_model` does;
|
||||
# pin `max_shard_size` above the total size so the output stays a single `model.safetensors`
|
||||
total_bytes = sum(t.numel() * t.element_size() for t in state_dict.values())
|
||||
save_torch_state_dict(state_dict, str(save_directory), max_shard_size=max(total_bytes, 1))
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
@@ -252,39 +291,92 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
peft_model=None,
|
||||
state_dict: dict[str, Tensor] | None = None,
|
||||
dataset_meta: LeRobotDatasetMetadata | None = None,
|
||||
) -> None:
|
||||
"""Publish this policy to the Hub.
|
||||
):
|
||||
api = HfApi()
|
||||
repo_id = api.create_repo(
|
||||
repo_id=self.config.repo_id, private=self.config.private, exist_ok=True
|
||||
).repo_id
|
||||
|
||||
Deprecated: use :func:`lerobot.common.train_utils.publish_trained_model` instead, which
|
||||
also publishes the pre/post-processors alongside the model.
|
||||
# Push the files to the repo in a single commit
|
||||
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
|
||||
saved_path = Path(tmp) / repo_id
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and
|
||||
used to render the model card.
|
||||
peft_model: The PEFT wrapper when training adapters, whose weights replace the full
|
||||
model weights in the published repo. Defaults to None.
|
||||
state_dict (dict[str, Tensor] | None): Ignored; weights are now gathered internally
|
||||
when the policy is sharded. Defaults to None.
|
||||
dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata for the model card,
|
||||
if available. Defaults to None.
|
||||
"""
|
||||
from lerobot.common.train_utils import publish_trained_model
|
||||
if peft_model is not None:
|
||||
# Since PEFT just forwards calls to `push_model_to_hub`, `self` is not the PeftModel wrapper
|
||||
# but the actual policy which is why we need the PEFT model passed to us to save the adapter.
|
||||
# That also means that we need to store the policy config ourselves since PEFT can't.
|
||||
peft_model.save_pretrained(saved_path)
|
||||
self.config.save_pretrained(saved_path)
|
||||
else:
|
||||
# Calls _save_pretrained and stores model tensors
|
||||
self.save_pretrained(saved_path, state_dict=state_dict)
|
||||
|
||||
warnings.warn(
|
||||
"PreTrainedPolicy.push_model_to_hub is deprecated and will be removed in a future "
|
||||
"version. Use lerobot.common.train_utils.publish_trained_model(cfg, model, "
|
||||
"preprocessor, postprocessor, dataset_meta) instead.",
|
||||
FutureWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if state_dict is not None:
|
||||
warnings.warn(
|
||||
"The `state_dict` argument is ignored: sharded weights are gathered internally "
|
||||
"when the policy is saved.",
|
||||
FutureWarning,
|
||||
stacklevel=2,
|
||||
card = self.generate_model_card(
|
||||
cfg.dataset.repo_id,
|
||||
self.config.type,
|
||||
self.config.license,
|
||||
self.config.tags,
|
||||
cfg=cfg,
|
||||
dataset_meta=dataset_meta,
|
||||
)
|
||||
publish_trained_model(cfg, self, None, None, dataset_meta, peft_model=peft_model)
|
||||
card.save(str(saved_path / "README.md"))
|
||||
|
||||
cfg.save_pretrained(saved_path) # Calls _save_pretrained and stores train config
|
||||
|
||||
commit_info = api.upload_folder(
|
||||
repo_id=repo_id,
|
||||
repo_type="model",
|
||||
folder_path=saved_path,
|
||||
commit_message="Upload policy weights, train config and readme",
|
||||
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
|
||||
ignore_patterns=["*.tmp", "*.log"],
|
||||
)
|
||||
|
||||
# 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 (it falls back to status polling if they drift).
|
||||
logging.info(f"Model pushed to {commit_info.repo_url.url}")
|
||||
|
||||
def generate_model_card(
|
||||
self,
|
||||
dataset_repo_id: str,
|
||||
model_type: str,
|
||||
license: str | None,
|
||||
tags: list[str] | None,
|
||||
cfg: TrainPipelineConfig | None = None,
|
||||
dataset_meta: LeRobotDatasetMetadata | None = None,
|
||||
) -> ModelCard:
|
||||
base_model_mapping = {
|
||||
"smolvla": "lerobot/smolvla_base",
|
||||
"pi0": "lerobot/pi0_base",
|
||||
"pi05": "lerobot/pi05_base",
|
||||
"pi0_fast": "lerobot/pi0fast-base",
|
||||
"xvla": "lerobot/xvla-base",
|
||||
}
|
||||
|
||||
card_data = ModelCardData(
|
||||
license=license or "apache-2.0",
|
||||
library_name="lerobot",
|
||||
pipeline_tag="robotics",
|
||||
tags=list(set(tags or []).union({"robotics", "lerobot", model_type})),
|
||||
model_name=model_type,
|
||||
datasets=dataset_repo_id,
|
||||
base_model=base_model_mapping.get(model_type),
|
||||
)
|
||||
|
||||
context = _build_card_context(
|
||||
cfg, dataset_meta, self.config.input_features, self.config.output_features
|
||||
)
|
||||
# Used by the template to pre-fill commands and the "Fine-tuned from" line.
|
||||
context["policy_repo_id"] = getattr(self.config, "repo_id", None)
|
||||
context["base_model"] = base_model_mapping.get(model_type)
|
||||
|
||||
template_card = (
|
||||
files("lerobot.templates").joinpath("lerobot_modelcard_template.md").read_text(encoding="utf-8")
|
||||
)
|
||||
card = ModelCard.from_template(card_data, template_str=template_card, **context)
|
||||
card.validate()
|
||||
return card
|
||||
|
||||
def wrap_with_peft(
|
||||
self,
|
||||
|
||||
@@ -46,11 +46,10 @@ class ActionQueue:
|
||||
Args:
|
||||
cfg (RTCConfig): Configuration for Real-Time Chunking behavior.
|
||||
|
||||
**Attributes**:
|
||||
- **queue** (`Tensor | None`) -- Processed actions for robot rollout (time_steps, action_dim).
|
||||
- **original_queue** (`Tensor | None`) -- Original actions for RTC computation (time_steps,
|
||||
action_dim).
|
||||
- **last_index** (`int`) -- Current consumption index in the queue.
|
||||
Attributes:
|
||||
queue (Tensor | None): Processed actions for robot rollout (time_steps, action_dim).
|
||||
original_queue (Tensor | None): Original actions for RTC computation (time_steps, action_dim).
|
||||
last_index (int): Current consumption index in the queue.
|
||||
"""
|
||||
|
||||
def __init__(self, cfg: RTCConfig):
|
||||
|
||||
@@ -27,19 +27,19 @@ from torch import Tensor
|
||||
class DebugStep:
|
||||
"""Container for debug information from a single denoising step.
|
||||
|
||||
**Attributes**:
|
||||
- **step_idx** (`int`) -- Step index/counter.
|
||||
- **x_t** (`Tensor | None`) -- Current latent/state tensor.
|
||||
- **v_t** (`Tensor | None`) -- Velocity from denoiser.
|
||||
- **x1_t** (`Tensor | None`) -- Denoised prediction (x_t - time * v_t).
|
||||
- **correction** (`Tensor | None`) -- Correction gradient tensor.
|
||||
- **err** (`Tensor | None`) -- Weighted error term.
|
||||
- **weights** (`Tensor | None`) -- Prefix attention weights.
|
||||
- **guidance_weight** (`float | Tensor | None`) -- Applied guidance weight.
|
||||
- **time** (`float | Tensor | None`) -- Time parameter.
|
||||
- **inference_delay** (`int | None`) -- Inference delay parameter.
|
||||
- **execution_horizon** (`int | None`) -- Execution horizon parameter.
|
||||
- **metadata** (`dict[str, Any]`) -- Additional metadata.
|
||||
Attributes:
|
||||
step_idx (int): Step index/counter.
|
||||
x_t (Tensor | None): Current latent/state tensor.
|
||||
v_t (Tensor | None): Velocity from denoiser.
|
||||
x1_t (Tensor | None): Denoised prediction (x_t - time * v_t).
|
||||
correction (Tensor | None): Correction gradient tensor.
|
||||
err (Tensor | None): Weighted error term.
|
||||
weights (Tensor | None): Prefix attention weights.
|
||||
guidance_weight (float | Tensor | None): Applied guidance weight.
|
||||
time (float | Tensor | None): Time parameter.
|
||||
inference_delay (int | None): Inference delay parameter.
|
||||
execution_horizon (int | None): Execution horizon parameter.
|
||||
metadata (dict[str, Any]): Additional metadata.
|
||||
"""
|
||||
|
||||
step_idx: int = 0
|
||||
|
||||
@@ -20,12 +20,11 @@ import torch
|
||||
|
||||
from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig
|
||||
from lerobot.processor import (
|
||||
EnvTransition,
|
||||
PolicyAction,
|
||||
PolicyActionProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
TransitionKey,
|
||||
UnnormalizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
@@ -33,22 +32,20 @@ from lerobot.processor import (
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="vla_jepa_clip_actions")
|
||||
class ClipActionsProcessorStep(ProcessorStep):
|
||||
class ClipActionsProcessorStep(PolicyActionProcessorStep):
|
||||
"""Clips action tensor to [-1, 1] before unnormalization."""
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is not None:
|
||||
transition = dict(transition)
|
||||
transition[TransitionKey.ACTION] = action.clamp(-1.0, 1.0)
|
||||
return transition
|
||||
skip_if_missing = True
|
||||
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
return action.clamp(-1.0, 1.0)
|
||||
|
||||
def transform_features(self, features):
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="vla_jepa_pre_snap_gripper")
|
||||
class PreSnapGripperProcessorStep(ProcessorStep):
|
||||
class PreSnapGripperProcessorStep(PolicyActionProcessorStep):
|
||||
"""Snaps a gripper dimension to {0, 1} BEFORE unnormalization.
|
||||
|
||||
Mirrors the original starVLA LIBERO eval:
|
||||
@@ -58,43 +55,49 @@ class PreSnapGripperProcessorStep(ProcessorStep):
|
||||
space where 0=open and 1=close.
|
||||
"""
|
||||
|
||||
skip_if_missing = True
|
||||
|
||||
def __init__(self, gripper_dim: int = 6, threshold: float = 0.5):
|
||||
self.gripper_dim = gripper_dim
|
||||
self.threshold = threshold
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is not None and action.shape[-1] > self.gripper_dim:
|
||||
transition = dict(transition)
|
||||
a = action.clone()
|
||||
a[..., self.gripper_dim] = (a[..., self.gripper_dim] >= self.threshold).float()
|
||||
transition[TransitionKey.ACTION] = a
|
||||
return transition
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
if action.shape[-1] <= self.gripper_dim:
|
||||
return action
|
||||
a = action.clone()
|
||||
a[..., self.gripper_dim] = (a[..., self.gripper_dim] >= self.threshold).float()
|
||||
return a
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"gripper_dim": self.gripper_dim, "threshold": self.threshold}
|
||||
|
||||
def transform_features(self, features):
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="vla_jepa_binarize_gripper")
|
||||
class BinarizeGripperProcessorStep(ProcessorStep):
|
||||
class BinarizeGripperProcessorStep(PolicyActionProcessorStep):
|
||||
"""Binarizes a gripper dimension after unnormalization.
|
||||
|
||||
Maps continuous value to {-1, 1}: > threshold → -1, <= threshold → 1 (matches starVLA convention).
|
||||
Only applied when action has more dimensions than gripper_dim.
|
||||
"""
|
||||
|
||||
skip_if_missing = True
|
||||
|
||||
def __init__(self, gripper_dim: int = 6, threshold: float = 0.5):
|
||||
self.gripper_dim = gripper_dim
|
||||
self.threshold = threshold
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is not None and action.shape[-1] > self.gripper_dim:
|
||||
transition = dict(transition)
|
||||
a = action.clone()
|
||||
a[..., self.gripper_dim] = 1.0 - 2.0 * (a[..., self.gripper_dim] > self.threshold).float()
|
||||
transition[TransitionKey.ACTION] = a
|
||||
return transition
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
if action.shape[-1] <= self.gripper_dim:
|
||||
return action
|
||||
a = action.clone()
|
||||
a[..., self.gripper_dim] = 1.0 - 2.0 * (a[..., self.gripper_dim] > self.threshold).float()
|
||||
return a
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"gripper_dim": self.gripper_dim, "threshold": self.threshold}
|
||||
|
||||
def transform_features(self, features):
|
||||
return features
|
||||
|
||||
@@ -21,10 +21,12 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import TransitionKey
|
||||
from lerobot.processor import (
|
||||
ComplementaryDataProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyActionProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
@@ -201,7 +203,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="xvla_image_scale")
|
||||
class XVLAImageScaleProcessorStep(ProcessorStep):
|
||||
class XVLAImageScaleProcessorStep(ObservationProcessorStep):
|
||||
"""Scale image observations by 255 to convert from [0, 1] to [0, 255] range.
|
||||
|
||||
This processor step multiplies all image observations by 255, which is required
|
||||
@@ -214,29 +216,22 @@ class XVLAImageScaleProcessorStep(ProcessorStep):
|
||||
|
||||
image_keys: list[str] | None = None
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
skip_if_missing = True
|
||||
|
||||
def observation(self, observation):
|
||||
"""Scale image observations by 255."""
|
||||
new_transition = transition.copy()
|
||||
obs = new_transition.get(TransitionKey.OBSERVATION, {})
|
||||
if obs is None:
|
||||
return new_transition
|
||||
|
||||
# Make a copy of observations to avoid modifying the original
|
||||
obs = obs.copy()
|
||||
|
||||
# Determine which keys to scale
|
||||
keys_to_scale = self.image_keys
|
||||
if keys_to_scale is None:
|
||||
# Auto-detect image keys
|
||||
keys_to_scale = [k for k in obs if k.startswith(OBS_IMAGES)]
|
||||
keys_to_scale = [k for k in observation if k.startswith(OBS_IMAGES)]
|
||||
|
||||
# Scale each image
|
||||
for key in keys_to_scale:
|
||||
if key in obs and isinstance(obs[key], torch.Tensor):
|
||||
obs[key] = obs[key] * 255
|
||||
if key in observation and isinstance(observation[key], torch.Tensor):
|
||||
observation[key] = observation[key] * 255
|
||||
|
||||
new_transition[TransitionKey.OBSERVATION] = obs
|
||||
return new_transition
|
||||
return observation
|
||||
|
||||
def transform_features(self, features):
|
||||
"""Image scaling doesn't change feature structure."""
|
||||
@@ -251,7 +246,7 @@ class XVLAImageScaleProcessorStep(ProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="xvla_image_to_float")
|
||||
class XVLAImageToFloatProcessorStep(ProcessorStep):
|
||||
class XVLAImageToFloatProcessorStep(ObservationProcessorStep):
|
||||
"""Convert image observations from [0, 255] to [0, 1] range.
|
||||
|
||||
This processor step divides image observations by 255 to convert from uint8-like
|
||||
@@ -270,32 +265,26 @@ class XVLAImageToFloatProcessorStep(ProcessorStep):
|
||||
image_keys: list[str] | None = None
|
||||
validate_range: bool = True
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
skip_if_missing = True
|
||||
|
||||
def observation(self, observation):
|
||||
"""Convert image observations from [0, 255] to [0, 1]."""
|
||||
new_transition = transition.copy()
|
||||
obs = new_transition.get(TransitionKey.OBSERVATION, {})
|
||||
if obs is None:
|
||||
return new_transition
|
||||
|
||||
# Make a copy of observations to avoid modifying the original
|
||||
obs = obs.copy()
|
||||
|
||||
# Determine which keys to convert
|
||||
keys_to_convert = self.image_keys
|
||||
if keys_to_convert is None:
|
||||
# Auto-detect image keys
|
||||
keys_to_convert = [k for k in obs if k.startswith(OBS_IMAGES)]
|
||||
keys_to_convert = [k for k in observation if k.startswith(OBS_IMAGES)]
|
||||
|
||||
# Convert each image
|
||||
for key in keys_to_convert:
|
||||
if key in obs and isinstance(obs[key], torch.Tensor):
|
||||
tensor = obs[key]
|
||||
if key in observation and isinstance(observation[key], torch.Tensor):
|
||||
tensor = observation[key]
|
||||
|
||||
min_val = tensor.min().item()
|
||||
max_val = tensor.max().item()
|
||||
|
||||
if max_val <= 1.0:
|
||||
obs[key] = tensor.float() # ensure float dtype, but no division
|
||||
observation[key] = tensor.float() # ensure float dtype, but no division
|
||||
continue
|
||||
# Validate that values are in [0, 255] range if requested
|
||||
if self.validate_range and (min_val < 0.0 or max_val > 255.0):
|
||||
@@ -306,10 +295,9 @@ class XVLAImageToFloatProcessorStep(ProcessorStep):
|
||||
)
|
||||
|
||||
# Convert to float and divide by 255
|
||||
obs[key] = tensor.float() / 255.0
|
||||
observation[key] = tensor.float() / 255.0
|
||||
|
||||
new_transition[TransitionKey.OBSERVATION] = obs
|
||||
return new_transition
|
||||
return observation
|
||||
|
||||
def transform_features(self, features):
|
||||
"""Image conversion doesn't change feature structure."""
|
||||
@@ -325,7 +313,7 @@ class XVLAImageToFloatProcessorStep(ProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="xvla_imagenet_normalize")
|
||||
class XVLAImageNetNormalizeProcessorStep(ProcessorStep):
|
||||
class XVLAImageNetNormalizeProcessorStep(ObservationProcessorStep):
|
||||
"""Normalize image observations using ImageNet statistics.
|
||||
|
||||
This processor step applies ImageNet normalization (mean and std) to image observations.
|
||||
@@ -343,26 +331,20 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep):
|
||||
|
||||
image_keys: list[str] | None = None
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
skip_if_missing = True
|
||||
|
||||
def observation(self, observation):
|
||||
"""Normalize image observations using ImageNet statistics."""
|
||||
new_transition = transition.copy()
|
||||
obs = new_transition.get(TransitionKey.OBSERVATION, {})
|
||||
if obs is None:
|
||||
return new_transition
|
||||
|
||||
# Make a copy of observations to avoid modifying the original
|
||||
obs = obs.copy()
|
||||
|
||||
# Determine which keys to normalize
|
||||
keys_to_normalize = self.image_keys
|
||||
if keys_to_normalize is None:
|
||||
# Auto-detect image keys
|
||||
keys_to_normalize = [k for k in obs if k.startswith(OBS_IMAGES)]
|
||||
keys_to_normalize = [k for k in observation if k.startswith(OBS_IMAGES)]
|
||||
|
||||
# Normalize each image
|
||||
for key in keys_to_normalize:
|
||||
if key in obs and isinstance(obs[key], torch.Tensor):
|
||||
tensor = obs[key]
|
||||
if key in observation and isinstance(observation[key], torch.Tensor):
|
||||
tensor = observation[key]
|
||||
|
||||
# Validate that values are in [0, 1] range
|
||||
min_val = tensor.min().item()
|
||||
@@ -384,10 +366,9 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep):
|
||||
std = std.unsqueeze(0)
|
||||
|
||||
# Normalize: (image - mean) / std
|
||||
obs[key] = (tensor - mean) / std
|
||||
observation[key] = (tensor - mean) / std
|
||||
|
||||
new_transition[TransitionKey.OBSERVATION] = obs
|
||||
return new_transition
|
||||
return observation
|
||||
|
||||
def transform_features(self, features):
|
||||
"""ImageNet normalization doesn't change feature structure."""
|
||||
@@ -402,38 +383,32 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="xvla_add_domain_id")
|
||||
class XVLAAddDomainIdProcessorStep(ProcessorStep):
|
||||
class XVLAAddDomainIdProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""Add domain_id to complementary data.
|
||||
|
||||
This processor step adds a domain_id tensor to the complementary data,
|
||||
which is used by XVLA to identify different robot embodiments or task domains.
|
||||
|
||||
Args:
|
||||
domain_id: The domain ID to add (default: 3)
|
||||
domain_id: The domain ID to add (default: 0)
|
||||
"""
|
||||
|
||||
domain_id: int = 0
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
def complementary_data(self, complementary_data):
|
||||
"""Add domain_id to complementary data."""
|
||||
new_transition = transition.copy()
|
||||
comp = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
comp = {} if comp is None else comp.copy()
|
||||
|
||||
# Infer batch size from observation tensors
|
||||
obs = new_transition.get(TransitionKey.OBSERVATION, {})
|
||||
obs = self.transition.get(TransitionKey.OBSERVATION) or {}
|
||||
batch_size = 1
|
||||
if obs:
|
||||
for v in obs.values():
|
||||
if isinstance(v, torch.Tensor):
|
||||
batch_size = v.shape[0]
|
||||
break
|
||||
for v in obs.values():
|
||||
if isinstance(v, torch.Tensor):
|
||||
batch_size = v.shape[0]
|
||||
break
|
||||
|
||||
# Add domain_id tensor
|
||||
comp["domain_id"] = torch.tensor([int(self.domain_id)] * batch_size, dtype=torch.long)
|
||||
complementary_data["domain_id"] = torch.tensor([int(self.domain_id)] * batch_size, dtype=torch.long)
|
||||
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = comp
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
def transform_features(self, features):
|
||||
"""Domain ID addition doesn't change feature structure."""
|
||||
@@ -448,7 +423,7 @@ class XVLAAddDomainIdProcessorStep(ProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="xvla_rotation_6d_to_axis_angle")
|
||||
class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep):
|
||||
class XVLARotation6DToAxisAngleProcessorStep(PolicyActionProcessorStep):
|
||||
"""Convert 6D rotation representation to axis-angle and reorganize action dimensions.
|
||||
|
||||
This processor step takes actions with 6D rotation representation and converts them to
|
||||
@@ -465,14 +440,10 @@ class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep):
|
||||
|
||||
expected_action_dim: int = 10
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
skip_if_missing = True
|
||||
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
"""Convert 6D rotation to axis-angle in action."""
|
||||
new_transition = transition.copy()
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
|
||||
if action is None or not isinstance(action, torch.Tensor):
|
||||
return new_transition
|
||||
|
||||
# Convert to numpy for processing
|
||||
device = action.device
|
||||
dtype = action.dtype
|
||||
@@ -494,10 +465,7 @@ class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep):
|
||||
action_np[:, -1] = np.where(action_np[:, -1] > 0.5, 1.0, -1.0)
|
||||
|
||||
# Convert back to tensor
|
||||
action = torch.from_numpy(action_np).to(device=device, dtype=dtype)
|
||||
|
||||
new_transition[TransitionKey.ACTION] = action
|
||||
return new_transition
|
||||
return torch.from_numpy(action_np).to(device=device, dtype=dtype)
|
||||
|
||||
def transform_features(self, features):
|
||||
"""Rotation conversion changes action dimension from 10 to 7."""
|
||||
|
||||
@@ -53,13 +53,15 @@ from .factory import (
|
||||
)
|
||||
from .gym_action_processor import (
|
||||
Numpy2TorchActionProcessorStep,
|
||||
Numpy2TorchTeleopActionProcessorStep,
|
||||
Torch2NumpyActionProcessorStep,
|
||||
)
|
||||
from .hil_processor import (
|
||||
AddTeleopActionAsComplimentaryDataStep,
|
||||
AddTeleopEventsAsInfoStep,
|
||||
GripperPenaltyProcessorStep,
|
||||
GymHILAdapterProcessorStep,
|
||||
GymHILInfoAdapterStep,
|
||||
GymHILTeleopDataAdapterStep,
|
||||
ImageCropResizeProcessorStep,
|
||||
InterventionActionProcessorStep,
|
||||
RewardClassifierProcessorStep,
|
||||
@@ -126,7 +128,8 @@ __all__ = [
|
||||
"DoneProcessorStep",
|
||||
"EnvAction",
|
||||
"EnvTransition",
|
||||
"GymHILAdapterProcessorStep",
|
||||
"GymHILInfoAdapterStep",
|
||||
"GymHILTeleopDataAdapterStep",
|
||||
"GripperPenaltyProcessorStep",
|
||||
"hotswap_stats",
|
||||
"IdentityProcessorStep",
|
||||
@@ -148,6 +151,7 @@ __all__ = [
|
||||
"NewLineTaskProcessorStep",
|
||||
"NormalizerProcessorStep",
|
||||
"Numpy2TorchActionProcessorStep",
|
||||
"Numpy2TorchTeleopActionProcessorStep",
|
||||
"ObservationProcessorStep",
|
||||
"PolicyAction",
|
||||
"PolicyActionProcessorStep",
|
||||
|
||||
@@ -175,6 +175,9 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
|
||||
if isinstance(task_index_value, Tensor) and task_index_value.dim() == 0:
|
||||
complementary_data["task_index"] = task_index_value.unsqueeze(0)
|
||||
|
||||
complementary_data.pop("language_persistent", None)
|
||||
complementary_data.pop("language_events", None)
|
||||
|
||||
if "messages" in complementary_data:
|
||||
messages = complementary_data["messages"]
|
||||
if isinstance(messages, list) and (not messages or isinstance(messages[0], dict)):
|
||||
@@ -217,12 +220,10 @@ class AddBatchDimensionProcessorStep(ProcessorStep):
|
||||
This step combines individual processors for actions, observations, and complementary data
|
||||
to create a batched transition (batch size 1) from a single-instance transition.
|
||||
|
||||
**Attributes**:
|
||||
- **to_batch_action_processor** (`AddBatchDimensionActionStep`) -- Processor for the action component.
|
||||
- **to_batch_observation_processor** (`AddBatchDimensionObservationStep`) -- Processor for the
|
||||
observation component.
|
||||
- **to_batch_complementary_data_processor** (`AddBatchDimensionComplementaryDataStep`) -- Processor
|
||||
for the complementary data component.
|
||||
Attributes:
|
||||
to_batch_action_processor: Processor for the action component.
|
||||
to_batch_observation_processor: Processor for the observation component.
|
||||
to_batch_complementary_data_processor: Processor for the complementary data component.
|
||||
"""
|
||||
|
||||
to_batch_action_processor: AddBatchDimensionActionStep = field(
|
||||
|
||||
@@ -32,8 +32,9 @@ class MapTensorToDeltaActionDictStep(ActionProcessorStep):
|
||||
It decomposes the vector into named components for delta movements of the
|
||||
end-effector (x, y, z) and optionally the gripper.
|
||||
|
||||
**Attributes**:
|
||||
- **use_gripper** (`bool`) -- If True, assumes the 4th element of the tensor is the gripper action.
|
||||
Attributes:
|
||||
use_gripper: If True, assumes the 4th element of the tensor is the
|
||||
gripper action.
|
||||
"""
|
||||
|
||||
use_gripper: bool = True
|
||||
@@ -80,10 +81,10 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
|
||||
into a target action format that includes an "enabled" flag and target
|
||||
end-effector positions. It also handles scaling and noise filtering.
|
||||
|
||||
**Attributes**:
|
||||
- **position_scale** (`float`) -- A factor to scale the delta position inputs.
|
||||
- **noise_threshold** (`float`) -- The magnitude below which delta inputs are considered noise and do
|
||||
not trigger an "enabled" state.
|
||||
Attributes:
|
||||
position_scale: A factor to scale the delta position inputs.
|
||||
noise_threshold: The magnitude below which delta inputs are considered noise
|
||||
and do not trigger an "enabled" state.
|
||||
"""
|
||||
|
||||
# Scale factors for delta movements
|
||||
|
||||
@@ -40,10 +40,10 @@ class DeviceProcessorStep(ProcessorStep):
|
||||
|
||||
This is crucial for preparing data for model training or inference on hardware like GPUs.
|
||||
|
||||
**Attributes**:
|
||||
- **device** (`str`) -- The target device for tensors (e.g., "cpu", "cuda", "cuda:0").
|
||||
- **float_dtype** (`str | None`) -- The target floating-point dtype as a string (e.g., "float32",
|
||||
"float16", "bfloat16"). If None, the dtype is not changed.
|
||||
Attributes:
|
||||
device: The target device for tensors (e.g., "cpu", "cuda", "cuda:0").
|
||||
float_dtype: The target floating-point dtype as a string (e.g., "float32", "float16", "bfloat16").
|
||||
If None, the dtype is not changed.
|
||||
"""
|
||||
|
||||
device: str = "cpu"
|
||||
|
||||
@@ -17,11 +17,11 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvAction, EnvTransition, PolicyAction, TransitionKey
|
||||
from lerobot.lerobot_types import EnvAction, PolicyAction
|
||||
|
||||
from .converters import to_tensor
|
||||
from .hil_processor import TELEOP_ACTION_KEY
|
||||
from .pipeline import ActionProcessorStep, ProcessorStep, ProcessorStepRegistry
|
||||
from .pipeline import ActionProcessorStep, ComplementaryDataProcessorStep, ProcessorStepRegistry
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("torch2numpy_action_processor")
|
||||
@@ -33,9 +33,10 @@ class Torch2NumpyActionProcessorStep(ActionProcessorStep):
|
||||
This step is useful when the output of a policy (typically a torch.Tensor)
|
||||
needs to be passed to an environment or component that expects a NumPy array.
|
||||
|
||||
**Attributes**:
|
||||
- **squeeze_batch_dim** (`bool`) -- If True, removes the first dimension of the array if it is of size
|
||||
1. This is useful for converting a batched action of size (1, D) to a single action of size (D,).
|
||||
Attributes:
|
||||
squeeze_batch_dim: If True, removes the first dimension of the array
|
||||
if it is of size 1. This is useful for converting a
|
||||
batched action of size (1, D) to a single action of size (D,).
|
||||
"""
|
||||
|
||||
squeeze_batch_dim: bool = True
|
||||
@@ -69,32 +70,36 @@ class Torch2NumpyActionProcessorStep(ActionProcessorStep):
|
||||
|
||||
@ProcessorStepRegistry.register("numpy2torch_action_processor")
|
||||
@dataclass
|
||||
class Numpy2TorchActionProcessorStep(ProcessorStep):
|
||||
class Numpy2TorchActionProcessorStep(ActionProcessorStep):
|
||||
"""Converts a NumPy array action to a PyTorch tensor when action is present."""
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
"""Converts numpy action to torch tensor if action exists, otherwise passes through."""
|
||||
self._current_transition = transition.copy()
|
||||
new_transition = self._current_transition
|
||||
skip_if_missing = True
|
||||
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is not None:
|
||||
if not isinstance(action, EnvAction):
|
||||
raise TypeError(
|
||||
f"Expected np.ndarray or None, got {type(action).__name__}. "
|
||||
"Use appropriate processor for non-tensor actions."
|
||||
)
|
||||
torch_action = to_tensor(action, dtype=None) # Preserve original dtype
|
||||
new_transition[TransitionKey.ACTION] = torch_action
|
||||
|
||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
if TELEOP_ACTION_KEY in complementary_data:
|
||||
teleop_action = complementary_data[TELEOP_ACTION_KEY]
|
||||
if isinstance(teleop_action, EnvAction):
|
||||
complementary_data[TELEOP_ACTION_KEY] = to_tensor(teleop_action)
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
||||
|
||||
return new_transition
|
||||
def action(self, action: EnvAction) -> PolicyAction:
|
||||
if not isinstance(action, EnvAction):
|
||||
raise TypeError(
|
||||
f"Expected np.ndarray or None, got {type(action).__name__}. "
|
||||
"Use appropriate processor for non-tensor actions."
|
||||
)
|
||||
return to_tensor(action, dtype=None) # Preserve original dtype
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("numpy2torch_teleop_action_processor")
|
||||
@dataclass
|
||||
class Numpy2TorchTeleopActionProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""Converts a NumPy teleop action in the complementary data to a PyTorch tensor."""
|
||||
|
||||
def complementary_data(self, complementary_data: dict) -> dict:
|
||||
if TELEOP_ACTION_KEY in complementary_data:
|
||||
teleop_action = complementary_data[TELEOP_ACTION_KEY]
|
||||
if isinstance(teleop_action, EnvAction):
|
||||
complementary_data[TELEOP_ACTION_KEY] = to_tensor(teleop_action)
|
||||
return complementary_data
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
|
||||
@@ -101,8 +101,8 @@ class AddTeleopActionAsComplimentaryDataStep(ComplementaryDataProcessorStep):
|
||||
be available to downstream processors, for example, to override a policy's action
|
||||
during an intervention.
|
||||
|
||||
**Attributes**:
|
||||
- **teleop_device** (`Teleoperator`) -- The teleoperator instance to get the action from.
|
||||
Attributes:
|
||||
teleop_device: The teleoperator instance to get the action from.
|
||||
"""
|
||||
|
||||
teleop_device: "Teleoperator"
|
||||
@@ -137,9 +137,9 @@ class AddTeleopEventsAsInfoStep(InfoProcessorStep):
|
||||
This step extracts control events from teleoperators that support event-based
|
||||
interaction, making these signals available to other parts of the system.
|
||||
|
||||
**Attributes**:
|
||||
- **teleop_device** (`TeleopWithEvents`) -- An instance of a teleoperator that implements the
|
||||
`HasTeleopEvents` protocol.
|
||||
Attributes:
|
||||
teleop_device: An instance of a teleoperator that implements the
|
||||
`HasTeleopEvents` protocol.
|
||||
"""
|
||||
|
||||
teleop_device: TeleopWithEvents
|
||||
@@ -180,10 +180,10 @@ class ImageCropResizeProcessorStep(ObservationProcessorStep):
|
||||
the specified transformations. It handles device placement, moving tensors to the
|
||||
CPU if necessary for operations not supported on certain accelerators like MPS.
|
||||
|
||||
**Attributes**:
|
||||
- **crop_params_dict** (`dict[str, tuple[int, int, int, int]] | None`) -- A dictionary mapping image
|
||||
keys to cropping parameters (top, left, height, width).
|
||||
- **resize_size** (`tuple[int, int] | None`) -- A tuple (height, width) to resize all images to.
|
||||
Attributes:
|
||||
crop_params_dict: A dictionary mapping image keys to cropping parameters
|
||||
(top, left, height, width).
|
||||
resize_size: A tuple (height, width) to resize all images to.
|
||||
"""
|
||||
|
||||
crop_params_dict: dict[str, tuple[int, int, int, int]] | None = None
|
||||
@@ -267,9 +267,9 @@ class TimeLimitProcessorStep(TruncatedProcessorStep):
|
||||
"""
|
||||
Tracks episode steps and enforces a time limit by truncating the episode.
|
||||
|
||||
**Attributes**:
|
||||
- **max_episode_steps** (`int`) -- The maximum number of steps allowed per episode.
|
||||
- **current_step** (`int`) -- The current step count for the active episode.
|
||||
Attributes:
|
||||
max_episode_steps: The maximum number of steps allowed per episode.
|
||||
current_step: The current step count for the active episode.
|
||||
"""
|
||||
|
||||
max_episode_steps: int
|
||||
@@ -312,20 +312,39 @@ class TimeLimitProcessorStep(TruncatedProcessorStep):
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("gym_hil_adapter_processor")
|
||||
class GymHILAdapterProcessorStep(ProcessorStep):
|
||||
@ProcessorStepRegistry.register("gym_hil_info_adapter")
|
||||
class GymHILInfoAdapterStep(InfoProcessorStep):
|
||||
"""
|
||||
Adapts the output of the `gym-hil` environment to the format expected by `lerobot` processors.
|
||||
Adapts the `info` dictionary of the `gym-hil` environment to the format expected by
|
||||
`lerobot` processors.
|
||||
|
||||
This step normalizes the `transition` object by:
|
||||
1. Copying `teleop_action` from `info` to `complementary_data`.
|
||||
2. Copying `is_intervention` from `info` (using the string key) to `info` (using the enum key).
|
||||
3. Copying `discrete_penalty` from `info` to `complementary_data`.
|
||||
Mirrors `is_intervention` from the string key to the `TeleopEvents.IS_INTERVENTION`
|
||||
enum key when present.
|
||||
"""
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
info = transition.get(TransitionKey.INFO, {})
|
||||
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
def info(self, info: dict) -> dict:
|
||||
if "is_intervention" in info:
|
||||
info[TeleopEvents.IS_INTERVENTION] = info["is_intervention"]
|
||||
return info
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("gym_hil_teleop_data_adapter")
|
||||
class GymHILTeleopDataAdapterStep(ComplementaryDataProcessorStep):
|
||||
"""
|
||||
Copies teleoperation data emitted by the `gym-hil` environment from `info` into the
|
||||
transition's complementary data.
|
||||
|
||||
Copies `teleop_action` and `discrete_penalty` from `info` to `complementary_data`
|
||||
when present.
|
||||
"""
|
||||
|
||||
def complementary_data(self, complementary_data: dict) -> dict:
|
||||
info = self.transition.get(TransitionKey.INFO) or {}
|
||||
|
||||
if TELEOP_ACTION_KEY in info:
|
||||
complementary_data[TELEOP_ACTION_KEY] = info[TELEOP_ACTION_KEY]
|
||||
@@ -333,13 +352,7 @@ class GymHILAdapterProcessorStep(ProcessorStep):
|
||||
if DISCRETE_PENALTY_KEY in info:
|
||||
complementary_data[DISCRETE_PENALTY_KEY] = info[DISCRETE_PENALTY_KEY]
|
||||
|
||||
if "is_intervention" in info:
|
||||
info[TeleopEvents.IS_INTERVENTION] = info["is_intervention"]
|
||||
|
||||
transition[TransitionKey.INFO] = info
|
||||
transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
||||
|
||||
return transition
|
||||
return complementary_data
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
@@ -349,7 +362,7 @@ class GymHILAdapterProcessorStep(ProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register("gripper_penalty_processor")
|
||||
class GripperPenaltyProcessorStep(ProcessorStep):
|
||||
class GripperPenaltyProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""
|
||||
Applies a small per-transition cost on the discrete gripper action.
|
||||
|
||||
@@ -358,11 +371,11 @@ class GripperPenaltyProcessorStep(ProcessorStep):
|
||||
This discourages gripper oscillation while leaving "stay" and saturating-further
|
||||
commands unpenalized.
|
||||
|
||||
**Attributes**:
|
||||
- **penalty** (`float`) -- The negative reward value to apply.
|
||||
- **max_gripper_pos** (`float`) -- The maximum position value for the gripper, used for normalization.
|
||||
- **open_threshold** (`float`) -- Normalized state below which the gripper is considered "open".
|
||||
- **closed_threshold** (`float`) -- Normalized state above which the gripper is considered "closed".
|
||||
Attributes:
|
||||
penalty: The negative reward value to apply.
|
||||
max_gripper_pos: The maximum position value for the gripper, used for normalization.
|
||||
open_threshold: Normalized state below which the gripper is considered "open".
|
||||
closed_threshold: Normalized state above which the gripper is considered "closed".
|
||||
"""
|
||||
|
||||
penalty: float = -0.02
|
||||
@@ -370,31 +383,30 @@ class GripperPenaltyProcessorStep(ProcessorStep):
|
||||
open_threshold: float = 0.1
|
||||
closed_threshold: float = 0.9
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
def complementary_data(self, complementary_data: dict) -> dict:
|
||||
"""
|
||||
Calculates the gripper penalty and adds it to the complementary data.
|
||||
|
||||
Args:
|
||||
transition: The incoming environment transition.
|
||||
complementary_data: The incoming complementary data dictionary.
|
||||
|
||||
Returns:
|
||||
The modified transition with the penalty added to complementary data.
|
||||
The complementary data with the penalty added under the
|
||||
`discrete_penalty` key.
|
||||
"""
|
||||
new_transition = transition.copy()
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
action = self.transition.get(TransitionKey.ACTION)
|
||||
|
||||
raw_joint_positions = complementary_data.get("raw_joint_positions")
|
||||
if raw_joint_positions is None:
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
current_gripper_pos = raw_joint_positions.get(f"{GRIPPER_KEY}.pos", None)
|
||||
if current_gripper_pos is None:
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
# During reset, the transition may not carry any action yet.
|
||||
if action is None:
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
# Gripper action is expected as the last action dimension.
|
||||
gripper_action = action[-1].item()
|
||||
@@ -414,12 +426,8 @@ class GripperPenaltyProcessorStep(ProcessorStep):
|
||||
|
||||
gripper_penalty = self.penalty * int(gripper_penalty_bool)
|
||||
|
||||
# Update complementary data with penalty info
|
||||
new_complementary_data = dict(complementary_data)
|
||||
new_complementary_data[DISCRETE_PENALTY_KEY] = gripper_penalty
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
|
||||
return new_transition
|
||||
complementary_data[DISCRETE_PENALTY_KEY] = gripper_penalty
|
||||
return complementary_data
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -436,10 +444,6 @@ class GripperPenaltyProcessorStep(ProcessorStep):
|
||||
"closed_threshold": self.closed_threshold,
|
||||
}
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the processor's internal state."""
|
||||
pass
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
@@ -456,10 +460,10 @@ class InterventionActionProcessorStep(ProcessorStep):
|
||||
this step replaces the policy's action with the human's teleoperated action.
|
||||
It also processes signals to terminate the episode or flag success.
|
||||
|
||||
**Attributes**:
|
||||
- **use_gripper** (`bool`) -- Whether to include the gripper in the teleoperated action.
|
||||
- **terminate_on_success** (`bool`) -- If True, automatically sets the `done` flag when a `success`
|
||||
event is received.
|
||||
Attributes:
|
||||
use_gripper: Whether to include the gripper in the teleoperated action.
|
||||
terminate_on_success: If True, automatically sets the `done` flag when a
|
||||
`success` event is received.
|
||||
"""
|
||||
|
||||
use_gripper: bool = False
|
||||
@@ -557,13 +561,13 @@ class RewardClassifierProcessorStep(ProcessorStep):
|
||||
This step uses a model to determine if the current state is successful, updating
|
||||
the reward and potentially terminating the episode.
|
||||
|
||||
**Attributes**:
|
||||
- **pretrained_path** (`str | None`) -- Path to the pretrained reward classifier model.
|
||||
- **device** (`str`) -- The device to run the classifier on.
|
||||
- **success_threshold** (`float`) -- The probability threshold to consider a prediction as successful.
|
||||
- **success_reward** (`float`) -- The reward value to assign on success.
|
||||
- **terminate_on_success** (`bool`) -- If True, terminates the episode upon successful classification.
|
||||
- **reward_classifier** (`Any`) -- The loaded classifier model instance.
|
||||
Attributes:
|
||||
pretrained_path: Path to the pretrained reward classifier model.
|
||||
device: The device to run the classifier on.
|
||||
success_threshold: The probability threshold to consider a prediction as successful.
|
||||
success_reward: The reward value to assign on success.
|
||||
terminate_on_success: If True, terminates the episode upon successful classification.
|
||||
reward_classifier: The loaded classifier model instance.
|
||||
"""
|
||||
|
||||
pretrained_path: str | None = None
|
||||
|
||||
@@ -647,15 +647,10 @@ def main():
|
||||
tags = set(tags).union({"robotics", "lerobot", policy_type})
|
||||
tags = list(tags)
|
||||
|
||||
# Generate model card through the free helper (PreTrainedPolicy.generate_model_card was
|
||||
# removed with the publisher redesign), then apply the metadata recovered above — the
|
||||
# migrated policy config does not carry the original repo's card fields.
|
||||
from lerobot.common.train_utils import generate_model_card
|
||||
|
||||
card = generate_model_card(policy.config)
|
||||
card.data.datasets = dataset_repo_id
|
||||
card.data.license = license
|
||||
card.data.tags = sorted(tags)
|
||||
# Generate model card
|
||||
card = policy.generate_model_card(
|
||||
dataset_repo_id=dataset_repo_id, model_type=policy_type, license=license, tags=tags
|
||||
)
|
||||
|
||||
# Save model card locally
|
||||
card.save(str(output_dir / "README.md"))
|
||||
|
||||
@@ -71,23 +71,22 @@ class _NormalizationMixin:
|
||||
)
|
||||
```
|
||||
|
||||
**Attributes**:
|
||||
- **features** (`dict[str, PolicyFeature]`) -- A dictionary mapping feature names to `PolicyFeature`
|
||||
objects, defining the data structure to be processed.
|
||||
- **norm_map** (`dict[FeatureType, NormalizationMode]`) -- A dictionary mapping `FeatureType` to
|
||||
`NormalizationMode`, specifying which normalization method to use for each type of feature.
|
||||
- **stats** (`dict[str, dict[str, Any]] | None`) -- A dictionary containing the normalization
|
||||
statistics (e.g., mean, std, min, max) for each feature.
|
||||
- **device** (`torch.device | str | None`) -- The PyTorch device on which to store and perform tensor
|
||||
operations.
|
||||
- **eps** (`float`) -- A small epsilon value to prevent division by zero in normalization
|
||||
calculations.
|
||||
- **normalize_observation_keys** (`set[str] | None`) -- An optional set of keys to selectively apply
|
||||
normalization to specific observation features.
|
||||
- **_tensor_stats** (`dict[str, dict[str, Tensor]]`) -- An internal dictionary holding the
|
||||
normalization statistics as PyTorch tensors.
|
||||
- **_stats_explicitly_provided** (`bool`) -- Internal flag tracking whether stats were explicitly
|
||||
provided during construction (used for override preservation).
|
||||
Attributes:
|
||||
features: A dictionary mapping feature names to `PolicyFeature` objects, defining
|
||||
the data structure to be processed.
|
||||
norm_map: A dictionary mapping `FeatureType` to `NormalizationMode`, specifying
|
||||
which normalization method to use for each type of feature.
|
||||
stats: A dictionary containing the normalization statistics (e.g., mean, std,
|
||||
min, max) for each feature.
|
||||
device: The PyTorch device on which to store and perform tensor operations.
|
||||
eps: A small epsilon value to prevent division by zero in normalization
|
||||
calculations.
|
||||
normalize_observation_keys: An optional set of keys to selectively apply
|
||||
normalization to specific observation features.
|
||||
_tensor_stats: An internal dictionary holding the normalization statistics as
|
||||
PyTorch tensors.
|
||||
_stats_explicitly_provided: Internal flag tracking whether stats were explicitly
|
||||
provided during construction (used for override preservation).
|
||||
"""
|
||||
|
||||
features: dict[str, PolicyFeature]
|
||||
|
||||
@@ -38,10 +38,10 @@ from collections.abc import Callable, Iterable, Sequence
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, TypedDict, TypeVar, cast
|
||||
from typing import Any, ClassVar, TypedDict, TypeVar, cast
|
||||
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download, snapshot_download
|
||||
from huggingface_hub import hf_hub_download
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
@@ -159,6 +159,14 @@ class ProcessorStep(ABC):
|
||||
|
||||
_current_transition: EnvTransition | None = None
|
||||
|
||||
# Consulted by the specialized single-field bases (ObservationProcessorStep, ActionProcessorStep,
|
||||
# etc.): when True, the step is skipped (the transition is returned unchanged) if its target field
|
||||
# is None, instead of raising a ValueError. Set it as a plain class attribute in subclasses
|
||||
# (`skip_if_missing = True`) so that dataclass steps don't pick it up as a field. Use it for steps
|
||||
# that must tolerate partial transitions, e.g. action steps in a preprocessor that also runs at
|
||||
# inference time (where the action is None) or steps in RL pipelines that run on reset transitions.
|
||||
skip_if_missing: ClassVar[bool] = False
|
||||
|
||||
@property
|
||||
def transition(self) -> EnvTransition:
|
||||
"""Provides access to the most recent transition being processed.
|
||||
@@ -212,10 +220,6 @@ class ProcessorStep(ABC):
|
||||
"""
|
||||
return None
|
||||
|
||||
def save_artifacts(self, save_directory: Path) -> dict[str, str]:
|
||||
"""Save non-tensor assets and map constructor arguments to relative paths."""
|
||||
return {}
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the processor step, if any."""
|
||||
return None
|
||||
@@ -269,18 +273,13 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
data processing workflow. It's generic, allowing for custom input and output types,
|
||||
which are handled by the `to_transition` and `to_output` converters.
|
||||
|
||||
**Attributes**:
|
||||
- **steps** (`Sequence[ProcessorStep]`) -- A sequence of `ProcessorStep` objects that make up the
|
||||
pipeline.
|
||||
- **name** (`str`) -- A descriptive name for the pipeline.
|
||||
- **to_transition** (`Callable[[TInput], EnvTransition]`) -- A function to convert raw input data into
|
||||
the standardized `EnvTransition` format.
|
||||
- **to_output** (`Callable[[EnvTransition], TOutput]`) -- A function to convert the final
|
||||
`EnvTransition` into the desired output format.
|
||||
- **before_step_hooks** (`list[Callable[[int, EnvTransition], None]]`) -- A list of functions to be
|
||||
called before each step is executed.
|
||||
- **after_step_hooks** (`list[Callable[[int, EnvTransition], None]]`) -- A list of functions to be
|
||||
called after each step is executed.
|
||||
Attributes:
|
||||
steps: A sequence of `ProcessorStep` objects that make up the pipeline.
|
||||
name: A descriptive name for the pipeline.
|
||||
to_transition: A function to convert raw input data into the standardized `EnvTransition` format.
|
||||
to_output: A function to convert the final `EnvTransition` into the desired output format.
|
||||
before_step_hooks: A list of functions to be called before each step is executed.
|
||||
after_step_hooks: A list of functions to be called after each step is executed.
|
||||
"""
|
||||
|
||||
steps: Sequence[ProcessorStep] = field(default_factory=list)
|
||||
@@ -565,22 +564,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
pipeline_config = self.get_config()
|
||||
pipeline_state_dict = self.state_dict()
|
||||
|
||||
for processor_step, step_entry in zip(self.steps, pipeline_config["steps"], strict=True):
|
||||
artifacts = processor_step.save_artifacts(save_directory)
|
||||
if artifacts:
|
||||
for config_key, relative_path in artifacts.items():
|
||||
artifact_path = Path(relative_path)
|
||||
if artifact_path.is_absolute() or ".." in artifact_path.parts:
|
||||
raise ValueError(
|
||||
f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
|
||||
)
|
||||
if not (save_directory / artifact_path).exists():
|
||||
raise FileNotFoundError(
|
||||
f"Processor step did not save declared artifact '{relative_path}'"
|
||||
)
|
||||
step_entry["config"][config_key] = artifact_path.as_posix()
|
||||
step_entry["artifacts"] = artifacts
|
||||
|
||||
for state_key, step_state_dict in pipeline_state_dict.items():
|
||||
state_filename = f"{state_key}.safetensors"
|
||||
save_file(step_state_dict, save_directory / state_filename)
|
||||
@@ -765,13 +748,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
|
||||
# 3. Build steps with overrides
|
||||
steps, validated_overrides = cls._build_steps_with_overrides(
|
||||
loaded_config,
|
||||
overrides or {},
|
||||
model_id,
|
||||
base_path,
|
||||
config_filename,
|
||||
hub_download_kwargs,
|
||||
is_local_source,
|
||||
loaded_config, overrides or {}, model_id, base_path, hub_download_kwargs, is_local_source
|
||||
)
|
||||
|
||||
# 4. Validate that all overrides were used
|
||||
@@ -967,7 +944,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
overrides: dict[str, Any],
|
||||
model_id: str,
|
||||
base_path: Path | None,
|
||||
config_filename: str,
|
||||
hub_download_kwargs: dict[str, Any],
|
||||
is_local_source: bool = False,
|
||||
) -> tuple[list[ProcessorStep], set[str]]:
|
||||
@@ -977,11 +953,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
|
||||
**For each step in loaded_config["steps"]**:
|
||||
|
||||
0. **Artifact Resolution** (via _resolve_artifact_paths):
|
||||
- Resolve declared relative artifact paths against a local checkpoint
|
||||
- Download declared artifacts when loading the pipeline from the Hub
|
||||
- Reject absolute paths and path traversal before step construction
|
||||
|
||||
1. **Class Resolution** (via _resolve_step_class):
|
||||
- **If "registry_name" exists**: Look up in ProcessorStepRegistry
|
||||
Example: {"registry_name": "normalize_step"} -> Get registered class
|
||||
@@ -1015,8 +986,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
overrides: User-provided parameter overrides (keyed by class/registry name)
|
||||
model_id: The model identifier (needed for Hub state file downloads)
|
||||
base_path: Local directory path for finding state files
|
||||
config_filename: Processor config path, used as the repository-relative
|
||||
base for state files and declared artifacts.
|
||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||
is_local_source: Whether model_id resolved to a local directory or config file.
|
||||
|
||||
@@ -1029,80 +998,15 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
ImportError: If a step class cannot be imported or found in registry
|
||||
ValueError: If a step cannot be instantiated with its configuration
|
||||
"""
|
||||
loaded_config = deepcopy(loaded_config)
|
||||
cls._resolve_artifact_paths(
|
||||
loaded_config,
|
||||
model_id,
|
||||
base_path,
|
||||
config_filename,
|
||||
hub_download_kwargs,
|
||||
)
|
||||
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
|
||||
|
||||
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
|
||||
cls._load_step_state(
|
||||
step_instance,
|
||||
step_entry,
|
||||
model_id,
|
||||
base_path,
|
||||
config_filename,
|
||||
hub_download_kwargs,
|
||||
is_local_source,
|
||||
step_instance, step_entry, model_id, base_path, hub_download_kwargs, is_local_source
|
||||
)
|
||||
|
||||
return steps, remaining_override_keys
|
||||
|
||||
@classmethod
|
||||
def _resolve_artifact_paths(
|
||||
cls,
|
||||
loaded_config: dict[str, Any],
|
||||
model_id: str,
|
||||
base_path: Path | None,
|
||||
config_filename: str,
|
||||
hub_download_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Resolve declared relative processor artifacts before step construction.
|
||||
|
||||
Args:
|
||||
loaded_config: Mutable processor configuration containing step artifact declarations.
|
||||
model_id: Local checkpoint path or Hub model identifier.
|
||||
base_path: Local directory containing the resolved processor configuration.
|
||||
config_filename: Processor config path, whose parent is the artifact root on the Hub.
|
||||
hub_download_kwargs: Authentication, revision, and cache arguments for Hub downloads.
|
||||
|
||||
Raises:
|
||||
ValueError: If a declared artifact path is absolute or escapes the checkpoint.
|
||||
FileNotFoundError: If a declared artifact cannot be found locally or downloaded.
|
||||
"""
|
||||
is_local = Path(model_id).is_dir() or Path(model_id).is_file()
|
||||
|
||||
for step_entry in loaded_config["steps"]:
|
||||
artifacts = step_entry.get("artifacts", {})
|
||||
for config_key, relative_path in artifacts.items():
|
||||
artifact_path = Path(relative_path)
|
||||
if artifact_path.is_absolute() or ".." in artifact_path.parts:
|
||||
raise ValueError(
|
||||
f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
|
||||
)
|
||||
|
||||
resolved_path = base_path / artifact_path if base_path is not None else artifact_path
|
||||
if not resolved_path.exists() and not is_local:
|
||||
repository_path = Path(config_filename).parent / artifact_path
|
||||
snapshot_download(
|
||||
repo_id=model_id,
|
||||
repo_type="model",
|
||||
allow_patterns=f"{repository_path.as_posix()}/**",
|
||||
**hub_download_kwargs,
|
||||
)
|
||||
|
||||
if not resolved_path.exists():
|
||||
step_name = step_entry.get("registry_name", step_entry.get("class", "unknown"))
|
||||
raise FileNotFoundError(
|
||||
f"Missing processor artifact '{relative_path}' for step '{step_name}' "
|
||||
f"next to '{config_filename}'. Checkpoint artifacts are incomplete."
|
||||
)
|
||||
step_entry["config"][config_key] = str(resolved_path)
|
||||
|
||||
@classmethod
|
||||
def _build_steps_from_config(
|
||||
cls,
|
||||
@@ -1262,7 +1166,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
step_entry: dict[str, Any],
|
||||
model_id: str,
|
||||
base_path: Path | None,
|
||||
config_filename: str,
|
||||
hub_download_kwargs: dict[str, Any],
|
||||
is_local_source: bool = False,
|
||||
) -> None:
|
||||
@@ -1303,8 +1206,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
step_entry: The step configuration dictionary (may contain "state_file")
|
||||
model_id: The model identifier (used for Hub downloads if needed)
|
||||
base_path: Local directory path for finding state files (None for Hub-only)
|
||||
config_filename: Processor config path, whose parent is used to resolve
|
||||
repository-relative state files on the Hub.
|
||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||
is_local_source: Whether model_id resolved to a local directory or config file.
|
||||
|
||||
@@ -1330,7 +1231,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
||||
# Download from Hub
|
||||
state_path = hf_hub_download(
|
||||
repo_id=model_id,
|
||||
filename=(Path(config_filename).parent / state_filename).as_posix(),
|
||||
filename=state_filename,
|
||||
repo_type="model",
|
||||
**hub_download_kwargs,
|
||||
)
|
||||
@@ -1860,7 +1761,12 @@ PolicyProcessorPipeline = DataProcessorPipeline[TInput, TOutput]
|
||||
|
||||
|
||||
class ObservationProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the observation in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the observation in a transition.
|
||||
|
||||
The `observation` hook may read other parts of the transition via `self.transition`, but only the
|
||||
observation may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of
|
||||
raising) when the transition carries no observation.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def observation(self, observation: RobotObservation) -> RobotObservation:
|
||||
@@ -1880,6 +1786,8 @@ class ObservationProcessorStep(ProcessorStep, ABC):
|
||||
new_transition = self._current_transition
|
||||
|
||||
observation = new_transition.get(TransitionKey.OBSERVATION)
|
||||
if observation is None and self.skip_if_missing:
|
||||
return new_transition
|
||||
if observation is None or not isinstance(observation, dict):
|
||||
raise ValueError("ObservationProcessorStep requires an observation in the transition.")
|
||||
|
||||
@@ -1889,7 +1797,12 @@ class ObservationProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class ActionProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the action in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the action in a transition.
|
||||
|
||||
The `action` hook may read other parts of the transition via `self.transition`, but only the action
|
||||
may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising)
|
||||
when the transition carries no action, e.g. for steps in pipelines that also run at inference time.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def action(
|
||||
@@ -1912,6 +1825,8 @@ class ActionProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
if self.skip_if_missing:
|
||||
return new_transition
|
||||
raise ValueError("ActionProcessorStep requires an action in the transition.")
|
||||
|
||||
processed_action = self.action(action)
|
||||
@@ -1920,7 +1835,12 @@ class ActionProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class RobotActionProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` for processing a `RobotAction` (a dictionary)."""
|
||||
"""An abstract `ProcessorStep` for processing a `RobotAction` (a dictionary).
|
||||
|
||||
The `action` hook may read other parts of the transition via `self.transition`, but only the action
|
||||
may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising)
|
||||
when the transition carries no action.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
@@ -1940,6 +1860,8 @@ class RobotActionProcessorStep(ProcessorStep, ABC):
|
||||
new_transition = self._current_transition
|
||||
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None and self.skip_if_missing:
|
||||
return new_transition
|
||||
if action is None or not isinstance(action, dict):
|
||||
raise ValueError(f"Action should be a RobotAction type (dict), but got {type(action)}")
|
||||
|
||||
@@ -1949,7 +1871,12 @@ class RobotActionProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class PolicyActionProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` for processing a `PolicyAction` (a tensor or dict of tensors)."""
|
||||
"""An abstract `ProcessorStep` for processing a `PolicyAction` (a tensor).
|
||||
|
||||
The `action` hook may read other parts of the transition via `self.transition`, but only the action
|
||||
may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising)
|
||||
when the transition carries no action, e.g. for steps in pipelines that also run at inference time.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
@@ -1969,6 +1896,8 @@ class PolicyActionProcessorStep(ProcessorStep, ABC):
|
||||
new_transition = self._current_transition
|
||||
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None and self.skip_if_missing:
|
||||
return new_transition
|
||||
if not isinstance(action, PolicyAction):
|
||||
raise ValueError(f"Action should be a PolicyAction type (tensor), but got {type(action)}")
|
||||
|
||||
@@ -1978,7 +1907,11 @@ class PolicyActionProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class RewardProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the reward in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the reward in a transition.
|
||||
|
||||
Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the
|
||||
transition carries no reward.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def reward(self, reward) -> float | torch.Tensor:
|
||||
@@ -1999,6 +1932,8 @@ class RewardProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
reward = new_transition.get(TransitionKey.REWARD)
|
||||
if reward is None:
|
||||
if self.skip_if_missing:
|
||||
return new_transition
|
||||
raise ValueError("RewardProcessorStep requires a reward in the transition.")
|
||||
|
||||
processed_reward = self.reward(reward)
|
||||
@@ -2007,7 +1942,11 @@ class RewardProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class DoneProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'done' flag in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'done' flag in a transition.
|
||||
|
||||
Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the
|
||||
transition carries no 'done' flag.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def done(self, done) -> bool | torch.Tensor:
|
||||
@@ -2028,6 +1967,8 @@ class DoneProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
done = new_transition.get(TransitionKey.DONE)
|
||||
if done is None:
|
||||
if self.skip_if_missing:
|
||||
return new_transition
|
||||
raise ValueError("DoneProcessorStep requires a done flag in the transition.")
|
||||
|
||||
processed_done = self.done(done)
|
||||
@@ -2036,7 +1977,11 @@ class DoneProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class TruncatedProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'truncated' flag in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'truncated' flag in a transition.
|
||||
|
||||
Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the
|
||||
transition carries no 'truncated' flag.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def truncated(self, truncated) -> bool | torch.Tensor:
|
||||
@@ -2057,6 +2002,8 @@ class TruncatedProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
truncated = new_transition.get(TransitionKey.TRUNCATED)
|
||||
if truncated is None:
|
||||
if self.skip_if_missing:
|
||||
return new_transition
|
||||
raise ValueError("TruncatedProcessorStep requires a truncated flag in the transition.")
|
||||
|
||||
processed_truncated = self.truncated(truncated)
|
||||
@@ -2065,7 +2012,11 @@ class TruncatedProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class InfoProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'info' dictionary in a transition."""
|
||||
"""An abstract `ProcessorStep` that specifically targets the 'info' dictionary in a transition.
|
||||
|
||||
The `info` hook may read other parts of the transition via `self.transition`, but only the info
|
||||
dictionary may be written.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def info(self, info) -> dict[str, Any]:
|
||||
@@ -2085,6 +2036,8 @@ class InfoProcessorStep(ProcessorStep, ABC):
|
||||
new_transition = self._current_transition
|
||||
|
||||
info = new_transition.get(TransitionKey.INFO)
|
||||
if info is None and self.skip_if_missing:
|
||||
return new_transition
|
||||
if info is None or not isinstance(info, dict):
|
||||
raise ValueError("InfoProcessorStep requires an info dictionary in the transition.")
|
||||
|
||||
@@ -2094,7 +2047,11 @@ class InfoProcessorStep(ProcessorStep, ABC):
|
||||
|
||||
|
||||
class ComplementaryDataProcessorStep(ProcessorStep, ABC):
|
||||
"""An abstract `ProcessorStep` that targets the 'complementary_data' in a transition."""
|
||||
"""An abstract `ProcessorStep` that targets the 'complementary_data' in a transition.
|
||||
|
||||
The `complementary_data` hook may read other parts of the transition via `self.transition` (e.g. an
|
||||
action or observation the step derives data from), but only the complementary data may be written.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def complementary_data(self, complementary_data) -> dict[str, Any]:
|
||||
@@ -2114,6 +2071,8 @@ class ComplementaryDataProcessorStep(ProcessorStep, ABC):
|
||||
new_transition = self._current_transition
|
||||
|
||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA)
|
||||
if complementary_data is None and self.skip_if_missing:
|
||||
return new_transition
|
||||
if complementary_data is None or not isinstance(complementary_data, dict):
|
||||
raise ValueError("ComplementaryDataProcessorStep requires complementary data in the transition.")
|
||||
|
||||
|
||||
@@ -20,11 +20,11 @@ import torch
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import EnvTransition, PolicyAction, TransitionKey
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
|
||||
from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep
|
||||
from .pipeline import ProcessorStep, ProcessorStepRegistry
|
||||
from .pipeline import PolicyActionProcessorStep, ProcessorStep, ProcessorStepRegistry
|
||||
|
||||
# Re-export for backward compatibility
|
||||
__all__ = [
|
||||
@@ -91,11 +91,11 @@ class RelativeActionsProcessorStep(ProcessorStep):
|
||||
Caches the last seen state so a paired AbsoluteActionsProcessorStep can reverse
|
||||
the conversion during postprocessing.
|
||||
|
||||
**Attributes**:
|
||||
- **enabled** (`bool`) -- Whether to apply the relative conversion.
|
||||
- **exclude_joints** (`list[str]`) -- Joint names to keep absolute (not converted to relative).
|
||||
- **action_names** (`list[str] | None`) -- Action dimension names from dataset metadata, used to build
|
||||
the mask from exclude_joints. If None, all dims are converted.
|
||||
Attributes:
|
||||
enabled: Whether to apply the relative conversion.
|
||||
exclude_joints: Joint names to keep absolute (not converted to relative).
|
||||
action_names: Action dimension names from dataset metadata, used to build
|
||||
the mask from exclude_joints. If None, all dims are converted.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
@@ -161,25 +161,26 @@ class RelativeActionsProcessorStep(ProcessorStep):
|
||||
|
||||
@ProcessorStepRegistry.register("absolute_actions_processor")
|
||||
@dataclass
|
||||
class AbsoluteActionsProcessorStep(ProcessorStep):
|
||||
class AbsoluteActionsProcessorStep(PolicyActionProcessorStep):
|
||||
"""Converts relative actions back to absolute actions (action += state) for all dimensions.
|
||||
|
||||
Mirrors OpenPI's AbsoluteActions transform. Applied during postprocessing so
|
||||
predicted relative offsets are converted back to absolute positions for execution.
|
||||
Reads the cached state from its paired RelativeActionsProcessorStep.
|
||||
|
||||
**Attributes**:
|
||||
- **enabled** (`bool`) -- Whether to apply the absolute conversion.
|
||||
- **relative_step** (`RelativeActionsProcessorStep | None`) -- Reference to the paired
|
||||
RelativeActionsProcessorStep that caches state.
|
||||
Attributes:
|
||||
enabled: Whether to apply the absolute conversion.
|
||||
relative_step: Reference to the paired RelativeActionsProcessorStep that caches state.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
relative_step: RelativeActionsProcessorStep | None = field(default=None, repr=False)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
skip_if_missing = True
|
||||
|
||||
def action(self, action: PolicyAction) -> PolicyAction:
|
||||
if not self.enabled:
|
||||
return transition
|
||||
return action
|
||||
|
||||
if self.relative_step is None:
|
||||
raise RuntimeError(
|
||||
@@ -194,14 +195,8 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
|
||||
"but no state has been cached. Ensure the preprocessor runs before the postprocessor."
|
||||
)
|
||||
|
||||
new_transition = transition.copy()
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
return new_transition
|
||||
|
||||
mask = self.relative_step._build_mask(action.shape[-1])
|
||||
new_transition[TransitionKey.ACTION] = to_absolute_actions(action, cached_state, mask)
|
||||
return new_transition
|
||||
return to_absolute_actions(action, cached_state, mask)
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"enabled": self.enabled}
|
||||
|
||||
@@ -32,9 +32,10 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
|
||||
from an environment's format to the format expected by a LeRobot policy or
|
||||
other downstream components.
|
||||
|
||||
**Attributes**:
|
||||
- **rename_map** (`dict[str, str]`) -- A dictionary mapping from old key names to new key names. Keys
|
||||
present in an observation that are not in this map will be kept with their original names.
|
||||
Attributes:
|
||||
rename_map: A dictionary mapping from old key names to new key names.
|
||||
Keys present in an observation that are not in this map will
|
||||
be kept with their original names.
|
||||
"""
|
||||
|
||||
rename_map: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
@@ -16,11 +16,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.configs.recipe import TrainingRecipe
|
||||
from lerobot.datasets.language import LANGUAGE_EVENTS, LANGUAGE_PERSISTENT
|
||||
@@ -34,46 +32,25 @@ from .pipeline import ProcessorStep, ProcessorStepRegistry
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="render_messages_processor")
|
||||
class RenderMessagesStep(ProcessorStep):
|
||||
"""Turn raw language columns into recipe-defined messages and supervision.
|
||||
"""Processor step that turns raw language columns into rendered chat messages.
|
||||
|
||||
Reads ``language_persistent`` and ``language_events`` from complementary
|
||||
data, renders them at each sample timestamp, and replaces the raw columns
|
||||
with ``messages``, ``message_streams``, and ``target_message_indices``.
|
||||
Batched inputs are filtered to samples with applicable supervision; samples
|
||||
without language annotations use their task string as low-level supervision
|
||||
when one is available.
|
||||
Reads ``language_persistent`` and ``language_events`` from the transition's
|
||||
complementary data, renders them through ``recipe`` at the sample timestamp,
|
||||
and replaces the raw columns with the resulting ``messages`` /
|
||||
``message_streams`` / ``target_message_indices`` keys.
|
||||
"""
|
||||
|
||||
recipe: TrainingRecipe
|
||||
dataset_ctx: Any | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if isinstance(self.recipe, dict):
|
||||
self.recipe = TrainingRecipe.from_dict(self.recipe)
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"recipe": asdict(self.recipe)}
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition | None:
|
||||
"""Render messages, preserving unannotated samples and dropping unmatched annotated ones."""
|
||||
"""Render messages for a single transition; return ``None`` to drop it."""
|
||||
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
persistent = complementary_data.get(LANGUAGE_PERSISTENT) or []
|
||||
events = complementary_data.get(LANGUAGE_EVENTS) or []
|
||||
|
||||
if not persistent and not events:
|
||||
# A dataset without language annotations remains usable: render its
|
||||
# task as low-level supervision, or pass it through when no task exists.
|
||||
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||
if rendered is None:
|
||||
return transition
|
||||
new_transition = transition.copy()
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data.update(rendered)
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
if _is_batched_language(persistent) or _is_batched_language(events):
|
||||
return self._call_batch(transition, complementary_data, persistent, events)
|
||||
return transition
|
||||
|
||||
timestamp = complementary_data.get("timestamp")
|
||||
if timestamp is None:
|
||||
@@ -90,171 +67,18 @@ class RenderMessagesStep(ProcessorStep):
|
||||
dataset_ctx=self.dataset_ctx,
|
||||
)
|
||||
if rendered is None:
|
||||
# Language is present but this sparse frame has no applicable recipe
|
||||
# branch. Keep it only when task-level action supervision is possible.
|
||||
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||
if rendered is None:
|
||||
return None
|
||||
return None
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data = dict(complementary_data)
|
||||
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||
new_complementary_data.update(rendered)
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
def _call_batch(
|
||||
self,
|
||||
transition: EnvTransition,
|
||||
complementary_data: dict[str, Any],
|
||||
persistent_batch: list,
|
||||
events_batch: list,
|
||||
) -> EnvTransition | None:
|
||||
"""Render a language batch.
|
||||
|
||||
Non-empty persistent and event batches must have the same size. Either
|
||||
list may be empty when that language column is absent from the batch.
|
||||
"""
|
||||
timestamp = complementary_data.get("timestamp")
|
||||
if timestamp is None:
|
||||
raise KeyError("RenderMessagesStep requires sample timestamp in complementary data.")
|
||||
|
||||
non_empty_batch_sizes = {len(batch) for batch in (persistent_batch, events_batch) if batch}
|
||||
if len(non_empty_batch_sizes) > 1:
|
||||
raise ValueError(
|
||||
"Batched language columns must have equal lengths when both are non-empty, "
|
||||
f"got persistent={len(persistent_batch)} and events={len(events_batch)}."
|
||||
)
|
||||
batch_size = next(iter(non_empty_batch_sizes), 0)
|
||||
messages: list[list[dict[str, Any]]] = []
|
||||
message_streams: list[list[str | None]] = []
|
||||
target_message_indices: list[list[int]] = []
|
||||
keep_indices: list[int] = []
|
||||
|
||||
for i in range(batch_size):
|
||||
rendered = render_sample(
|
||||
recipe=self.recipe,
|
||||
persistent=persistent_batch[i] if i < len(persistent_batch) else [],
|
||||
events=events_batch[i] if i < len(events_batch) else [],
|
||||
t=_batch_value(timestamp, i),
|
||||
sample_idx=int(_batch_value(complementary_data.get("index", 0), i)),
|
||||
task=_batch_value(complementary_data.get("task"), i),
|
||||
dataset_ctx=self.dataset_ctx,
|
||||
)
|
||||
if rendered is None:
|
||||
rendered = _fallback_low_level_render(_batch_value(complementary_data.get("task"), i))
|
||||
if rendered is None:
|
||||
continue
|
||||
keep_indices.append(i)
|
||||
messages.append(rendered["messages"])
|
||||
message_streams.append(rendered["message_streams"])
|
||||
target_message_indices.append(rendered["target_message_indices"])
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
new_transition = (
|
||||
_select_batch_indices(transition, keep_indices, batch_size)
|
||||
if len(keep_indices) != batch_size
|
||||
else transition.copy()
|
||||
)
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||
new_complementary_data["messages"] = messages
|
||||
new_complementary_data["message_streams"] = message_streams
|
||||
new_complementary_data["target_message_indices"] = target_message_indices
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Pass features through unchanged; rendering only touches complementary data."""
|
||||
return features
|
||||
|
||||
|
||||
def _is_batched_language(value: Any) -> bool:
|
||||
return isinstance(value, list) and bool(value) and isinstance(value[0], list)
|
||||
|
||||
|
||||
def _batch_value(value: Any, index: int) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, list):
|
||||
return value[index]
|
||||
if hasattr(value, "ndim") and value.ndim > 0:
|
||||
return unwrap_scalar(value[index])
|
||||
return unwrap_scalar(value)
|
||||
|
||||
|
||||
def _select_batch_indices(transition: EnvTransition, indices: list[int], batch_size: int) -> EnvTransition:
|
||||
selected = transition.copy()
|
||||
for key in (TransitionKey.OBSERVATION, TransitionKey.COMPLEMENTARY_DATA):
|
||||
data = selected.get(key)
|
||||
if isinstance(data, dict):
|
||||
selected[key] = {
|
||||
name: _select_value(value, indices, batch_size, f"{key}.{name}")
|
||||
for name, value in data.items()
|
||||
}
|
||||
action = selected.get(TransitionKey.ACTION)
|
||||
if action is not None:
|
||||
selected[TransitionKey.ACTION] = _select_value(action, indices, batch_size, str(TransitionKey.ACTION))
|
||||
return selected
|
||||
|
||||
|
||||
def _select_value(value: Any, indices: list[int], batch_size: int, path: str) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {key: _select_value(item, indices, batch_size, f"{path}.{key}") for key, item in value.items()}
|
||||
if isinstance(value, list):
|
||||
if len(value) != batch_size:
|
||||
raise ValueError(
|
||||
f"Cannot filter batched field {path!r}: expected {batch_size} values, got {len(value)}."
|
||||
)
|
||||
return [value[i] for i in indices]
|
||||
if isinstance(value, np.ndarray) and value.ndim > 0:
|
||||
return value[indices]
|
||||
if hasattr(value, "index_select") and hasattr(value, "new_tensor") and getattr(value, "ndim", 0) > 0:
|
||||
return value.index_select(0, value.new_tensor(indices).long())
|
||||
return value
|
||||
|
||||
|
||||
def _fallback_low_level_render(task: Any) -> dict[str, Any] | None:
|
||||
"""Keep action-only samples trainable when no recipe branch matches."""
|
||||
if hasattr(task, "item"):
|
||||
task = task.item()
|
||||
if isinstance(task, list):
|
||||
if not task:
|
||||
return None
|
||||
messages = []
|
||||
message_streams = []
|
||||
target_message_indices = []
|
||||
missing_indices = []
|
||||
for index, t in enumerate(task):
|
||||
rendered = _fallback_low_level_render(t)
|
||||
if rendered is None:
|
||||
missing_indices.append(index)
|
||||
continue
|
||||
messages.append(rendered["messages"])
|
||||
message_streams.append(rendered["message_streams"])
|
||||
target_message_indices.append(rendered["target_message_indices"])
|
||||
if missing_indices:
|
||||
if len(missing_indices) == len(task):
|
||||
return None
|
||||
raise ValueError(
|
||||
"Batched low-level fallback requires a non-empty task for every sample; "
|
||||
f"missing task at indices {missing_indices}."
|
||||
)
|
||||
return {
|
||||
"messages": messages,
|
||||
"message_streams": message_streams,
|
||||
"target_message_indices": target_message_indices,
|
||||
}
|
||||
if not isinstance(task, str) or not task:
|
||||
return None
|
||||
return {
|
||||
"messages": [{"role": "user", "content": task}],
|
||||
"message_streams": ["low_level"],
|
||||
"target_message_indices": [],
|
||||
}
|
||||
|
||||
@@ -25,7 +25,6 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
@@ -33,7 +32,6 @@ import torch
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, RobotObservation, TransitionKey
|
||||
from lerobot.utils.constants import (
|
||||
ACTION_CODE_TOKEN_MASK,
|
||||
ACTION_TOKEN_MASK,
|
||||
ACTION_TOKENS,
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
@@ -43,7 +41,7 @@ from lerobot.utils.constants import (
|
||||
)
|
||||
from lerobot.utils.import_utils import _transformers_available
|
||||
|
||||
from .pipeline import ActionProcessorStep, ObservationProcessorStep, ProcessorStepRegistry
|
||||
from .pipeline import ComplementaryDataProcessorStep, ObservationProcessorStep, ProcessorStepRegistry
|
||||
|
||||
# Conditional import for type checking and lazy loading
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
@@ -65,17 +63,15 @@ class TokenizerProcessorStep(ObservationProcessorStep):
|
||||
|
||||
Requires the `transformers` library to be installed.
|
||||
|
||||
**Attributes**:
|
||||
- **tokenizer_name** (`str | None`) -- The name of a pretrained tokenizer from the Hugging Face Hub
|
||||
(e.g., "bert-base-uncased").
|
||||
- **tokenizer** (`Any | None`) -- A pre-initialized tokenizer object. If provided, `tokenizer_name` is
|
||||
ignored.
|
||||
- **max_length** (`int`) -- The maximum length to pad or truncate sequences to.
|
||||
- **task_key** (`str`) -- The key in `complementary_data` where the task string is stored.
|
||||
- **padding_side** (`str`) -- The side to pad on ('left' or 'right').
|
||||
- **padding** (`str`) -- The padding strategy ('max_length', 'longest', etc.).
|
||||
- **truncation** (`bool`) -- Whether to truncate sequences longer than `max_length`.
|
||||
- **input_tokenizer** (`Any`) -- The internal tokenizer instance, loaded during initialization.
|
||||
Attributes:
|
||||
tokenizer_name: The name of a pretrained tokenizer from the Hugging Face Hub (e.g., "bert-base-uncased").
|
||||
tokenizer: A pre-initialized tokenizer object. If provided, `tokenizer_name` is ignored.
|
||||
max_length: The maximum length to pad or truncate sequences to.
|
||||
task_key: The key in `complementary_data` where the task string is stored.
|
||||
padding_side: The side to pad on ('left' or 'right').
|
||||
padding: The padding strategy ('max_length', 'longest', etc.).
|
||||
truncation: Whether to truncate sequences longer than `max_length`.
|
||||
input_tokenizer: The internal tokenizer instance, loaded during initialization.
|
||||
"""
|
||||
|
||||
tokenizer_name: str | None = None
|
||||
@@ -140,7 +136,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
|
||||
# Standardize to a list of strings for the tokenizer
|
||||
if isinstance(task, str):
|
||||
return [task]
|
||||
elif isinstance(task, list | tuple) and all(isinstance(t, str) for t in task):
|
||||
elif isinstance(task, (list, tuple)) and all(isinstance(t, str) for t in task):
|
||||
return list(task)
|
||||
|
||||
return None
|
||||
@@ -297,15 +293,6 @@ class TokenizerProcessorStep(ObservationProcessorStep):
|
||||
|
||||
return config
|
||||
|
||||
def save_artifacts(self, save_directory: Path) -> dict[str, str]:
|
||||
"""Save the tokenizer so object-provided instances reload without overrides."""
|
||||
artifact_path = Path("tokenizer")
|
||||
save_pretrained = getattr(self.input_tokenizer, "save_pretrained", None)
|
||||
if save_pretrained is None:
|
||||
raise TypeError("Tokenizer must implement save_pretrained() to save a portable pipeline.")
|
||||
save_pretrained(save_directory / artifact_path)
|
||||
return {"tokenizer_name": artifact_path.as_posix()}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
@@ -338,27 +325,22 @@ class TokenizerProcessorStep(ObservationProcessorStep):
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="action_tokenizer_processor")
|
||||
class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
class ActionTokenizerProcessorStep(ComplementaryDataProcessorStep):
|
||||
"""
|
||||
Processor step to tokenize action data using a fast action tokenizer.
|
||||
|
||||
This step takes action tensors from an `EnvTransition`, tokenizes them using
|
||||
This step reads the action tensor from the `EnvTransition`, tokenizes it using
|
||||
a Hugging Face `transformers` AutoProcessor (such as the Physical Intelligence "fast" tokenizer),
|
||||
and returns the tokenized action.
|
||||
and stores the resulting token IDs and mask in the transition's complementary data.
|
||||
|
||||
Requires the `transformers` library to be installed.
|
||||
|
||||
**Attributes**:
|
||||
- **tokenizer_name** -- The name of a pretrained processor from the Hugging Face Hub (e.g.,
|
||||
"lerobot/fast-action-tokenizer").
|
||||
- **tokenizer** -- A pre-initialized processor/tokenizer object. If provided, `tokenizer_name` is
|
||||
ignored.
|
||||
- **trust_remote_code** (`bool`) -- Whether to trust remote code when loading the tokenizer (required
|
||||
for some tokenizers).
|
||||
- **action_tokenizer** (`Any`) -- The internal tokenizer/processor instance, loaded during
|
||||
initialization.
|
||||
- **paligemma_tokenizer_name** (`str`) -- The name of a pretrained PaliGemma tokenizer from the
|
||||
Hugging Face Hub (e.g., "google/paligemma-3b-pt-224").
|
||||
Attributes:
|
||||
action_tokenizer_name: The name of a pretrained processor from the Hugging Face Hub (e.g., "lerobot/fast-action-tokenizer").
|
||||
action_tokenizer_input_object: A pre-initialized processor/tokenizer object. If provided, `action_tokenizer_name` is ignored.
|
||||
trust_remote_code: Whether to trust remote code when loading the tokenizer (required for some tokenizers).
|
||||
action_tokenizer: The internal tokenizer/processor instance, loaded during initialization.
|
||||
paligemma_tokenizer_name: The name of a pretrained PaliGemma tokenizer from the Hugging Face Hub (e.g., "google/paligemma-3b-pt-224").
|
||||
"""
|
||||
|
||||
action_tokenizer_name: str | None = None
|
||||
@@ -367,7 +349,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
max_action_tokens: int = 256
|
||||
fast_skip_tokens: int = 128
|
||||
paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224"
|
||||
allow_truncation: bool = True
|
||||
# Internal tokenizer instance (not part of the config)
|
||||
action_tokenizer: Any = field(default=None, init=False, repr=False)
|
||||
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
|
||||
@@ -411,38 +392,27 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
add_bos_token=False,
|
||||
)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Applies action tokenization to the transition.
|
||||
|
||||
This overrides the base class to handle both tokens and mask.
|
||||
Tokenizes the transition's action and adds the tokens and mask to the complementary data.
|
||||
|
||||
Args:
|
||||
transition: The input transition with action data.
|
||||
complementary_data: The input complementary data dictionary.
|
||||
|
||||
Returns:
|
||||
The processed transition with tokenized actions and mask in complementary data.
|
||||
The complementary data with tokenized actions and mask added.
|
||||
"""
|
||||
self._current_transition = transition.copy()
|
||||
new_transition = self._current_transition
|
||||
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
action = self.transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
# During inference, no action is available, skip tokenization
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
# Tokenize and get masks for the full formatted sequence and the discrete action codes.
|
||||
tokens, mask, code_mask = self._tokenize_action(action)
|
||||
# Tokenize and get both tokens and mask
|
||||
tokens, mask = self._tokenize_action(action)
|
||||
|
||||
# Store mask in complementary data
|
||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
if complementary_data is None:
|
||||
complementary_data = {}
|
||||
complementary_data[ACTION_TOKEN_MASK] = mask
|
||||
complementary_data[ACTION_CODE_TOKEN_MASK] = code_mask
|
||||
complementary_data[ACTION_TOKENS] = tokens
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
||||
return new_transition
|
||||
return complementary_data
|
||||
|
||||
def _act_tokens_to_paligemma_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
@@ -450,7 +420,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
"""
|
||||
return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens
|
||||
|
||||
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Tokenizes the action tensor and creates a mask.
|
||||
|
||||
@@ -479,7 +449,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
# The fast tokenizer expects action data and returns token IDs
|
||||
tokens_list = []
|
||||
masks_list = []
|
||||
code_masks_list = []
|
||||
|
||||
for i in range(batch_size):
|
||||
# Tokenize single action (move to CPU first as tokenizer uses scipy which requires numpy)
|
||||
@@ -497,83 +466,58 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
if tokens.dim() > 1:
|
||||
tokens = tokens.flatten()
|
||||
|
||||
action_code_tokens = self._act_tokens_to_paligemma_tokens(tokens)
|
||||
bos_id = self._paligemma_tokenizer.bos_token_id
|
||||
prompt_tokens = torch.tensor(
|
||||
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
||||
device=action.device,
|
||||
)
|
||||
end_tokens = torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device)
|
||||
|
||||
code_start = 1 + len(prompt_tokens)
|
||||
code_end = code_start + len(action_code_tokens)
|
||||
# add bos
|
||||
tokens = torch.cat(
|
||||
[
|
||||
torch.tensor([bos_id], device=action.device),
|
||||
prompt_tokens,
|
||||
action_code_tokens,
|
||||
end_tokens,
|
||||
torch.tensor(
|
||||
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
||||
device=action.device,
|
||||
),
|
||||
self._act_tokens_to_paligemma_tokens(tokens),
|
||||
torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device),
|
||||
]
|
||||
)
|
||||
code_mask = torch.zeros(len(tokens), dtype=torch.bool, device=action.device)
|
||||
code_mask[code_start:code_end] = True
|
||||
|
||||
# Truncate or pad to max_action_tokens
|
||||
if len(tokens) > self.max_action_tokens:
|
||||
if not self.allow_truncation:
|
||||
raise ValueError(
|
||||
f"FAST action sequence has {len(tokens)} tokens, exceeding "
|
||||
f"max_action_tokens={self.max_action_tokens}."
|
||||
)
|
||||
logging.warning(
|
||||
f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. "
|
||||
"Consider increasing the `max_action_tokens` in your model config if this happens frequently."
|
||||
)
|
||||
tokens = tokens[: self.max_action_tokens]
|
||||
code_mask = code_mask[: self.max_action_tokens]
|
||||
mask = torch.ones(self.max_action_tokens, dtype=torch.bool, device=action.device)
|
||||
else:
|
||||
pad_len = self.max_action_tokens - len(tokens)
|
||||
mask = torch.cat(
|
||||
[
|
||||
torch.ones(len(tokens), dtype=torch.bool, device=action.device),
|
||||
torch.zeros(pad_len, dtype=torch.bool, device=action.device),
|
||||
torch.zeros(
|
||||
self.max_action_tokens - len(tokens), dtype=torch.bool, device=action.device
|
||||
),
|
||||
]
|
||||
)
|
||||
code_mask = torch.nn.functional.pad(code_mask, (0, pad_len), value=False)
|
||||
# Pad tokens with zeros
|
||||
tokens = torch.nn.functional.pad(tokens, (0, pad_len), value=0)
|
||||
tokens = torch.nn.functional.pad(tokens, (0, self.max_action_tokens - len(tokens)), value=0)
|
||||
|
||||
tokens_list.append(tokens)
|
||||
masks_list.append(mask)
|
||||
code_masks_list.append(code_mask)
|
||||
|
||||
# Stack into batched tensors
|
||||
tokens_batch = torch.stack(tokens_list, dim=0) # (B, max_action_tokens)
|
||||
masks_batch = torch.stack(masks_list, dim=0) # (B, max_action_tokens)
|
||||
code_masks_batch = torch.stack(code_masks_list, dim=0) # (B, max_action_tokens)
|
||||
|
||||
# Remove batch dimension if input was single sample
|
||||
if single_sample:
|
||||
tokens_batch = tokens_batch.squeeze(0)
|
||||
masks_batch = masks_batch.squeeze(0)
|
||||
code_masks_batch = code_masks_batch.squeeze(0)
|
||||
|
||||
# Move to the same device as the input
|
||||
if device is not None:
|
||||
tokens_batch = tokens_batch.to(device)
|
||||
masks_batch = masks_batch.to(device)
|
||||
code_masks_batch = code_masks_batch.to(device)
|
||||
|
||||
return tokens_batch, masks_batch, code_masks_batch
|
||||
|
||||
def action(self, action: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
This method is not used since we override __call__.
|
||||
Required by ActionProcessorStep ABC.
|
||||
"""
|
||||
tokens, _, _ = self._tokenize_action(action)
|
||||
return tokens
|
||||
return tokens_batch, masks_batch
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -590,7 +534,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
"max_action_tokens": self.max_action_tokens,
|
||||
"fast_skip_tokens": self.fast_skip_tokens,
|
||||
"paligemma_tokenizer_name": self.paligemma_tokenizer_name,
|
||||
"allow_truncation": self.allow_truncation,
|
||||
}
|
||||
|
||||
# Only save tokenizer_name if it was used to create the tokenizer
|
||||
@@ -599,27 +542,19 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
|
||||
return config
|
||||
|
||||
def save_artifacts(self, save_directory: Path) -> dict[str, str]:
|
||||
artifact_path = Path("action_tokenizer")
|
||||
save_pretrained = getattr(self.action_tokenizer, "save_pretrained", None)
|
||||
if save_pretrained is None:
|
||||
raise TypeError("Action tokenizer must implement save_pretrained() to save a portable pipeline.")
|
||||
save_pretrained(save_directory / artifact_path)
|
||||
return {"action_tokenizer_name": artifact_path.as_posix()}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""
|
||||
Updates feature definitions to reflect tokenized actions.
|
||||
Returns the policy features unchanged.
|
||||
|
||||
This updates the policy features dictionary to indicate that the action
|
||||
has been tokenized into a sequence of token IDs with shape (max_action_tokens,).
|
||||
The tokenized actions and mask are stored in complementary data, which is not
|
||||
tracked in the policy features dictionary.
|
||||
|
||||
Args:
|
||||
features: The dictionary of existing policy features.
|
||||
|
||||
Returns:
|
||||
The updated dictionary of policy features.
|
||||
The dictionary of policy features, unchanged.
|
||||
"""
|
||||
return features
|
||||
|
||||
@@ -16,11 +16,12 @@ import abc
|
||||
import builtins
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
from importlib.resources import files
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, Any, TypeVar
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
from safetensors.torch import load_model as load_model_as_safetensor, save_model as save_model_as_safetensor
|
||||
@@ -60,22 +61,6 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
|
||||
raise TypeError(f"Class {cls.__name__} must define 'name'")
|
||||
|
||||
def _save_pretrained(self, save_directory: Path) -> None:
|
||||
"""Serialize this reward model's parameters (and config) into `save_directory`.
|
||||
|
||||
Safe to call on every rank: replicas carry identical weights, so only the main process
|
||||
writes (sharded reward models are rejected at config validation — no collective gather).
|
||||
|
||||
Args:
|
||||
save_directory (Path): Target directory for the reward model config (`config.json`)
|
||||
and `model.safetensors`.
|
||||
"""
|
||||
from lerobot.distributed.utils import is_main_process
|
||||
|
||||
# save_checkpoint calls this on every rank; replicas carry identical
|
||||
# weights, so the main process is the only writer. Sharded reward models are rejected
|
||||
# at config validation, so no collective gather is needed here.
|
||||
if not is_main_process():
|
||||
return
|
||||
self.config._save_pretrained(save_directory)
|
||||
model_to_save = self.module if hasattr(self, "module") else self
|
||||
save_model_as_safetensor(model_to_save, str(save_directory / SAFETENSORS_SINGLE_FILE))
|
||||
@@ -190,22 +175,53 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
|
||||
"""
|
||||
return type(self).forward is not PreTrainedRewardModel.forward
|
||||
|
||||
def push_model_to_hub(self, cfg: "TrainPipelineConfig") -> None:
|
||||
"""Publish this reward model to the Hub.
|
||||
def push_model_to_hub(self, cfg: "TrainPipelineConfig"):
|
||||
api = HfApi()
|
||||
repo_id = api.create_repo(
|
||||
repo_id=self.config.repo_id, private=self.config.private, exist_ok=True
|
||||
).repo_id
|
||||
|
||||
Deprecated: use :func:`lerobot.common.train_utils.publish_trained_model` instead.
|
||||
# Push the files to the repo in a single commit
|
||||
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
|
||||
saved_path = Path(tmp) / repo_id
|
||||
|
||||
Args:
|
||||
cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and
|
||||
used to render the model card.
|
||||
"""
|
||||
from lerobot.common.train_utils import publish_trained_model
|
||||
self.save_pretrained(saved_path) # Calls _save_pretrained and stores model tensors
|
||||
|
||||
warnings.warn(
|
||||
"PreTrainedRewardModel.push_model_to_hub is deprecated and will be removed in a "
|
||||
"future version. Use lerobot.common.train_utils.publish_trained_model(cfg, model, "
|
||||
"preprocessor, postprocessor, dataset_meta) instead.",
|
||||
FutureWarning,
|
||||
stacklevel=2,
|
||||
card = self.generate_model_card(
|
||||
cfg.dataset.repo_id, self.config.type, self.config.license, self.config.tags
|
||||
)
|
||||
card.save(str(saved_path / "README.md"))
|
||||
|
||||
cfg.save_pretrained(saved_path) # Calls _save_pretrained and stores train config
|
||||
|
||||
commit_info = api.upload_folder(
|
||||
repo_id=repo_id,
|
||||
repo_type="model",
|
||||
folder_path=saved_path,
|
||||
commit_message="Upload reward model weights, train config and readme",
|
||||
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
|
||||
ignore_patterns=["*.tmp", "*.log"],
|
||||
)
|
||||
|
||||
logging.info(f"Model pushed to {commit_info.repo_url.url}")
|
||||
|
||||
def generate_model_card(
|
||||
self, dataset_repo_id: str, model_type: str, license: str | None, tags: list[str] | None
|
||||
) -> ModelCard:
|
||||
card_data = ModelCardData(
|
||||
license=license or "apache-2.0",
|
||||
library_name="lerobot",
|
||||
pipeline_tag="robotics",
|
||||
tags=list(set(tags or []).union({"robotics", "lerobot", "reward-model", model_type})),
|
||||
model_name=model_type,
|
||||
datasets=dataset_repo_id,
|
||||
)
|
||||
publish_trained_model(cfg, self, None, None, None)
|
||||
|
||||
template_card = (
|
||||
files("lerobot.templates")
|
||||
.joinpath("lerobot_rewardmodel_modelcard_template.md")
|
||||
.read_text(encoding="utf-8")
|
||||
)
|
||||
card = ModelCard.from_template(card_data, template_str=template_card)
|
||||
card.validate()
|
||||
return card
|
||||
|
||||
@@ -25,13 +25,13 @@ from PIL import Image
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import TransitionKey
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
policy_action_to_transition,
|
||||
)
|
||||
@@ -105,7 +105,7 @@ def _expand_tasks(task: Any, *, batch_size: int, default: str | None) -> list[st
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="robometer_encoder")
|
||||
class RobometerEncoderProcessorStep(ProcessorStep):
|
||||
class RobometerEncoderProcessorStep(ObservationProcessorStep):
|
||||
"""Encode raw frames + task into Qwen-VL tensors for the Robometer model.
|
||||
|
||||
Loads a :class:`~transformers.AutoProcessor` matching ``base_model_id`` and
|
||||
@@ -160,11 +160,8 @@ class RobometerEncoderProcessorStep(ProcessorStep):
|
||||
if token not in tokenizer.get_vocab():
|
||||
tokenizer.add_special_tokens({"additional_special_tokens": [token]})
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
observation = transition.get(TransitionKey.OBSERVATION)
|
||||
complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
if not isinstance(observation, dict):
|
||||
raise ValueError("RobometerEncoderProcessorStep requires an observation dict")
|
||||
def observation(self, observation: dict[str, Any]) -> dict[str, Any]:
|
||||
complementary = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
|
||||
if self.image_key not in observation:
|
||||
raise KeyError(f"Robometer expected image key {self.image_key!r} in observation")
|
||||
@@ -190,13 +187,9 @@ class RobometerEncoderProcessorStep(ProcessorStep):
|
||||
]
|
||||
encoded = self.encode_samples(samples)
|
||||
|
||||
new_observation = dict(observation)
|
||||
for key, value in encoded.items():
|
||||
new_observation[f"{ROBOMETER_FEATURE_PREFIX}{key}"] = value
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_transition[TransitionKey.OBSERVATION] = new_observation
|
||||
return new_transition
|
||||
observation[f"{ROBOMETER_FEATURE_PREFIX}{key}"] = value
|
||||
return observation
|
||||
|
||||
def encode_samples(self, samples: list[tuple[np.ndarray, str]]) -> dict[str, Tensor]:
|
||||
"""Run the Qwen-VL processor on a list of ``(frames, task)`` samples."""
|
||||
|
||||
@@ -48,13 +48,13 @@ else:
|
||||
Faker = None # type: ignore[assignment, misc]
|
||||
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, PolicyAction, TransitionKey
|
||||
from lerobot.lerobot_types import PolicyAction, TransitionKey
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
RenameObservationsProcessorStep,
|
||||
from_tensor_to_numpy,
|
||||
policy_action_to_transition,
|
||||
@@ -73,7 +73,7 @@ from .sarm_utils import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SARMEncodingProcessorStep(ProcessorStep):
|
||||
class SARMEncodingProcessorStep(ObservationProcessorStep):
|
||||
"""ProcessorStep that encodes images and text with CLIP and generates stage and progress labels for SARM."""
|
||||
|
||||
def __init__(
|
||||
@@ -257,9 +257,9 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
||||
|
||||
return annotations
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
def observation(self, observation: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Encode images, text, and normalize states in the transition.
|
||||
Encode images, text, and normalize states in the observation.
|
||||
|
||||
Implements SARM training data preparation:
|
||||
- Applies language perturbation (20% probability)
|
||||
@@ -267,9 +267,7 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
||||
- Generates stage+tau targets for all frames
|
||||
- Outputs lengths tensor for valid sequence masking
|
||||
"""
|
||||
new_transition = transition.copy() if hasattr(transition, "copy") else dict(transition)
|
||||
observation = new_transition.get(TransitionKey.OBSERVATION)
|
||||
comp_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
comp_data = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
|
||||
frame_index = comp_data.get("index")
|
||||
episode_index = comp_data.get("episode_index")
|
||||
@@ -392,8 +390,7 @@ class SARMEncodingProcessorStep(ProcessorStep):
|
||||
)
|
||||
observation["dense_targets"] = dense_targets
|
||||
|
||||
new_transition[TransitionKey.OBSERVATION] = observation
|
||||
return new_transition
|
||||
return observation
|
||||
|
||||
def _compute_batch_targets(
|
||||
self,
|
||||
|
||||
@@ -58,11 +58,12 @@ import builtins
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, Any, TypeVar
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
from huggingface_hub import HfApi, hf_hub_download
|
||||
from huggingface_hub.constants import CONFIG_NAME
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
from torch import Tensor
|
||||
@@ -74,6 +75,9 @@ from lerobot.rewards.topreward.configuration_topreward import TOPRewardConfig
|
||||
from lerobot.rewards.topreward.processor_topreward import TOPREWARD_FEATURE_PREFIX, TOPREWARD_INPUT_KEYS
|
||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers import Qwen3VLForConditionalGeneration
|
||||
else:
|
||||
@@ -201,3 +205,34 @@ class TOPRewardModel(PreTrainedRewardModel):
|
||||
instance.to(config.device)
|
||||
instance.eval()
|
||||
return instance
|
||||
|
||||
def push_model_to_hub(self, cfg: TrainPipelineConfig):
|
||||
"""Push the TOPReward ``config.json`` + model card to the Hub."""
|
||||
api = HfApi()
|
||||
repo_id = api.create_repo(
|
||||
repo_id=self.config.repo_id, private=self.config.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)
|
||||
|
||||
self.config._save_pretrained(saved_path)
|
||||
|
||||
card = self.generate_model_card(
|
||||
cfg.dataset.repo_id, self.config.type, self.config.license, self.config.tags
|
||||
)
|
||||
card.save(str(saved_path / "README.md"))
|
||||
|
||||
cfg.save_pretrained(saved_path)
|
||||
|
||||
commit_info = api.upload_folder(
|
||||
repo_id=repo_id,
|
||||
repo_type="model",
|
||||
folder_path=saved_path,
|
||||
commit_message="Upload TOPReward config and readme",
|
||||
allow_patterns=["*.json", "*.yaml", "*.md"],
|
||||
ignore_patterns=["*.tmp", "*.log", "*.safetensors"],
|
||||
)
|
||||
|
||||
logger.info(f"Model pushed to {commit_info.repo_url.url}")
|
||||
|
||||
@@ -23,13 +23,13 @@ import torch
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.lerobot_types import EnvTransition, TransitionKey
|
||||
from lerobot.lerobot_types import TransitionKey
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
policy_action_to_transition,
|
||||
)
|
||||
@@ -107,7 +107,7 @@ def _expand_tasks(task: Any, *, batch_size: int, default: str | None) -> list[st
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="topreward_encoder")
|
||||
class TOPRewardEncoderProcessorStep(ProcessorStep):
|
||||
class TOPRewardEncoderProcessorStep(ObservationProcessorStep):
|
||||
"""Encode raw frames + task into Qwen-VL tensors for the TOPReward model.
|
||||
|
||||
Loads a :class:`~transformers.AutoProcessor` matching ``vlm_name`` and
|
||||
@@ -142,9 +142,8 @@ class TOPRewardEncoderProcessorStep(ProcessorStep):
|
||||
require_package("transformers", extra="topreward")
|
||||
self._processor = AutoProcessor.from_pretrained(self.vlm_name, trust_remote_code=True)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
observation = transition.get(TransitionKey.OBSERVATION)
|
||||
complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
def observation(self, observation: dict[str, Any]) -> dict[str, Any]:
|
||||
complementary = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
if self.image_key not in observation:
|
||||
raise KeyError(f"TOPReward expected image key {self.image_key!r} in observation")
|
||||
|
||||
@@ -161,13 +160,9 @@ class TOPRewardEncoderProcessorStep(ProcessorStep):
|
||||
|
||||
encoded = self._encode_batch(videos, tasks, batch_size)
|
||||
|
||||
new_observation = dict(observation)
|
||||
for key, value in encoded.items():
|
||||
new_observation[f"{TOPREWARD_FEATURE_PREFIX}{key}"] = value
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_transition[TransitionKey.OBSERVATION] = new_observation
|
||||
return new_transition
|
||||
observation[f"{TOPREWARD_FEATURE_PREFIX}{key}"] = value
|
||||
return observation
|
||||
|
||||
def _encode_batch(self, videos: Tensor, tasks: list[str], batch_size) -> dict[str, Any]:
|
||||
"""Tokenise a batch of (frames, task) pairs into Qwen-VL tensors.
|
||||
|
||||
+23
-59
@@ -13,7 +13,8 @@
|
||||
# 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.
|
||||
"""Actor server runner for distributed HILSerl robot policy training.
|
||||
"""
|
||||
Actor server runner for distributed HILSerl robot policy training.
|
||||
|
||||
This script implements the actor component of the distributed HILSerl architecture.
|
||||
It executes the policy in the robot environment, collects experience,
|
||||
@@ -118,16 +119,6 @@ from .train_rl import TrainRLServerPipelineConfig
|
||||
|
||||
@parser.wrap()
|
||||
def actor_cli(cfg: TrainRLServerPipelineConfig):
|
||||
"""CLI entry point for the HILSerl actor server.
|
||||
|
||||
Connects to the learner server over gRPC, then launches (as threads or processes, depending on
|
||||
`cfg.policy.concurrency.multiprocessing_context`) the background workers that receive updated
|
||||
policy parameters and stream transitions/interactions back to the learner, while running the
|
||||
policy-environment interaction loop (`act_with_policy`) on the main thread/process.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Parsed from the CLI.
|
||||
"""
|
||||
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
||||
require_package("grpcio", extra="hilserl", import_name="grpc")
|
||||
cfg.validate()
|
||||
@@ -243,19 +234,18 @@ def act_with_policy(
|
||||
transitions_queue: Queue,
|
||||
interactions_queue: Queue,
|
||||
):
|
||||
"""Executes policy interaction within the environment.
|
||||
"""
|
||||
Executes policy interaction within the environment.
|
||||
|
||||
This function rolls out the policy in the environment, collecting interaction data and pushing it to a queue for streaming to the learner.
|
||||
Once an episode is completed, updated network parameters received from the learner are retrieved from a queue and loaded into the network.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Training configuration.
|
||||
shutdown_event (`Event`): Set to stop the policy loop.
|
||||
parameters_queue (`Queue`): Queue of serialized learner weights, drained via
|
||||
`update_policy_parameters`.
|
||||
transitions_queue (`Queue`): Queue transitions are pushed to for streaming to the learner.
|
||||
interactions_queue (`Queue`): Queue interaction messages are pushed to for streaming to the
|
||||
learner.
|
||||
cfg: Configuration settings for the interaction process.
|
||||
shutdown_event: Event to check if the process should shutdown.
|
||||
parameters_queue: Queue to receive updated network parameters from the learner.
|
||||
transitions_queue: Queue to send transitions to the learner.
|
||||
interactions_queue: Queue to send interactions to the learner.
|
||||
"""
|
||||
# Initialize logging for multiprocessing
|
||||
if not use_threads(cfg):
|
||||
@@ -450,8 +440,7 @@ def establish_learner_connection(
|
||||
Args:
|
||||
stub (services_pb2_grpc.LearnerServiceStub): The stub to use for the connection.
|
||||
shutdown_event (Event): The event to check if the connection should be established.
|
||||
attempts (int, *optional*, defaults to 30): The number of attempts to establish the connection.
|
||||
|
||||
attempts (int): The number of attempts to establish the connection.
|
||||
Returns:
|
||||
bool: True if the connection is established, False otherwise.
|
||||
"""
|
||||
@@ -484,6 +473,7 @@ def learner_service_client(
|
||||
Returns:
|
||||
tuple[services_pb2_grpc.LearnerServiceStub, grpc.Channel]: The stub and the channel.
|
||||
"""
|
||||
|
||||
channel = grpc.insecure_channel(
|
||||
f"{host}:{port}",
|
||||
grpc_channel_options(),
|
||||
@@ -506,8 +496,8 @@ def receive_policy(
|
||||
cfg (TrainRLServerPipelineConfig): The configuration for the actor.
|
||||
parameters_queue (Queue): The queue to receive the parameters.
|
||||
shutdown_event (Event): The event to check if the process should shutdown.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None, *optional*): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None, *optional*): Optional pre-created channel.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None): Optional pre-created channel.
|
||||
"""
|
||||
logging.info("[ACTOR] Start receiving parameters from the Learner")
|
||||
if not use_threads(cfg):
|
||||
@@ -567,9 +557,10 @@ def send_transitions(
|
||||
cfg (TrainRLServerPipelineConfig): The configuration for the actor.
|
||||
transitions_queue (Queue): The queue to receive the transitions.
|
||||
shutdown_event (Event): The event to check if the process should shutdown.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None, *optional*): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None, *optional*): Optional pre-created channel.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None): Optional pre-created channel.
|
||||
"""
|
||||
|
||||
if not use_threads(cfg):
|
||||
# Create a process-specific log file
|
||||
log_dir = os.path.join(cfg.output_dir, "logs")
|
||||
@@ -621,9 +612,10 @@ def send_interactions(
|
||||
cfg (TrainRLServerPipelineConfig): The configuration for the actor.
|
||||
interactions_queue (Queue): The queue to receive the interactions.
|
||||
shutdown_event (Event): The event to check if the process should shutdown.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None, *optional*): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None, *optional*): Optional pre-created channel.
|
||||
learner_client (services_pb2_grpc.LearnerServiceStub | None): Optional pre-created stub.
|
||||
grpc_channel (grpc.Channel | None): Optional pre-created channel.
|
||||
"""
|
||||
|
||||
if not use_threads(cfg):
|
||||
# Create a process-specific log file
|
||||
log_dir = os.path.join(cfg.output_dir, "logs")
|
||||
@@ -665,17 +657,6 @@ def transitions_stream(
|
||||
transitions_queue: Queue,
|
||||
timeout: float,
|
||||
) -> "Generator[Any, None, services_pb2.Empty]":
|
||||
"""GRPC client-streaming generator that forwards queued transitions to the learner.
|
||||
|
||||
Args:
|
||||
shutdown_event (`Event`): Set to stop streaming and return.
|
||||
transitions_queue (`Queue`): Queue of serialized transition batches, filled by
|
||||
`push_transitions_to_transport_queue`.
|
||||
timeout (`float`): Seconds to wait for a queue item before checking `shutdown_event` again.
|
||||
|
||||
Yields:
|
||||
Chunks of a `services_pb2.Transition` message, produced by `send_bytes_in_chunks`.
|
||||
"""
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
message = transitions_queue.get(block=True, timeout=timeout)
|
||||
@@ -695,16 +676,6 @@ def interactions_stream(
|
||||
interactions_queue: Queue,
|
||||
timeout: float,
|
||||
) -> "Generator[Any, None, services_pb2.Empty]":
|
||||
"""GRPC client-streaming generator that forwards queued interaction messages to the learner.
|
||||
|
||||
Args:
|
||||
shutdown_event (`Event`): Set to stop streaming and return.
|
||||
interactions_queue (`Queue`): Queue of serialized interaction messages.
|
||||
timeout (`float`): Seconds to wait for a queue item before checking `shutdown_event` again.
|
||||
|
||||
Yields:
|
||||
Chunks of a `services_pb2.InteractionMessage`, produced by `send_bytes_in_chunks`.
|
||||
"""
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
message = interactions_queue.get(block=True, timeout=timeout)
|
||||
@@ -747,11 +718,12 @@ def update_policy_parameters(algorithm: RLAlgorithm, parameters_queue: Queue, de
|
||||
|
||||
|
||||
def push_transitions_to_transport_queue(transitions: list, transitions_queue):
|
||||
"""Move `transitions` to CPU, check for NaNs, and enqueue them for the learner.
|
||||
"""Send transitions to learner in smaller chunks to avoid network issues.
|
||||
|
||||
Args:
|
||||
transitions (`list`): Transitions to send, as produced by the actor's rollout loop.
|
||||
transitions_queue (`Queue`): Queue drained by `transitions_stream`.
|
||||
transitions: List of transitions to send
|
||||
message_queue: Queue to send messages to learner
|
||||
chunk_size: Size of each chunk to send
|
||||
"""
|
||||
transition_to_send_to_learner = []
|
||||
for transition in transitions:
|
||||
@@ -788,13 +760,6 @@ def get_frequency_stats(timer: TimerManager) -> dict[str, float]:
|
||||
|
||||
|
||||
def log_policy_frequency_issue(policy_fps: float, cfg: TrainRLServerPipelineConfig, interaction_step: int):
|
||||
"""Log a warning if `policy_fps` is below the environment's target `cfg.env.fps`.
|
||||
|
||||
Args:
|
||||
policy_fps (`float`): Measured policy loop frequency.
|
||||
cfg (`TrainRLServerPipelineConfig`): Provides the target `cfg.env.fps` to compare against.
|
||||
interaction_step (`int`): Current interaction step, included in the warning message.
|
||||
"""
|
||||
if policy_fps < cfg.env.fps:
|
||||
logging.warning(
|
||||
f"[ACTOR] Policy FPS {policy_fps:.1f} below required {cfg.env.fps} at step {interaction_step}"
|
||||
@@ -802,7 +767,6 @@ def log_policy_frequency_issue(policy_fps: float, cfg: TrainRLServerPipelineConf
|
||||
|
||||
|
||||
def use_threads(cfg: TrainRLServerPipelineConfig) -> bool:
|
||||
"""Whether the actor's background workers should run as threads instead of processes."""
|
||||
return cfg.policy.concurrency.actor == "threads"
|
||||
|
||||
|
||||
|
||||
@@ -101,7 +101,6 @@ class RLAlgorithm(HubMixin, abc.ABC):
|
||||
|
||||
@optimization_step.setter
|
||||
def optimization_step(self, value: int) -> None:
|
||||
"""Set the current learner optimization step."""
|
||||
self._optimization_step = int(value)
|
||||
|
||||
def get_weights(self) -> dict[str, Any]:
|
||||
|
||||
@@ -45,6 +45,7 @@ class TrainingStats:
|
||||
|
||||
def to_log_dict(self) -> dict[str, float]:
|
||||
"""Flatten all stats into a single dict for logging."""
|
||||
|
||||
d: dict[str, float] = {}
|
||||
for name, val in self.losses.items():
|
||||
d[name] = val
|
||||
@@ -97,35 +98,6 @@ class RLAlgorithmConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC):
|
||||
revision: str | None = None,
|
||||
**algo_kwargs: Any,
|
||||
) -> T:
|
||||
"""Load an algorithm config from a local directory or the Hugging Face Hub.
|
||||
|
||||
Args:
|
||||
pretrained_name_or_path (`str | Path`):
|
||||
Local directory containing `config.json`, or a Hub repo id.
|
||||
force_download (`bool`, *optional*, defaults to `False`):
|
||||
Whether to force re-download the config even if it's cached.
|
||||
resume_download (`bool | None`, *optional*):
|
||||
Whether to resume an interrupted download.
|
||||
proxies (`dict[Any, Any] | None`, *optional*):
|
||||
Proxies to use for the download request.
|
||||
token (`str | bool | None`, *optional*):
|
||||
Hugging Face Hub authentication token.
|
||||
cache_dir (`str | Path | None`, *optional*):
|
||||
Directory to cache the downloaded config in.
|
||||
local_files_only (`bool`, *optional*, defaults to `False`):
|
||||
Whether to only look for files locally, without querying the Hub.
|
||||
revision (`str | None`, *optional*):
|
||||
Hub revision (branch, tag, or commit hash) to load from.
|
||||
**algo_kwargs: Attribute overrides applied to the loaded config instance.
|
||||
|
||||
Returns:
|
||||
RLAlgorithmConfig: The loaded config, as the concrete registered subclass.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If no `config.json` is found locally or on the Hub.
|
||||
TypeError: If loaded via a specific subclass but the config's registered type doesn't
|
||||
match it.
|
||||
"""
|
||||
model_id = str(pretrained_name_or_path)
|
||||
config_file: str | None = None
|
||||
if Path(model_id).is_dir():
|
||||
|
||||
@@ -24,8 +24,8 @@ def make_algorithm_config(algorithm_type: str, **kwargs) -> RLAlgorithmConfig:
|
||||
"""Instantiate an `RLAlgorithmConfig` from its registered type name.
|
||||
|
||||
Args:
|
||||
algorithm_type (`str`): Registry key of the algorithm (e.g. `"sac"`).
|
||||
kwargs (`Any`, *optional*): Keyword arguments forwarded to the config class constructor.
|
||||
algorithm_type: Registry key of the algorithm (e.g. ``"sac"``).
|
||||
**kwargs: Keyword arguments forwarded to the config class constructor.
|
||||
|
||||
Returns:
|
||||
An instance of the matching ``RLAlgorithmConfig`` subclass.
|
||||
@@ -44,13 +44,14 @@ def make_algorithm_config(algorithm_type: str, **kwargs) -> RLAlgorithmConfig:
|
||||
|
||||
|
||||
def get_algorithm_class(name: str) -> type[RLAlgorithm]:
|
||||
"""Retrieves an RL algorithm class by its registered name.
|
||||
"""
|
||||
Retrieves an RL algorithm class by its registered name.
|
||||
|
||||
This function uses dynamic imports to avoid loading all algorithm classes into
|
||||
memory at once, improving startup time and reducing dependencies.
|
||||
|
||||
Args:
|
||||
name (`str`): The name of the algorithm. Supported names are "sac".
|
||||
name: The name of the algorithm. Supported names are "sac".
|
||||
|
||||
Returns:
|
||||
The algorithm class corresponding to the given name.
|
||||
@@ -69,7 +70,8 @@ def get_algorithm_class(name: str) -> type[RLAlgorithm]:
|
||||
|
||||
|
||||
def make_algorithm(cfg: RLAlgorithmConfig, policy: torch.nn.Module) -> RLAlgorithm:
|
||||
"""Instantiate an RL algorithm.
|
||||
"""
|
||||
Instantiate an RL algorithm.
|
||||
|
||||
This factory function looks up the :class:`RLAlgorithm` subclass that matches
|
||||
``cfg.type`` and instantiates it with the provided policy. It also enforces
|
||||
@@ -77,8 +79,8 @@ def make_algorithm(cfg: RLAlgorithmConfig, policy: torch.nn.Module) -> RLAlgorit
|
||||
normally handled by :meth:`TrainRLServerPipelineConfig.validate`).
|
||||
|
||||
Args:
|
||||
cfg (`RLAlgorithmConfig`): The algorithm configuration. Must have `policy_config` set.
|
||||
policy (`torch.nn.Module`): The policy module the algorithm will train.
|
||||
cfg: The algorithm configuration. Must have ``policy_config`` set.
|
||||
policy: The policy module the algorithm will train.
|
||||
|
||||
Returns:
|
||||
An instantiated :class:`RLAlgorithm`.
|
||||
|
||||
@@ -39,73 +39,52 @@ class SACAlgorithmConfig(RLAlgorithmConfig):
|
||||
update loop. The policy-side (actor + observation encoder) lives in
|
||||
:class:`~lerobot.policies.gaussian_actor.GaussianActorConfig` and is
|
||||
referenced via :attr:`policy_config`.
|
||||
|
||||
Args:
|
||||
actor_lr (`float`, *optional*, defaults to 0.0003):
|
||||
Learning rate for the actor network.
|
||||
critic_lr (`float`, *optional*, defaults to 0.0003):
|
||||
Learning rate for the critic network.
|
||||
temperature_lr (`float`, *optional*, defaults to 0.0003):
|
||||
Learning rate for the temperature parameter.
|
||||
discount (`float`, *optional*, defaults to 0.99):
|
||||
Discount factor for the Bellman update.
|
||||
use_backup_entropy (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use backup entropy in the Bellman target.
|
||||
critic_target_update_weight (`float`, *optional*, defaults to 0.005):
|
||||
Polyak-averaging weight for the critic target update.
|
||||
num_critics (`int`, *optional*, defaults to 2):
|
||||
Number of critics in the ensemble.
|
||||
num_subsample_critics (`int | None`, *optional*):
|
||||
Number of critics to subsample from the ensemble for each Bellman target computation.
|
||||
`None` uses the full ensemble.
|
||||
critic_network_kwargs (`CriticNetworkConfig`, *optional*):
|
||||
Configuration for the (continuous-action) critic network architecture.
|
||||
discrete_critic_network_kwargs (`CriticNetworkConfig`, *optional*):
|
||||
Configuration for the discrete-action critic network architecture.
|
||||
temperature_init (`float`, *optional*, defaults to 1.0):
|
||||
Initial value of the entropy temperature.
|
||||
target_entropy (`float | None`, *optional*):
|
||||
Target entropy for automatic temperature tuning. If `None`, defaults to `-|A|/2` where
|
||||
`|A|` is the total action dimension (continuous + 1 if there is a discrete action head).
|
||||
utd_ratio (`int`, *optional*, defaults to 1):
|
||||
Update-to-data ratio. Set to `>1` to enable extra critic updates per env step.
|
||||
policy_update_freq (`int`, *optional*, defaults to 1):
|
||||
Frequency of policy updates, in units of critic updates.
|
||||
grad_clip_norm (`float`, *optional*, defaults to 40.0):
|
||||
Gradient-clipping norm applied during optimization.
|
||||
use_torch_compile (`bool`, *optional*, defaults to `False`):
|
||||
Whether to `torch.compile` the algorithm's forward passes. Currently disabled by default.
|
||||
policy_config (`PreTrainedConfig | None`, *optional*):
|
||||
The policy (actor) config this algorithm trains. Populated via `from_policy_config` or by
|
||||
`TrainRLServerPipelineConfig.validate` before the algorithm is constructed.
|
||||
"""
|
||||
|
||||
# Optimizer learning rates
|
||||
# Learning rate for the actor network
|
||||
actor_lr: float = 3e-4
|
||||
# Learning rate for the critic network
|
||||
critic_lr: float = 3e-4
|
||||
# Learning rate for the temperature parameter
|
||||
temperature_lr: float = 3e-4
|
||||
|
||||
# Bellman update
|
||||
# Discount factor for the SAC algorithm
|
||||
discount: float = 0.99
|
||||
# Whether to use backup entropy for the SAC algorithm
|
||||
use_backup_entropy: bool = True
|
||||
# Weight for the critic target update
|
||||
critic_target_update_weight: float = 0.005
|
||||
|
||||
# Critic ensemble
|
||||
# Number of critics in the ensemble
|
||||
num_critics: int = 2
|
||||
# Number of subsampled critics for training
|
||||
num_subsample_critics: int | None = None
|
||||
# Configuration for the critic network architecture
|
||||
critic_network_kwargs: CriticNetworkConfig = field(default_factory=CriticNetworkConfig)
|
||||
# Configuration for the discrete critic network
|
||||
discrete_critic_network_kwargs: CriticNetworkConfig = field(default_factory=CriticNetworkConfig)
|
||||
|
||||
# Temperature / entropy
|
||||
# Initial temperature value
|
||||
temperature_init: float = 1.0
|
||||
# Target entropy for automatic temperature tuning. If ``None``, defaults to
|
||||
# ``-|A|/2`` where ``|A|`` is the total action dimension (continuous + 1 if
|
||||
# there is a discrete action head).
|
||||
target_entropy: float | None = None
|
||||
|
||||
# Update loop
|
||||
# Update-to-data ratio. Set to >1 to enable extra critic updates per env step.
|
||||
utd_ratio: int = 1
|
||||
# Frequency of policy updates
|
||||
policy_update_freq: int = 1
|
||||
# Gradient clipping norm for the SAC algorithm
|
||||
grad_clip_norm: float = 40.0
|
||||
|
||||
# Optimizations
|
||||
# torch.compile is currently disabled by default
|
||||
use_torch_compile: bool = False
|
||||
|
||||
# Policy config
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user