mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| fde5db8406 | |||
| 39c4e746f1 | |||
| d3ee0b820c | |||
| 072c697c0e | |||
| 266be2bd17 | |||
| ff7cc3de1d | |||
| 31fedfd9dd | |||
| b1bf24f565 |
@@ -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"
|
||||
|
||||
@@ -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
|
||||
@@ -62,7 +62,10 @@ Reference data points on a 4×H100 80 GB cluster (`accelerate launch --num_proce
|
||||
| `smolvla` | 27m 49s | 0.312 | 0.011 | ~80% | `--policy.path=lerobot/smolvla_base`, `freeze_vision_encoder=false`, `train_expert_only=false` |
|
||||
| `pi05` | 3h 41m | 2.548 | 0.014 | ~95% | `--policy.pretrained_path=lerobot/pi05_base`, `gradient_checkpointing=true`, `dtype=bfloat16`, vision encoder + expert trained |
|
||||
|
||||
The `dataloading_s` vs. `update_s` ratio is the diagnostic that matters: when `dataloading_s` approaches `update_s`, more GPUs stop helping — your dataloader is the bottleneck and you should look at `--num_workers`, image resolution, and disk speed before adding compute.
|
||||
Training logs separate the full iteration into `dataloading_s` (`next(dl_iter)`), `preprocessing_s`
|
||||
(image conversion and the policy pipeline), and `update_s` (the optimizer update). `step_s` covers all
|
||||
three and drives `samples_per_s`. The benchmark above predates this split, so its `dataloading_s` includes
|
||||
preprocessing.
|
||||
|
||||
### Schedule and checkpoints
|
||||
|
||||
|
||||
+81
-16
@@ -241,24 +241,89 @@ See the [Real-Time Chunking](./rtc) guide for details on tuning RTC parameters.
|
||||
|
||||
---
|
||||
|
||||
## Interactive Sessions
|
||||
|
||||
Add `--interactive=true` to drive the rollout from the terminal instead of starting immediately. Hardware connects and the policy loads as usual, but **the robot stays still until you type `/start`** — useful when you want to position the scene first, re-instruct the policy between attempts, or run several takes without paying the load time again.
|
||||
|
||||
```bash
|
||||
lerobot-rollout \
|
||||
--strategy.type=base \
|
||||
--policy.path=${HF_USER}/my_smolvla_policy \
|
||||
--robot.type=so100_follower \
|
||||
--robot.port=/dev/ttyACM0 \
|
||||
--robot.cameras="{ front: {type: opencv, index_or_path: 0, width: 640, height: 480, fps: 30}}" \
|
||||
--task="pick up the cube" \
|
||||
--interactive=true
|
||||
```
|
||||
|
||||
| Command | Action |
|
||||
| ----------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `/start` | Start (or restart) the policy control loop |
|
||||
| `/subtask <text>` | Change the instruction the policy follows, without stopping. No argument prints the current task. Only affects policies that condition on language (SmolVLA, π0/π0.5, and similar) |
|
||||
| `/ask <question>` | Ask a supported policy text head about its latest view. The answer is generated in the background without pausing the session |
|
||||
| `/reset` | Stop movement, return the robot to its startup position, and restore the `--task` instruction |
|
||||
| `/stop` | End the session and run the normal shutdown routines |
|
||||
| `/help` | List the commands |
|
||||
|
||||
```text
|
||||
> /start
|
||||
Rollout running — task 'pick up the cube'. /subtask <text> to change it, ...
|
||||
> /subtask put the cube in the box
|
||||
Task: 'pick up the cube' → 'put the cube in the box' (applies from the next policy inference)
|
||||
> /ask where is the red cube?
|
||||
Question queued: 'where is the red cube?' (the rollout keeps running)
|
||||
[policy] The red cube is beside the bowl.
|
||||
> /reset
|
||||
Task restored to 'pick up the cube'
|
||||
Resetting — returning the robot to its initial position...
|
||||
Robot reset — holding at initial position. /start to run.
|
||||
> /stop
|
||||
```
|
||||
|
||||
`Ctrl-C` still shuts down as usual, and closing stdin (`Ctrl-D`, or the end of a piped script) ends the session — so a piped script must keep stdin open for the intended duration:
|
||||
|
||||
```bash
|
||||
(printf '/start\n'; sleep 60; printf '/stop\n') | lerobot-rollout ... --interactive=true
|
||||
```
|
||||
|
||||
**How `/subtask` reaches the policy.** The stdin reader publishes the new instruction to the inference engine, which picks it up on its own inference thread, so nothing is mutated across threads while the robot is moving. How quickly the behavior changes depends on the backend:
|
||||
|
||||
- **Sync** (`--inference.type=sync`) — precomputed chunk actions are dropped, so the new instruction applies on the very next control tick. Without this a chunking policy would keep executing up to `chunk_size` stale actions (seconds of the old behavior). Only the queued actions are discarded, so observation history and the rest of the episode state are preserved.
|
||||
- **RTC** (`--inference.type=rtc`) — the next chunk is generated under the new instruction and merged over the previous chunk's leftover prefix, so the switch lands within one inference and the motion stays continuous. The queue is deliberately not cleared: that would leave the robot without commands for a full inference latency. (With blending turned off via `--inference.rtc.enabled=false` the queued chunk drains first, so the switch lands up to one chunk later.)
|
||||
|
||||
With `--use_torch_compile=true`, a switch whose instruction tokenizes to a different length can trigger a recompilation on the next forward pass, pausing inference for as long as the original warm-up took. Prefer leaving compilation off for sessions where you expect to re-instruct the policy often.
|
||||
|
||||
**How `/ask` runs without taking over the rollout.** During an active rollout, the inference engine caches the latest policy-ready observation, so the command reader never touches cameras, processors, or robot hardware. A single background worker sends that snapshot to the optional `PreTrainedPolicy.generate_text(..., kind=TextKind.VQA, user_text=question)` hook and prints the result when ready. Questions are independent turns; there is no conversation history, and a second question is rejected while one is running so stale image tensors cannot accumulate. WALL-OSS (`wall_x`) is the first policy implementing this hook; policies without a compatible text head report that `/ask` is unsupported.
|
||||
|
||||
Text and action calls share one policy safely: the engine gives a pending question priority after the current action inference finishes, while action inference uses a non-blocking gate. The hardware loop therefore keeps ticking and `/ask` never clears an action queue. RTC continues dispatching its buffered actions while text is decoded. Sync keeps the robot on its last commanded target until the policy is available again. Text generation still consumes model/GPU capacity, so response generation can reduce action freshness; RTC is preferred when uninterrupted action buffering matters.
|
||||
|
||||
`/stop` suppresses any late answer and gives an active decoder five seconds to finish cleanly. If it is stuck, hardware teardown continues rather than leaving the robot session open indefinitely; the daemon may retain its model/GPU resources until it returns.
|
||||
|
||||
**Console logs are muted while the session runs** so they don't interleave with what you're typing; they resume when it ends. A fatal inference error is still printed. Run without `--interactive` to watch the live log.
|
||||
|
||||
Interactive sessions currently require `--strategy.type=base`: the recording strategies finalize their dataset when their loop exits, so they cannot be restarted by `/start`, and their keyboard controls would compete for the same terminal.
|
||||
|
||||
---
|
||||
|
||||
## Common Flags
|
||||
|
||||
| Flag | Description | Default |
|
||||
| --------------------------------- | ----------------------------------------------------------------- | ------- |
|
||||
| `--policy.path` | **Required.** HF Hub model ID or local checkpoint path | -- |
|
||||
| `--robot.type` | **Required.** Robot type (e.g. `so100_follower`, `koch_follower`) | -- |
|
||||
| `--robot.port` | Serial port for the robot | -- |
|
||||
| `--robot.cameras` | Camera configuration (JSON dict) | -- |
|
||||
| `--fps` | Control loop frequency | 30 |
|
||||
| `--duration` | Run time in seconds (0 = infinite) | 0 |
|
||||
| `--device` | Torch device (`cpu`, `cuda`, `mps`) | auto |
|
||||
| `--task` | Task description (used when no dataset is provided) | -- |
|
||||
| `--display_data` | Stream telemetry to Rerun visualization | false |
|
||||
| `--display_ip` / `--display_port` | Remote Rerun server address | -- |
|
||||
| `--interpolation_multiplier` | Action interpolation factor | 1 |
|
||||
| `--use_torch_compile` | Enable `torch.compile` for inference | false |
|
||||
| `--resume` | Resume a previous recording session | false |
|
||||
| `--play_sounds` | Vocal synthesis for events | true |
|
||||
| Flag | Description | Default |
|
||||
| --------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------- | ------- |
|
||||
| `--policy.path` | **Required.** HF Hub model ID or local checkpoint path | -- |
|
||||
| `--robot.type` | **Required.** Robot type (e.g. `so100_follower`, `koch_follower`) | -- |
|
||||
| `--robot.port` | Serial port for the robot | -- |
|
||||
| `--robot.cameras` | Camera configuration (JSON dict) | -- |
|
||||
| `--fps` | Control loop frequency | 30 |
|
||||
| `--duration` | Run time in seconds (0 = infinite) | 0 |
|
||||
| `--device` | Torch device (`cpu`, `cuda`, `mps`) | auto |
|
||||
| `--task` | Task description (used when no dataset is provided) | -- |
|
||||
| `--display_data` | Stream telemetry to Rerun visualization | false |
|
||||
| `--display_ip` / `--display_port` | Remote Rerun server address | -- |
|
||||
| `--interpolation_multiplier` | Action interpolation factor | 1 |
|
||||
| `--interactive` | Chat-style stdin session (see [Interactive Sessions](#interactive-sessions)); the robot stays idle until `/start`. Base strategy only | false |
|
||||
| `--use_torch_compile` | Enable `torch.compile` for inference | false |
|
||||
| `--resume` | Resume a previous recording session | false |
|
||||
| `--play_sounds` | Vocal synthesis for events | true |
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -242,6 +242,17 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
||||
--repo-id=your_dataset
|
||||
```
|
||||
|
||||
Recording, resuming, and merging aggregate quantiles from per-episode summaries, so `meta/stats.json` ends up holding a conservative envelope (`min` for `q <= 50`, `max` for `q > 50`) rather than whole-dataset quantiles. To estimate the latter, scan every episode with a running histogram:
|
||||
|
||||
```bash
|
||||
python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
||||
--repo-id=your_dataset \
|
||||
--overwrite \
|
||||
--skip-images
|
||||
```
|
||||
|
||||
`--skip-images` keeps the existing image statistics and avoids video decoding when only `STATE`/`ACTION` need recomputing, and `--root` reads a local dataset instead of the Hub. These values are histogram estimates, subject to discretization and rebinning error, so they can differ from the conservative ones — which changes MolmoAct2's normalized targets and therefore its loss scale. Statistics already saved inside an existing checkpoint are not affected.
|
||||
|
||||
Alternatively, train MolmoAct2 with mean/std normalization:
|
||||
|
||||
```bash
|
||||
|
||||
@@ -127,6 +127,17 @@ lerobot-edit-dataset \
|
||||
|
||||
Or keep the dataset as-is and pass `--policy.normalization_mapping='{"ACTION": "MEAN_STD", "STATE": "MEAN_STD", "VISUAL": "IDENTITY"}'`.
|
||||
|
||||
Recording, resuming, and merging aggregate quantiles from per-episode summaries, so `meta/stats.json` ends up holding a conservative envelope (`min` for `q <= 50`, `max` for `q > 50`) rather than whole-dataset quantiles. To estimate the latter, scan every episode with a running histogram:
|
||||
|
||||
```bash
|
||||
python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
||||
--repo-id=your_dataset \
|
||||
--overwrite \
|
||||
--skip-images
|
||||
```
|
||||
|
||||
`--skip-images` keeps the existing image statistics and avoids video decoding when only `STATE`/`ACTION` need recomputing, and `--root` reads a local dataset instead of the Hub. These values are histogram estimates, subject to discretization and rebinning error, so they can differ from the conservative ones — which changes π₀.₅'s normalized targets and therefore its loss scale. Statistics already saved inside an existing checkpoint are not affected.
|
||||
|
||||
### Training Command Example
|
||||
|
||||
The same finetune with the VLM frozen: less memory, at some cost in success rate. Swap `--dataset.repo_id` for your own dataset.
|
||||
|
||||
@@ -2,6 +2,25 @@
|
||||
|
||||
https://diffusion-policy.cs.columbia.edu
|
||||
|
||||
## Training
|
||||
|
||||
The reference implementation maintains an exponential moving average (EMA) of the policy weights during training and evaluates the EMA weights. To reproduce this behavior, enable the trainer's EMA shadow:
|
||||
|
||||
```bash
|
||||
lerobot-train \
|
||||
--policy.type=diffusion \
|
||||
--ema.enable=true \
|
||||
...
|
||||
```
|
||||
|
||||
Checkpoints then contain a directly loadable copy of the EMA weights next to the live ones, e.g. for evaluation:
|
||||
|
||||
```bash
|
||||
lerobot-eval --policy.path=outputs/train/.../checkpoints/last/pretrained_model_ema ...
|
||||
```
|
||||
|
||||
The EMA decay schedule (`--ema.inv_gamma`, `--ema.power`, ...) defaults to the reference implementation's values. For a constant decay instead of the warmup schedule (e.g. to match openpi's pi0/pi05 training), set `--ema.decay=0.99`.
|
||||
|
||||
## Citation
|
||||
|
||||
```bibtex
|
||||
|
||||
@@ -59,6 +59,22 @@ When `use_relative_actions=true`, the training script automatically:
|
||||
|
||||
---
|
||||
|
||||
## EMA of the policy weights
|
||||
|
||||
OpenPI maintains an exponential moving average of the weights during training (`ema_decay=0.99` by default) and keeps the EMA copy for inference. To reproduce this with the LeRobot trainer, enable the EMA shadow with a constant decay:
|
||||
|
||||
```bash
|
||||
python -m lerobot.scripts.lerobot_train \
|
||||
--policy.type=pi05 \
|
||||
--dataset.repo_id=your_org/your_dataset \
|
||||
--ema.enable=true \
|
||||
--ema.decay=0.99
|
||||
```
|
||||
|
||||
Checkpoints then contain a directly loadable copy of the EMA weights in `pretrained_model_ema/` next to the live ones. Note that the shadow is a full extra copy of the parameters on the GPU. Like OpenPI (which disables EMA in its LoRA configs), EMA is not supported together with PEFT adapters.
|
||||
|
||||
---
|
||||
|
||||
## Citation
|
||||
|
||||
If you use this work, please cite both **OpenPI** and the π₀.₅ paper:
|
||||
|
||||
@@ -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.
|
||||
+17
-64
@@ -401,63 +401,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,24 +457,21 @@ 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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -22,7 +22,7 @@ Import them directly: ``from lerobot.configs.train import TrainPipelineConfig``
|
||||
"""
|
||||
|
||||
from .dataset import DatasetRecordConfig
|
||||
from .default import DatasetConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig
|
||||
from .default import DatasetConfig, EMAConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig
|
||||
from .policies import PreTrainedConfig
|
||||
from .recipe import MessageTurn, TrainingRecipe, load_recipe
|
||||
from .types import (
|
||||
@@ -31,6 +31,7 @@ from .types import (
|
||||
PipelineFeatureType,
|
||||
PolicyFeature,
|
||||
RTCAttentionSchedule,
|
||||
TextKind,
|
||||
)
|
||||
from .video import (
|
||||
DEFAULT_DEPTH_UNIT,
|
||||
@@ -54,9 +55,11 @@ __all__ = [
|
||||
"PipelineFeatureType",
|
||||
"PolicyFeature",
|
||||
"RTCAttentionSchedule",
|
||||
"TextKind",
|
||||
# Config classes
|
||||
"DatasetRecordConfig",
|
||||
"DatasetConfig",
|
||||
"EMAConfig",
|
||||
"EvalConfig",
|
||||
"JobConfig",
|
||||
"MessageTurn",
|
||||
|
||||
@@ -139,6 +139,59 @@ class EvalConfig:
|
||||
return min(by_cpu, self.n_episodes, 64)
|
||||
|
||||
|
||||
@dataclass
|
||||
class EMAConfig:
|
||||
"""Exponential moving average (EMA) of the policy weights.
|
||||
|
||||
Standard practice for diffusion-style policies (Chi et al. 2023, "Diffusion Policy", section V.D):
|
||||
the reference implementation enables it in every config and evaluates the EMA weights. Off by
|
||||
default here because it keeps a second full copy of the parameters in memory.
|
||||
|
||||
The decay follows the warmup schedule from diffusers' `EMAModel`:
|
||||
`decay_t = 1 - (1 + t / inv_gamma) ** -power`, clamped to `[min_decay, max_decay]`.
|
||||
The defaults mirror the reference implementation. Alternatively, set `decay` for a constant
|
||||
decay at every step, as used by openpi for pi0/pi05 (`ema_decay=0.99`).
|
||||
"""
|
||||
|
||||
enable: bool = False
|
||||
# Constant decay coefficient (openpi-style, e.g. 0.99 for pi0/pi05). When set, the warmup
|
||||
# schedule below is bypassed and the shadow uses this decay at every step.
|
||||
decay: float | None = None
|
||||
# Number of optimizer steps during which the shadow stays a hard copy of the live weights.
|
||||
update_after_step: int = 0
|
||||
# Warmup schedule parameters (see class docstring).
|
||||
inv_gamma: float = 1.0
|
||||
power: float = 0.75
|
||||
min_decay: float = 0.0
|
||||
max_decay: float = 0.9999
|
||||
# Evaluate the EMA weights (instead of the live ones) during periodic env eval.
|
||||
# Offline eval-loss (--eval_steps) always uses the live weights: it runs on every rank
|
||||
# while the EMA shadow only lives on the main process.
|
||||
use_for_eval: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not (0.0 <= self.min_decay <= self.max_decay <= 1.0):
|
||||
raise ValueError(
|
||||
"Expected 0 <= ema.min_decay <= ema.max_decay <= 1, got "
|
||||
f"min_decay={self.min_decay} and max_decay={self.max_decay}."
|
||||
)
|
||||
if self.inv_gamma <= 0:
|
||||
raise ValueError(f"ema.inv_gamma must be positive, got {self.inv_gamma}.")
|
||||
if self.power <= 0:
|
||||
raise ValueError(f"ema.power must be positive, got {self.power}.")
|
||||
if self.update_after_step < 0:
|
||||
raise ValueError(f"ema.update_after_step must be >= 0, got {self.update_after_step}.")
|
||||
if self.decay is not None:
|
||||
if not 0.0 <= self.decay <= 1.0:
|
||||
raise ValueError(f"ema.decay must be in [0, 1], got {self.decay}.")
|
||||
# Keep the literals in sync with the field defaults above.
|
||||
if self.min_decay != 0.0 or self.max_decay != 0.9999:
|
||||
raise ValueError(
|
||||
"ema.decay (constant decay) and ema.min_decay/ema.max_decay (schedule clamp) are "
|
||||
"mutually exclusive: set one or the other."
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PeftConfig:
|
||||
# PEFT offers many fine-tuning methods, layer adapters being the most common and currently also the most
|
||||
|
||||
@@ -67,6 +67,12 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
||||
# Whether the policy employed PEFT for training.
|
||||
use_peft: bool = False
|
||||
|
||||
# Decoding defaults for policies that implement `generate_text`. They live
|
||||
# in config.json so a text head uses the settings it was trained/evaluated
|
||||
# with; policy-specific decoding knobs belong on the concrete config.
|
||||
text_temperature: float = 0.0 # 0.0 = greedy; > 0 enables sampling
|
||||
text_top_p: float = 1.0
|
||||
|
||||
push_to_hub: bool = True # type: ignore[assignment] # TODO: use a different name to avoid override
|
||||
repo_id: str | None = None
|
||||
|
||||
|
||||
@@ -35,7 +35,7 @@ from lerobot.utils.hub import HubMixin, find_latest_hub_checkpoint
|
||||
from lerobot.utils.sample_weighting import SampleWeightingConfig
|
||||
|
||||
from . import parser
|
||||
from .default import DatasetConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig
|
||||
from .default import DatasetConfig, EMAConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig
|
||||
from .policies import PreTrainedConfig
|
||||
from .rewards import RewardModelConfig
|
||||
|
||||
@@ -163,6 +163,8 @@ class TrainPipelineConfig(HubMixin):
|
||||
# FSDP/DDP tuning knobs, compile & activation-checkpointing placeholders.
|
||||
accelerator: AcceleratorConfig = field(default_factory=AcceleratorConfig)
|
||||
eval: EvalConfig = field(default_factory=EvalConfig)
|
||||
# Maintain an EMA shadow of the policy weights during training (see EMAConfig).
|
||||
ema: EMAConfig = field(default_factory=EMAConfig)
|
||||
wandb: WandBConfig = field(default_factory=WandBConfig)
|
||||
peft: PeftConfig | None = None
|
||||
|
||||
|
||||
@@ -31,6 +31,13 @@ class PipelineFeatureType(str, Enum):
|
||||
OBSERVATION = "OBSERVATION"
|
||||
|
||||
|
||||
class TextKind(str, Enum):
|
||||
"""Text-generation requests understood by interactive policy hooks."""
|
||||
|
||||
SUBTASK = "subtask"
|
||||
VQA = "vqa"
|
||||
|
||||
|
||||
class NormalizationMode(str, Enum):
|
||||
MIN_MAX = "MIN_MAX"
|
||||
MEAN_STD = "MEAN_STD"
|
||||
|
||||
@@ -613,8 +613,15 @@ def aggregate_feature_stats(stats_ft_list: list[dict[str, dict]]) -> dict[str, d
|
||||
for q_key in quantile_keys:
|
||||
if all(q_key in s for s in stats_ft_list):
|
||||
quantile_values = np.stack([s[q_key] for s in stats_ft_list])
|
||||
weighted_quantiles = quantile_values * counts
|
||||
aggregated[q_key] = weighted_quantiles.sum(axis=0) / total_count
|
||||
# Exact global quantiles cannot be recovered from quantile summaries.
|
||||
# Keep a conservative envelope of the available estimates: min
|
||||
# for lower quantiles and max for upper quantiles. The resulting
|
||||
# values are bounds across the inputs, not global quantile estimates.
|
||||
q_percent = int(q_key[1:])
|
||||
if q_percent <= 50:
|
||||
aggregated[q_key] = np.min(quantile_values, axis=0)
|
||||
else:
|
||||
aggregated[q_key] = np.max(quantile_values, axis=0)
|
||||
|
||||
return aggregated
|
||||
|
||||
|
||||
@@ -33,11 +33,7 @@ if TYPE_CHECKING:
|
||||
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 behind the config system's back, making train_config.json lie about what ran.
|
||||
_ACCELERATE_ENV_VARS = (
|
||||
"ACCELERATE_USE_FSDP",
|
||||
"ACCELERATE_USE_PARALLELISM_CONFIG",
|
||||
@@ -59,11 +55,7 @@ def guard_against_env_interference() -> None:
|
||||
"""
|
||||
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)
|
||||
)
|
||||
offending = sorted(name for name in _ACCELERATE_ENV_VARS if name in os.environ)
|
||||
if offending:
|
||||
raise RuntimeError(
|
||||
f"Accelerate-configuring environment variables are set: {', '.join(offending)}. "
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -28,7 +28,8 @@ from huggingface_hub.errors import HfHubHTTPError
|
||||
from safetensors.torch import load_model as load_model_as_safetensor
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.configs import PreTrainedConfig
|
||||
from lerobot.configs import PreTrainedConfig, TextKind
|
||||
from lerobot.utils.constants import ACTION
|
||||
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
|
||||
@@ -210,6 +211,51 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def drop_queued_actions(self) -> None:
|
||||
"""Discard actions precomputed by earlier ``select_action`` calls.
|
||||
|
||||
Chunking policies answer most control ticks from a queue filled by an
|
||||
earlier forward pass, so a mid-episode change to the conditioning —
|
||||
e.g. a new language instruction — would otherwise only take effect
|
||||
once that queue drains (up to ``chunk_size`` ticks). Dropping the
|
||||
queue forces a fresh forward pass on the next ``select_action``.
|
||||
|
||||
Unlike :meth:`reset` this keeps the rest of the episode state (e.g.
|
||||
observation history), so it does not perturb policies that condition
|
||||
on it. Call it from the thread that calls ``select_action``: it
|
||||
mutates the same queues that thread pops from.
|
||||
|
||||
Policies that keep no action queue inherit a no-op.
|
||||
"""
|
||||
queues = getattr(self, "_queues", None)
|
||||
if isinstance(queues, dict) and ACTION in queues:
|
||||
queues[ACTION].clear()
|
||||
action_queue = getattr(self, "_action_queue", None)
|
||||
if action_queue is not None:
|
||||
action_queue.clear()
|
||||
|
||||
def supports_text_generation(self) -> bool:
|
||||
"""Whether this policy implements the optional :meth:`generate_text` hook."""
|
||||
return type(self).generate_text is not PreTrainedPolicy.generate_text
|
||||
|
||||
def generate_text(
|
||||
self,
|
||||
batch: dict[str, Tensor],
|
||||
*,
|
||||
kind: TextKind = TextKind.SUBTASK,
|
||||
user_text: str | None = None,
|
||||
) -> str:
|
||||
"""Generate one string from a policy's optional language head.
|
||||
|
||||
Interactive rollout calls this with a policy-ready observation batch.
|
||||
Implementations must treat the batch as read-only and avoid mutating
|
||||
action queues or episode state: text generation runs on a background
|
||||
worker while the control loop remains active.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
f"{type(self).__name__} has no text head. Implement `generate_text` to support /ask."
|
||||
)
|
||||
|
||||
def supports_rtc(self) -> bool:
|
||||
"""Whether this policy implements Real-Time Chunking inference semantics."""
|
||||
return False
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -52,6 +52,7 @@ from torch.nn import CrossEntropyLoss
|
||||
from torchvision.transforms import InterpolationMode
|
||||
from torchvision.transforms.v2 import functional as tv_functional
|
||||
|
||||
from lerobot.configs import TextKind
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
from lerobot.utils.import_utils import (
|
||||
_wallx_deps_available,
|
||||
@@ -107,6 +108,7 @@ else:
|
||||
|
||||
from .utils import (
|
||||
get_wallx_normal_text,
|
||||
img_key_mapping,
|
||||
preprocesser_call,
|
||||
process_grounding_points,
|
||||
replace_action_token,
|
||||
@@ -1585,6 +1587,25 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
- Handles special cases for input_embeds, generation methods, and GPU synchronization
|
||||
- Manages vision inputs to avoid unnecessary forward passes
|
||||
"""
|
||||
if cache_position is None:
|
||||
past_length = 0
|
||||
if past_key_values is not None and hasattr(past_key_values, "get_seq_length"):
|
||||
past_length = int(past_key_values.get_seq_length())
|
||||
input_length = input_ids.shape[1]
|
||||
end = input_length if input_length > past_length else past_length + input_length
|
||||
cache_position = torch.arange(
|
||||
past_length,
|
||||
end,
|
||||
dtype=torch.long,
|
||||
device=input_ids.device,
|
||||
)
|
||||
if cache_position.numel() == 0:
|
||||
cache_position = torch.arange(
|
||||
input_length,
|
||||
dtype=torch.long,
|
||||
device=input_ids.device,
|
||||
)
|
||||
|
||||
# Initialize MoE token types if not provided
|
||||
if moe_token_types is None:
|
||||
moe_token_types = torch.zeros_like(
|
||||
@@ -1851,6 +1872,23 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
"""Get parameters for optimization."""
|
||||
return self.parameters()
|
||||
|
||||
@staticmethod
|
||||
def _observation_prompt(img_keys: list[str]) -> str:
|
||||
prompt = "Observation:"
|
||||
for label in img_key_mapping(img_keys):
|
||||
prompt += f" {label}: <|vision_start|><|image_pad|><|vision_end|>"
|
||||
return prompt
|
||||
|
||||
def _format_text_prompt(self, instruction: str, kind: str, img_keys: list[str]) -> str:
|
||||
if kind == TextKind.SUBTASK:
|
||||
instruction = f"{instruction}\nPredict the next action in language."
|
||||
return (
|
||||
"<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n"
|
||||
f"<|im_start|>user\n{self._observation_prompt(img_keys)}\n"
|
||||
f"Instruction: {instruction}<|im_end|>\n"
|
||||
"<|im_start|>assistant\n"
|
||||
)
|
||||
|
||||
def preprocess_inputs(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
@@ -2080,6 +2118,118 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
|
||||
return loss, loss_dict
|
||||
|
||||
def _build_text_inputs(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
*,
|
||||
kind: str,
|
||||
user_text: str | list[str] | None,
|
||||
) -> BatchFeature:
|
||||
batch_size = batch[OBS_STATE].shape[0]
|
||||
img_keys = [key for key in self.config.image_features if key in batch]
|
||||
if not img_keys:
|
||||
raise ValueError("Wall-X text generation requires at least one image feature.")
|
||||
|
||||
image_inputs, dimensions_by_key = _prepare_wall_x_image_inputs(batch, img_keys)
|
||||
orig_height, orig_width, resized_height, resized_width = dimensions_by_key[img_keys[-1]]
|
||||
tasks = batch["task"] if isinstance(batch["task"], list) else [batch["task"]] * batch_size
|
||||
if user_text is None:
|
||||
instructions = tasks
|
||||
elif isinstance(user_text, str):
|
||||
instructions = [user_text] * batch_size
|
||||
elif len(user_text) == batch_size:
|
||||
instructions = user_text
|
||||
else:
|
||||
raise ValueError(f"Expected one text prompt for each of the {batch_size} samples.")
|
||||
|
||||
texts = [
|
||||
process_grounding_points(
|
||||
self._format_text_prompt(str(instruction), kind, img_keys),
|
||||
orig_height,
|
||||
orig_width,
|
||||
resized_height,
|
||||
resized_width,
|
||||
MODEL_TYPE,
|
||||
)
|
||||
for instruction in instructions
|
||||
]
|
||||
inputs = preprocesser_call(
|
||||
processor=self.model.processor,
|
||||
text=texts,
|
||||
images=image_inputs,
|
||||
videos=None,
|
||||
device=batch[OBS_STATE].device,
|
||||
padding=True,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
max_length=TOKENIZER_MAX_LENGTH,
|
||||
)
|
||||
inputs.pop("labels", None)
|
||||
inputs["moe_token_types"] = torch.zeros_like(inputs.input_ids, dtype=torch.bool)
|
||||
for key, value in inputs.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
inputs[key] = value.to(batch[OBS_STATE].device)
|
||||
return inputs
|
||||
|
||||
@torch.no_grad()
|
||||
def generate_text(
|
||||
self,
|
||||
batch: dict[str, Tensor],
|
||||
*,
|
||||
kind: TextKind = TextKind.SUBTASK,
|
||||
user_text: str | None = None,
|
||||
) -> str:
|
||||
"""Generate one grounded language response from the WALL-OSS VLM."""
|
||||
outputs = self.generate_texts(
|
||||
batch,
|
||||
kind=kind,
|
||||
user_text=user_text,
|
||||
temperature=self.config.text_temperature,
|
||||
top_p=self.config.text_top_p,
|
||||
)
|
||||
if len(outputs) != 1:
|
||||
raise ValueError(f"Interactive rollout expected one Wall-X output, got {len(outputs)}.")
|
||||
return outputs[0]
|
||||
|
||||
@torch.no_grad()
|
||||
def generate_texts(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
*,
|
||||
kind: TextKind = TextKind.VQA,
|
||||
user_text: str | list[str] | None = None,
|
||||
max_new_tokens: int = 100,
|
||||
min_new_tokens: int = 0,
|
||||
temperature: float = 0.0,
|
||||
top_p: float = 1.0,
|
||||
) -> list[str]:
|
||||
"""Generate grounded Wall-X text for one or more observations."""
|
||||
self.eval()
|
||||
if kind not in {TextKind.VQA, TextKind.SUBTASK}:
|
||||
raise ValueError("Unsupported Wall-X text kind.")
|
||||
inputs = self._build_text_inputs(batch, kind=kind, user_text=user_text)
|
||||
prompt_length = inputs.input_ids.shape[1]
|
||||
sampling = temperature > 0
|
||||
generation_kwargs: dict[str, Any] = {
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"min_new_tokens": min_new_tokens,
|
||||
"do_sample": sampling,
|
||||
"eos_token_id": self.model.processor.tokenizer.eos_token_id,
|
||||
"pad_token_id": self.model.processor.tokenizer.pad_token_id,
|
||||
"use_cache": True,
|
||||
}
|
||||
if sampling:
|
||||
generation_kwargs.update(temperature=temperature, top_p=top_p)
|
||||
output_ids = self.model.generate(**inputs, **generation_kwargs)
|
||||
return [
|
||||
value.strip()
|
||||
for value in self.model.processor.tokenizer.batch_decode(
|
||||
output_ids[:, prompt_length:],
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=True,
|
||||
)
|
||||
]
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
"""Predict action chunk for evaluation."""
|
||||
|
||||
@@ -217,12 +217,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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -358,11 +358,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
|
||||
@@ -456,10 +456,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 +557,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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -269,18 +269,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)
|
||||
|
||||
@@ -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
|
||||
@@ -168,10 +168,9 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
|
||||
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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -65,17 +65,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
|
||||
@@ -348,17 +346,12 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
|
||||
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:
|
||||
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: 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
|
||||
|
||||
+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
|
||||
|
||||
@@ -55,15 +55,6 @@ class SACAlgorithm(RLAlgorithm):
|
||||
policy: GaussianActorPolicy,
|
||||
config: SACAlgorithmConfig,
|
||||
):
|
||||
"""Build the critic ensemble, target networks, and temperature from `config`.
|
||||
|
||||
Args:
|
||||
policy (`GaussianActorPolicy`):
|
||||
The actor policy this algorithm trains. Its observation encoder is shared with the
|
||||
critics.
|
||||
config (`SACAlgorithmConfig`):
|
||||
Algorithm configuration.
|
||||
"""
|
||||
self.config = config
|
||||
self.policy_config = config.policy_config
|
||||
self.policy = policy
|
||||
@@ -153,18 +144,17 @@ class SACAlgorithm(RLAlgorithm):
|
||||
use_target: bool = False,
|
||||
observation_features: Tensor | None = None,
|
||||
) -> Tensor:
|
||||
"""Forward pass through a critic network ensemble.
|
||||
"""Forward pass through a critic network ensemble
|
||||
|
||||
Args:
|
||||
observations: Dictionary of observations
|
||||
actions: Action tensor
|
||||
use_target: If True, use target critics, otherwise use ensemble critics
|
||||
observation_features: Optional pre-computed observation features to avoid recomputing
|
||||
encoder output
|
||||
|
||||
Returns:
|
||||
Tensor of Q-values from all critics
|
||||
"""
|
||||
|
||||
critics = self.critic_target if use_target else self.critic_ensemble
|
||||
q_values = critics(observations, actions, observation_features)
|
||||
return q_values
|
||||
@@ -172,7 +162,7 @@ class SACAlgorithm(RLAlgorithm):
|
||||
def _discrete_critic_forward(
|
||||
self, observations, use_target=False, observation_features=None
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass through a discrete critic network.
|
||||
"""Forward pass through a discrete critic network
|
||||
|
||||
Args:
|
||||
observations: Dictionary of observations
|
||||
@@ -418,7 +408,7 @@ class SACAlgorithm(RLAlgorithm):
|
||||
return actor_loss
|
||||
|
||||
def _compute_loss_temperature(self, batch: dict[str, Any]) -> Tensor:
|
||||
"""Compute the temperature loss."""
|
||||
"""Compute the temperature loss"""
|
||||
observations = batch["state"]
|
||||
observation_features = batch.get("observation_feature")
|
||||
|
||||
@@ -430,7 +420,7 @@ class SACAlgorithm(RLAlgorithm):
|
||||
return temperature_loss
|
||||
|
||||
def _update_target_networks(self) -> None:
|
||||
"""Update target networks with exponential moving average."""
|
||||
"""Update target networks with exponential moving average"""
|
||||
for target_p, p in zip(
|
||||
self.critic_target.parameters(), self.critic_ensemble.parameters(), strict=True
|
||||
):
|
||||
@@ -471,7 +461,8 @@ class SACAlgorithm(RLAlgorithm):
|
||||
return forward_batch
|
||||
|
||||
def make_optimizers_and_scheduler(self) -> dict[str, Optimizer]:
|
||||
"""Creates and returns optimizers for the actor, critic, and temperature components of a reinforcement learning policy.
|
||||
"""
|
||||
Creates and returns optimizers for the actor, critic, and temperature components of a reinforcement learning policy.
|
||||
|
||||
This function sets up Adam optimizers for:
|
||||
- The **actor network**, ensuring that only relevant parameters are optimized.
|
||||
@@ -480,7 +471,7 @@ class SACAlgorithm(RLAlgorithm):
|
||||
|
||||
It also initializes a learning rate scheduler, though currently, it is set to `None`.
|
||||
|
||||
Note:
|
||||
NOTE:
|
||||
- If the encoder is shared, its parameters are excluded from the actor's optimization process.
|
||||
- The policy's log temperature (`log_alpha`) is wrapped in a list to ensure proper optimization as a standalone tensor.
|
||||
|
||||
@@ -505,7 +496,6 @@ class SACAlgorithm(RLAlgorithm):
|
||||
return self.optimizers
|
||||
|
||||
def get_optimizers(self) -> dict[str, Optimizer]:
|
||||
"""See [`~rl.algorithms.RLAlgorithm.get_optimizers`]."""
|
||||
return self.optimizers
|
||||
|
||||
def get_weights(self) -> dict[str, Any]:
|
||||
@@ -570,18 +560,20 @@ class SACAlgorithm(RLAlgorithm):
|
||||
def get_observation_features(
|
||||
self, observations: Tensor, next_observations: Tensor
|
||||
) -> tuple[Tensor | None, Tensor | None]:
|
||||
"""Get observation features from the policy encoder, acting as a cache.
|
||||
|
||||
When the encoder is frozen, the observation features are not updated, so we can save compute
|
||||
by caching them here instead of recomputing on every critic/actor forward pass.
|
||||
"""
|
||||
Get observation features from the policy encoder. It act as cache for the observation features.
|
||||
when the encoder is frozen, the observation features are not updated.
|
||||
We can save compute by caching the observation features.
|
||||
|
||||
Args:
|
||||
policy: The policy model
|
||||
observations: The current observations
|
||||
next_observations: The next observations
|
||||
|
||||
Returns:
|
||||
tuple: observation_features, next_observation_features
|
||||
"""
|
||||
|
||||
if self.policy.config.vision_encoder_name is None or not self.policy.config.freeze_vision_encoder:
|
||||
return None, None
|
||||
|
||||
@@ -603,8 +595,6 @@ def _split_prefix(state: dict[str, torch.Tensor], prefix: str) -> dict[str, torc
|
||||
|
||||
|
||||
class CriticHead(nn.Module):
|
||||
"""A single Q-value head: an MLP followed by a scalar linear output layer."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
@@ -615,23 +605,6 @@ class CriticHead(nn.Module):
|
||||
init_final: float | None = None,
|
||||
final_activation: Callable[[torch.Tensor], torch.Tensor] | str | None = None,
|
||||
):
|
||||
"""Build the MLP trunk and scalar output layer.
|
||||
|
||||
Args:
|
||||
input_dim (`int`): Dimension of the concatenated observation-encoding + action input.
|
||||
hidden_dims (`list[int]`): Hidden layer widths of the MLP trunk.
|
||||
activations (`Callable[[torch.Tensor], torch.Tensor] | str`, *optional*, defaults to `nn.SiLU()`):
|
||||
Activation used between hidden layers.
|
||||
activate_final (`bool`, *optional*, defaults to `False`): Whether to apply `activations`
|
||||
after the last hidden layer.
|
||||
dropout_rate (`float | None`, *optional*): Dropout probability applied between hidden
|
||||
layers. `None` disables dropout.
|
||||
init_final (`float | None`, *optional*): When set, the output layer's weight and bias are
|
||||
initialized uniformly in `[-init_final, init_final]` instead of the default
|
||||
orthogonal initialization.
|
||||
final_activation (`Callable[[torch.Tensor], torch.Tensor] | str | None`, *optional*):
|
||||
Activation applied after the MLP trunk's last hidden layer, before the output layer.
|
||||
"""
|
||||
super().__init__()
|
||||
self.net = MLP(
|
||||
input_dim=input_dim,
|
||||
@@ -649,17 +622,17 @@ class CriticHead(nn.Module):
|
||||
orthogonal_init()(self.output_layer.weight)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Compute the scalar Q-value for `x` (a concatenated observation-encoding + action tensor)."""
|
||||
return self.output_layer(self.net(x))
|
||||
|
||||
|
||||
class CriticEnsemble(nn.Module):
|
||||
"""CriticEnsemble wraps multiple CriticHead modules into an ensemble.
|
||||
"""
|
||||
CriticEnsemble wraps multiple CriticHead modules into an ensemble.
|
||||
|
||||
Args:
|
||||
encoder (GaussianActorObservationEncoder): encoder for observations.
|
||||
ensemble (List[CriticHead]): list of critic heads.
|
||||
init_final (float | None, *optional*): optional initializer scale for final layers.
|
||||
init_final (float | None): optional initializer scale for final layers.
|
||||
|
||||
Forward returns a tensor of shape (num_critics, batch_size) containing Q-values.
|
||||
"""
|
||||
@@ -670,14 +643,6 @@ class CriticEnsemble(nn.Module):
|
||||
ensemble: list[CriticHead],
|
||||
init_final: float | None = None,
|
||||
):
|
||||
"""Wrap `ensemble` behind the shared `encoder`.
|
||||
|
||||
Args:
|
||||
encoder (`GaussianActorObservationEncoder`): Shared observation encoder for all critics.
|
||||
ensemble (`list[CriticHead]`): The critic heads making up the ensemble.
|
||||
init_final (`float | None`, *optional*): Stored for introspection; each `CriticHead` is
|
||||
already initialized with it before being passed in here.
|
||||
"""
|
||||
super().__init__()
|
||||
self.encoder = encoder
|
||||
self.init_final = init_final
|
||||
@@ -689,19 +654,6 @@ class CriticEnsemble(nn.Module):
|
||||
actions: torch.Tensor,
|
||||
observation_features: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Encode `observations` and return each ensemble member's Q-value for `actions`.
|
||||
|
||||
Args:
|
||||
observations (`dict[str, torch.Tensor]`): Raw observation tensors, moved to the module's
|
||||
device.
|
||||
actions (`torch.Tensor`): Action tensor to evaluate.
|
||||
observation_features (`torch.Tensor | None`, *optional*): Pre-computed encoder output,
|
||||
e.g. from `SACAlgorithm.get_observation_features`. Bypasses re-encoding when the
|
||||
vision encoder is frozen.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Q-values of shape `(num_critics, batch_size)`.
|
||||
"""
|
||||
device = get_device_from_parameters(self)
|
||||
# Move each tensor in observations to device
|
||||
observations = {k: v.to(device) for k, v in observations.items()}
|
||||
|
||||
+27
-34
@@ -30,19 +30,6 @@ from lerobot.utils.transition import Transition
|
||||
|
||||
|
||||
class BatchTransition(TypedDict):
|
||||
"""A batch of transitions sampled from a `ReplayBuffer`.
|
||||
|
||||
**Attributes**:
|
||||
- **state** (`dict[str, torch.Tensor]`) -- Batched observation tensors at time `t`.
|
||||
- **action** (`torch.Tensor`) -- Batched actions taken at time `t`.
|
||||
- **reward** (`torch.Tensor`) -- Batched rewards received after `action`.
|
||||
- **next_state** (`dict[str, torch.Tensor]`) -- Batched observation tensors at time `t+1`.
|
||||
- **done** (`torch.Tensor`) -- Batched episode-termination flags.
|
||||
- **truncated** (`torch.Tensor`) -- Batched episode-truncation flags.
|
||||
- **complementary_info** (`dict[str, torch.Tensor | float | int] | None`) -- Optional extra
|
||||
per-transition data (e.g. intervention flags), when present in the underlying dataset.
|
||||
"""
|
||||
|
||||
state: dict[str, torch.Tensor]
|
||||
action: torch.Tensor
|
||||
reward: torch.Tensor
|
||||
@@ -53,7 +40,10 @@ class BatchTransition(TypedDict):
|
||||
|
||||
|
||||
def random_crop_vectorized(images: torch.Tensor, output_size: tuple) -> torch.Tensor:
|
||||
"""Perform a per-image random crop over a batch of images in a vectorized way."""
|
||||
"""
|
||||
Perform a per-image random crop over a batch of images in a vectorized way.
|
||||
(Same as shown previously.)
|
||||
"""
|
||||
B, C, H, W = images.shape # noqa: N806
|
||||
crop_h, crop_w = output_size
|
||||
|
||||
@@ -82,15 +72,13 @@ def random_crop_vectorized(images: torch.Tensor, output_size: tuple) -> torch.Te
|
||||
|
||||
|
||||
def random_shift(images: torch.Tensor, pad: int = 4):
|
||||
"""Vectorized random shift. `images` has shape `(B, C, H, W)`; `pad` is the shift range in pixels."""
|
||||
"""Vectorized random shift, imgs: (B,C,H,W), pad: #pixels"""
|
||||
_, _, h, w = images.shape
|
||||
images = F.pad(input=images, pad=(pad, pad, pad, pad), mode="replicate")
|
||||
return random_crop_vectorized(images=images, output_size=(h, w))
|
||||
|
||||
|
||||
class ReplayBuffer:
|
||||
"""In-memory replay buffer of `Transition`s, sampled in batches for off-policy RL training."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
capacity: int,
|
||||
@@ -101,12 +89,11 @@ class ReplayBuffer:
|
||||
storage_device: str = "cpu",
|
||||
optimize_memory: bool = False,
|
||||
):
|
||||
"""Replay buffer for storing transitions.
|
||||
|
||||
"""
|
||||
Replay buffer for storing transitions.
|
||||
It will allocate tensors on the specified device, when the first transition is added.
|
||||
NOTE: If you encounter memory issues, you can try to use the `optimize_memory` flag to save memory or
|
||||
and use the `storage_device` flag to store the buffer on a different device.
|
||||
|
||||
Args:
|
||||
capacity (int): Maximum number of transitions to store in the buffer.
|
||||
device (str): The device where the tensors will be moved when sampling ("cuda:0" or "cpu").
|
||||
@@ -200,7 +187,6 @@ class ReplayBuffer:
|
||||
self.initialized = True
|
||||
|
||||
def __len__(self):
|
||||
"""Number of transitions currently stored in the buffer."""
|
||||
return self.size
|
||||
|
||||
def add(
|
||||
@@ -319,8 +305,8 @@ class ReplayBuffer:
|
||||
async_prefetch: bool = True,
|
||||
queue_size: int = 2,
|
||||
):
|
||||
"""Creates an infinite iterator that yields batches of transitions.
|
||||
|
||||
"""
|
||||
Creates an infinite iterator that yields batches of transitions.
|
||||
Will automatically restart when internal iterator is exhausted.
|
||||
|
||||
Args:
|
||||
@@ -343,9 +329,10 @@ class ReplayBuffer:
|
||||
yield from iterator
|
||||
|
||||
def _get_async_iterator(self, batch_size: int, queue_size: int = 2):
|
||||
"""Create an iterator that continuously yields prefetched batches in a background thread.
|
||||
|
||||
The design is intentionally simple and avoids busy waiting / complex state management.
|
||||
"""
|
||||
Create an iterator that continuously yields prefetched batches in a
|
||||
background thread. The design is intentionally simple and avoids busy
|
||||
waiting / complex state management.
|
||||
|
||||
Args:
|
||||
batch_size (int): Size of batches to sample.
|
||||
@@ -396,7 +383,8 @@ class ReplayBuffer:
|
||||
producer_thread.join(timeout=1.0)
|
||||
|
||||
def _get_naive_iterator(self, batch_size: int, queue_size: int = 2):
|
||||
"""Creates a simple non-threaded iterator that yields batches.
|
||||
"""
|
||||
Creates a simple non-threaded iterator that yields batches.
|
||||
|
||||
Args:
|
||||
batch_size (int): Size of batches to sample
|
||||
@@ -410,7 +398,6 @@ class ReplayBuffer:
|
||||
queue = collections.deque()
|
||||
|
||||
def enqueue(n):
|
||||
"""Sample `n` more batches and append them to `queue`."""
|
||||
for _ in range(n):
|
||||
data = self.sample(batch_size)
|
||||
queue.append(data)
|
||||
@@ -432,7 +419,8 @@ class ReplayBuffer:
|
||||
storage_device: str = "cpu",
|
||||
optimize_memory: bool = False,
|
||||
) -> "ReplayBuffer":
|
||||
"""Convert a LeRobotDataset into a ReplayBuffer.
|
||||
"""
|
||||
Convert a LeRobotDataset into a ReplayBuffer.
|
||||
|
||||
Args:
|
||||
lerobot_dataset (LeRobotDataset): The dataset to convert.
|
||||
@@ -521,7 +509,9 @@ class ReplayBuffer:
|
||||
root=None,
|
||||
task_name="from_replay_buffer",
|
||||
) -> LeRobotDataset:
|
||||
"""Converts all transitions in this ReplayBuffer into a single LeRobotDataset object."""
|
||||
"""
|
||||
Converts all transitions in this ReplayBuffer into a single LeRobotDataset object.
|
||||
"""
|
||||
if self.size == 0:
|
||||
raise ValueError("The replay buffer is empty. Cannot convert to a dataset.")
|
||||
|
||||
@@ -622,7 +612,8 @@ class ReplayBuffer:
|
||||
dataset: LeRobotDataset,
|
||||
state_keys: Sequence[str] | None = None,
|
||||
) -> list[Transition]:
|
||||
"""Convert a LeRobotDataset into a list of RL (s, a, r, s', done) transitions.
|
||||
"""
|
||||
Convert a LeRobotDataset into a list of RL (s, a, r, s', done) transitions.
|
||||
|
||||
Args:
|
||||
dataset (LeRobotDataset):
|
||||
@@ -742,11 +733,12 @@ class ReplayBuffer:
|
||||
|
||||
# Utility function to guess shapes/dtypes from a tensor
|
||||
def guess_feature_info(t, name: str):
|
||||
"""Return a dictionary with the 'dtype' and 'shape' for a given tensor or scalar value.
|
||||
|
||||
"""
|
||||
Return a dictionary with the 'dtype' and 'shape' for a given tensor or scalar value.
|
||||
If it looks like a 3D (C,H,W) shape, we might consider it an 'image'.
|
||||
Otherwise default to appropriate dtype for numeric.
|
||||
"""
|
||||
|
||||
shape = tuple(t.shape)
|
||||
# Basic guess: if we have exactly 3 dims and shape[0] in {1, 3}, guess 'image'
|
||||
if len(shape) == 3 and shape[0] in [1, 3]:
|
||||
@@ -765,7 +757,8 @@ def guess_feature_info(t, name: str):
|
||||
def concatenate_batch_transitions(
|
||||
left_batch_transitions: BatchTransition, right_batch_transition: BatchTransition
|
||||
) -> BatchTransition:
|
||||
"""Concatenates two BatchTransition objects into one.
|
||||
"""
|
||||
Concatenates two BatchTransition objects into one.
|
||||
|
||||
This function merges the right BatchTransition into the left one by concatenating
|
||||
all corresponding tensors along dimension 0. The operation modifies the left_batch_transitions
|
||||
|
||||
@@ -29,7 +29,8 @@ from lerobot.utils.constants import DONE, REWARD
|
||||
|
||||
|
||||
def select_rect_roi(img):
|
||||
"""Allows the user to draw a rectangular ROI on the image.
|
||||
"""
|
||||
Allows the user to draw a rectangular ROI on the image.
|
||||
|
||||
The user must click and drag to draw the rectangle.
|
||||
- While dragging, the rectangle is dynamically drawn.
|
||||
@@ -51,7 +52,6 @@ def select_rect_roi(img):
|
||||
index_x, index_y = -1, -1 # Initial click coordinates
|
||||
|
||||
def mouse_callback(event, x, y, flags, param):
|
||||
"""`cv2.setMouseCallback` handler that drives the click-and-drag ROI selection."""
|
||||
nonlocal index_x, index_y, drawing, roi, working_img
|
||||
|
||||
if event == cv2.EVENT_LBUTTONDOWN:
|
||||
@@ -118,11 +118,12 @@ def select_rect_roi(img):
|
||||
|
||||
|
||||
def select_square_roi_for_images(images: dict) -> dict:
|
||||
"""For each image in the provided dictionary, open a window to allow the user to select a ROI.
|
||||
"""
|
||||
For each image in the provided dictionary, open a window to allow the user
|
||||
to select a rectangular ROI. Returns a dictionary mapping each key to a tuple
|
||||
(top, left, height, width) representing the ROI.
|
||||
|
||||
Returns a dictionary mapping each key to a tuple (top, left, height, width) representing the ROI.
|
||||
|
||||
Args:
|
||||
Parameters:
|
||||
images (dict): Dictionary where keys are identifiers and values are OpenCV images.
|
||||
|
||||
Returns:
|
||||
@@ -148,7 +149,9 @@ def select_square_roi_for_images(images: dict) -> dict:
|
||||
|
||||
|
||||
def get_image_from_lerobot_dataset(dataset: LeRobotDataset):
|
||||
"""Find the first row in the dataset and extract the image in order to be used for the crop."""
|
||||
"""
|
||||
Find the first row in the dataset and extract the image in order to be used for the crop.
|
||||
"""
|
||||
row = dataset[0]
|
||||
image_dict = {}
|
||||
for k in row:
|
||||
@@ -166,23 +169,19 @@ def convert_lerobot_dataset_to_cropped_lerobot_dataset(
|
||||
push_to_hub: bool = False,
|
||||
task: str = "",
|
||||
) -> LeRobotDataset:
|
||||
"""Converts an existing LeRobotDataset to a new one with cropped/resized image observations.
|
||||
|
||||
Iterates over the source dataset's episodes and frames, applying cropping and resizing to image
|
||||
observations, and saves a new dataset with the transformed data.
|
||||
"""
|
||||
Converts an existing LeRobotDataset by iterating over its episodes and frames,
|
||||
applying cropping and resizing to image observations, and saving a new dataset
|
||||
with the transformed data.
|
||||
|
||||
Args:
|
||||
original_dataset (`LeRobotDataset`): The source dataset.
|
||||
crop_params_dict (`dict[str, tuple[int, int, int, int]]`):
|
||||
original_dataset (LeRobotDataset): The source dataset.
|
||||
crop_params_dict (dict[str, Tuple[int, int, int, int]]):
|
||||
A dictionary mapping observation keys to crop parameters (top, left, height, width).
|
||||
new_repo_id (`str`): Repository id for the new dataset.
|
||||
new_dataset_root (`str`): The root directory where the new dataset will be written.
|
||||
resize_size (`tuple[int, int]`, *optional*, defaults to `(128, 128)`): The target size
|
||||
(height, width) after cropping.
|
||||
push_to_hub (`bool`, *optional*, defaults to `False`): Whether to push the new dataset to the
|
||||
Hugging Face Hub.
|
||||
task (`str`, *optional*, defaults to `""`): Task description recorded on every frame of the
|
||||
new dataset.
|
||||
new_repo_id (str): Repository id for the new dataset.
|
||||
new_dataset_root (str): The root directory where the new dataset will be written.
|
||||
resize_size (tuple[int, int], optional): The target size (height, width) after cropping.
|
||||
Defaults to (128, 128).
|
||||
|
||||
Returns:
|
||||
LeRobotDataset: A new LeRobotDataset where the specified image observations have been cropped
|
||||
|
||||
@@ -49,18 +49,6 @@ class OnlineOfflineMixer(DataMixer):
|
||||
offline_buffer: ReplayBuffer | None = None,
|
||||
online_ratio: float = 1.0,
|
||||
):
|
||||
"""Create the mixer.
|
||||
|
||||
Args:
|
||||
online_buffer (`ReplayBuffer`): Buffer of transitions collected online during training.
|
||||
offline_buffer (`ReplayBuffer | None`, *optional*): Buffer of pre-collected offline
|
||||
transitions. When `None`, every batch is drawn from `online_buffer` alone.
|
||||
online_ratio (`float`, *optional*, defaults to 1.0): Fraction of each batch drawn from
|
||||
`online_buffer`; the remainder comes from `offline_buffer`. Must be in `[0, 1]`.
|
||||
|
||||
Raises:
|
||||
ValueError: If `online_ratio` is not in `[0, 1]`.
|
||||
"""
|
||||
if not 0.0 <= online_ratio <= 1.0:
|
||||
raise ValueError(f"online_ratio must be in [0, 1], got {online_ratio}")
|
||||
self.online_buffer = online_buffer
|
||||
@@ -68,7 +56,6 @@ class OnlineOfflineMixer(DataMixer):
|
||||
self.online_ratio = online_ratio
|
||||
|
||||
def sample(self, batch_size: int) -> BatchType:
|
||||
"""See [`~rl.data_sources.DataMixer.sample`]."""
|
||||
if self.offline_buffer is None:
|
||||
return self.online_buffer.sample(batch_size)
|
||||
|
||||
@@ -86,6 +73,7 @@ class OnlineOfflineMixer(DataMixer):
|
||||
queue_size: int = 2,
|
||||
):
|
||||
"""Yield batches by composing buffer async iterators."""
|
||||
|
||||
n_online = max(1, int(batch_size * self.online_ratio))
|
||||
|
||||
online_iter = self.online_buffer.get_iterator(
|
||||
|
||||
@@ -36,13 +36,6 @@ logging.basicConfig(level=logging.INFO)
|
||||
|
||||
|
||||
def eval_policy(env, policy, n_episodes):
|
||||
"""Roll out `policy` in `env` for `n_episodes` and log the per-episode and average reward.
|
||||
|
||||
Args:
|
||||
env (`gymnasium.Env`): A robot environment, built via `make_robot_env`.
|
||||
policy (`PreTrainedPolicy`): A policy exposing `select_action(obs) -> action`.
|
||||
n_episodes (`int`): Number of episodes to run.
|
||||
"""
|
||||
sum_reward_episode = []
|
||||
for _ in range(n_episodes):
|
||||
obs, _ = env.reset()
|
||||
@@ -61,12 +54,6 @@ def eval_policy(env, policy, n_episodes):
|
||||
|
||||
@parser.wrap()
|
||||
def main(cfg: TrainRLServerPipelineConfig):
|
||||
"""CLI entry point: load a pretrained policy and evaluate it for 10 episodes.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Parsed from the CLI. `cfg.env.pretrained_policy_name_or_path`
|
||||
selects the checkpoint to load; `cfg.dataset.repo_id` provides normalization stats.
|
||||
"""
|
||||
env_cfg = cfg.env
|
||||
env = make_robot_env(env_cfg)
|
||||
dataset_cfg = cfg.dataset
|
||||
|
||||
@@ -305,9 +305,7 @@ def make_robot_env(cfg: HILSerlRobotEnvConfig) -> tuple[gym.Env, Any]:
|
||||
"""Create robot environment from configuration.
|
||||
|
||||
Args:
|
||||
cfg (`HILSerlRobotEnvConfig`): Environment configuration. `cfg.name == "gym_hil"` selects the
|
||||
GymHIL simulation environment; otherwise a real-robot `RobotEnv` is built from
|
||||
`cfg.robot`/`cfg.teleop`.
|
||||
cfg: Environment configuration.
|
||||
|
||||
Returns:
|
||||
Tuple of (gym environment, teleoperator device).
|
||||
@@ -365,13 +363,10 @@ def make_processors(
|
||||
"""Create environment and action processors.
|
||||
|
||||
Args:
|
||||
env (`Env`): The environment returned by `make_robot_env`.
|
||||
teleop_device (`lerobot.teleoperators.teleoperator.Teleoperator | None`): The teleoperator
|
||||
device returned by `make_robot_env`, used to configure intervention-related processor
|
||||
steps. `None` for simulation environments.
|
||||
cfg (`HILSerlRobotEnvConfig`): Environment configuration; provides the reward classifier,
|
||||
gripper, and reset-behavior settings for the built processor steps.
|
||||
device (`str`, *optional*, defaults to `"cpu"`): Torch device the processors run on.
|
||||
env: Robot environment instance.
|
||||
teleop_device: Teleoperator device for intervention.
|
||||
cfg: Processor configuration.
|
||||
device: Target device for computations.
|
||||
|
||||
Returns:
|
||||
Tuple of (environment processor, action processor).
|
||||
@@ -541,21 +536,20 @@ def step_env_and_process_transition(
|
||||
env_processor: DataProcessorPipeline[EnvTransition, EnvTransition],
|
||||
action_processor: DataProcessorPipeline[EnvTransition, EnvTransition],
|
||||
) -> EnvTransition:
|
||||
"""Execute one step with processor pipeline.
|
||||
"""
|
||||
Execute one step with processor pipeline.
|
||||
|
||||
Args:
|
||||
env (`Env`): The environment to step.
|
||||
transition (`EnvTransition`): The current transition; its observation is overwritten with the
|
||||
action processor's input before dispatch, then discarded.
|
||||
action (`Tensor`): The raw action to process and send to `env`.
|
||||
env_processor (`DataProcessorPipeline`): Post-processes the environment-produced transition
|
||||
(e.g. reward shaping, termination overrides).
|
||||
action_processor (`DataProcessorPipeline`): Pre-processes `action` before it reaches `env`
|
||||
(e.g. intervention overrides, gripper handling).
|
||||
env: The robot environment
|
||||
transition: Current transition state
|
||||
action: Action to execute
|
||||
env_processor: Environment processor
|
||||
action_processor: Action processor
|
||||
|
||||
Returns:
|
||||
Processed transition with updated state.
|
||||
"""
|
||||
|
||||
# Create action transition
|
||||
transition[TransitionKey.ACTION] = action
|
||||
transition[TransitionKey.OBSERVATION] = (
|
||||
@@ -624,16 +618,14 @@ def control_loop(
|
||||
cfg: GymManipulatorConfig,
|
||||
) -> None:
|
||||
"""Main control loop for robot environment interaction.
|
||||
|
||||
When `cfg.mode == "record"`, a dataset is created and recorded.
|
||||
if cfg.mode == "record": then a dataset will be created and recorded
|
||||
|
||||
Args:
|
||||
env (`Env`): The environment to control, built via `make_robot_env`.
|
||||
env_processor (`DataProcessorPipeline`): Post-processes environment-produced transitions.
|
||||
action_processor (`DataProcessorPipeline`): Pre-processes teleoperator actions before they
|
||||
reach `env`.
|
||||
teleop_device (`Teleoperator`): Teleoperator device driving the robot.
|
||||
cfg (`GymManipulatorConfig`): Control-loop configuration (mode, fps, episode/dataset settings).
|
||||
env: The robot environment
|
||||
env_processor: Environment processor
|
||||
action_processor: Action processor
|
||||
teleop_device: Teleoperator device
|
||||
cfg: gym_manipulator configuration
|
||||
"""
|
||||
dt = 1.0 / cfg.env.fps
|
||||
|
||||
|
||||
@@ -31,17 +31,18 @@ from lerobot.utils.constants import OBS_STATE
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register("joint_velocity_processor")
|
||||
class JointVelocityProcessorStep(ObservationProcessorStep):
|
||||
"""Calculates and appends joint velocity information to the observation state.
|
||||
"""
|
||||
Calculates and appends joint velocity information to the observation state.
|
||||
|
||||
This step computes the velocity of each joint by calculating the finite
|
||||
difference between the current and the last observed joint positions. The
|
||||
resulting velocity vector is then concatenated to the original state vector.
|
||||
|
||||
**Attributes**:
|
||||
- **dt** (`float`) -- The time step (delta time) in seconds between observations, used for calculating
|
||||
velocity.
|
||||
- **last_joint_positions** (`torch.Tensor | None`) -- Stores the joint positions from the previous
|
||||
step to enable velocity calculation.
|
||||
Attributes:
|
||||
dt: The time step (delta time) in seconds between observations, used for
|
||||
calculating velocity.
|
||||
last_joint_positions: Stores the joint positions from the previous step
|
||||
to enable velocity calculation.
|
||||
"""
|
||||
|
||||
dt: float = 0.1
|
||||
@@ -49,7 +50,8 @@ class JointVelocityProcessorStep(ObservationProcessorStep):
|
||||
last_joint_positions: torch.Tensor | None = None
|
||||
|
||||
def observation(self, observation: dict) -> dict:
|
||||
"""Computes joint velocities and adds them to the observation state.
|
||||
"""
|
||||
Computes joint velocities and adds them to the observation state.
|
||||
|
||||
Args:
|
||||
observation: The input observation dictionary, expected to contain
|
||||
@@ -87,7 +89,8 @@ class JointVelocityProcessorStep(ObservationProcessorStep):
|
||||
return new_observation
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
"""Returns the configuration of the step for serialization.
|
||||
"""
|
||||
Returns the configuration of the step for serialization.
|
||||
|
||||
Returns:
|
||||
A dictionary containing the time step `dt`.
|
||||
@@ -103,7 +106,8 @@ class JointVelocityProcessorStep(ObservationProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Updates the `observation.state` feature to reflect the added velocities.
|
||||
"""
|
||||
Updates the `observation.state` feature to reflect the added velocities.
|
||||
|
||||
This method doubles the size of the first dimension of the `observation.state`
|
||||
shape to account for the concatenation of position and velocity vectors.
|
||||
@@ -128,20 +132,22 @@ class JointVelocityProcessorStep(ObservationProcessorStep):
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register("current_processor")
|
||||
class MotorCurrentProcessorStep(ObservationProcessorStep):
|
||||
"""Reads motor currents from a robot and appends them to the observation state.
|
||||
"""
|
||||
Reads motor currents from a robot and appends them to the observation state.
|
||||
|
||||
This step queries the robot's hardware interface to get the present current
|
||||
for each motor and concatenates this information to the existing state vector.
|
||||
|
||||
**Attributes**:
|
||||
- **robot** (`Robot | None`) -- An instance of a `lerobot` Robot class that provides access to the
|
||||
hardware bus.
|
||||
Attributes:
|
||||
robot: An instance of a `lerobot` Robot class that provides access to
|
||||
the hardware bus.
|
||||
"""
|
||||
|
||||
robot: Robot | None = None
|
||||
|
||||
def observation(self, observation: dict) -> dict:
|
||||
"""Fetches motor currents and adds them to the observation state.
|
||||
"""
|
||||
Fetches motor currents and adds them to the observation state.
|
||||
|
||||
Args:
|
||||
observation: The input observation dictionary.
|
||||
@@ -178,7 +184,8 @@ class MotorCurrentProcessorStep(ObservationProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Updates the `observation.state` feature to reflect the added motor currents.
|
||||
"""
|
||||
Updates the `observation.state` feature to reflect the added motor currents.
|
||||
|
||||
This method increases the size of the first dimension of the `observation.state`
|
||||
shape by the number of motors in the robot.
|
||||
|
||||
+70
-85
@@ -14,7 +14,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.
|
||||
"""Learner server runner for distributed HILSerl robot policy training.
|
||||
"""
|
||||
Learner server runner for distributed HILSerl robot policy training.
|
||||
|
||||
This script implements the learner component of the distributed HILSerl architecture.
|
||||
It initializes the policy network, maintains replay buffers, and updates
|
||||
@@ -120,11 +121,6 @@ from .trainer import RLTrainer
|
||||
|
||||
@parser.wrap()
|
||||
def train_cli(cfg: TrainRLServerPipelineConfig):
|
||||
"""CLI entry point for the HILSerl learner server.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Parsed from the CLI, forwarded to `train`.
|
||||
"""
|
||||
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
||||
require_package("grpcio", extra="hilserl", import_name="grpc")
|
||||
if not use_threads(cfg):
|
||||
@@ -140,13 +136,14 @@ def train_cli(cfg: TrainRLServerPipelineConfig):
|
||||
|
||||
|
||||
def train(cfg: TrainRLServerPipelineConfig, job_name: str | None = None):
|
||||
"""Main training function that initializes and runs the training process.
|
||||
"""
|
||||
Main training function that initializes and runs the training process.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): The training configuration.
|
||||
job_name (`str | None`, *optional*): Job name for logging. Defaults to `cfg.job_name` when
|
||||
unset.
|
||||
cfg (TrainRLServerPipelineConfig): The training configuration
|
||||
job_name (str | None, optional): Job name for logging. Defaults to None.
|
||||
"""
|
||||
|
||||
cfg.validate()
|
||||
|
||||
if job_name is None:
|
||||
@@ -201,12 +198,13 @@ def start_learner_threads(
|
||||
wandb_logger: WandBLogger | None,
|
||||
shutdown_event: Any, # Event
|
||||
) -> None:
|
||||
"""Start the learner threads for training.
|
||||
"""
|
||||
Start the learner threads for training.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Training configuration.
|
||||
wandb_logger (`WandBLogger | None`): Logger for metrics.
|
||||
shutdown_event (`Event`): Event signaling the learner and its background workers to stop.
|
||||
cfg (TrainRLServerPipelineConfig): Training configuration
|
||||
wandb_logger (WandBLogger | None): Logger for metrics
|
||||
shutdown_event: Event to signal shutdown
|
||||
"""
|
||||
# Create multiprocessing queues
|
||||
transition_queue = Queue()
|
||||
@@ -277,7 +275,9 @@ def add_actor_information_and_train(
|
||||
interaction_message_queue: Queue,
|
||||
parameters_queue: Queue,
|
||||
):
|
||||
"""Handles data transfer from the actor to the learner, manages training updates, and logs progress.
|
||||
"""
|
||||
Handles data transfer from the actor to the learner, manages training updates,
|
||||
and logs training progress in an online reinforcement learning setup.
|
||||
|
||||
This function continuously:
|
||||
- Transfers transitions from the actor to the replay buffer.
|
||||
@@ -482,18 +482,17 @@ def start_learner(
|
||||
shutdown_event: Any, # Event
|
||||
cfg: TrainRLServerPipelineConfig,
|
||||
):
|
||||
"""Start the learner server for training.
|
||||
|
||||
Receives transitions and interaction messages from the actor server, and sends policy parameters
|
||||
to the actor server.
|
||||
"""
|
||||
Start the learner server for training.
|
||||
It will receive transitions and interaction messages from the actor server,
|
||||
and send policy parameters to the actor server.
|
||||
|
||||
Args:
|
||||
parameters_queue (`Queue`): Queue of serialized policy weights, drained and streamed to the
|
||||
actor by `LearnerService.StreamParameters`.
|
||||
transition_queue (`Queue`): Queue filled by `LearnerService.SendTransitions`.
|
||||
interaction_message_queue (`Queue`): Queue filled by `LearnerService.SendInteractions`.
|
||||
shutdown_event (`Event`): Event signaling this process/thread to stop.
|
||||
cfg (`TrainRLServerPipelineConfig`): Training configuration.
|
||||
parameters_queue: Queue for sending policy parameters to the actor
|
||||
transition_queue: Queue for receiving transitions from the actor
|
||||
interaction_message_queue: Queue for receiving interaction messages from the actor
|
||||
shutdown_event: Event to signal shutdown
|
||||
cfg: Training configuration
|
||||
"""
|
||||
if not use_threads(cfg):
|
||||
# Create a process-specific log file
|
||||
@@ -561,7 +560,8 @@ def save_training_checkpoint(
|
||||
preprocessor=None,
|
||||
postprocessor=None,
|
||||
) -> None:
|
||||
"""Save training checkpoint and associated data.
|
||||
"""
|
||||
Save training checkpoint and associated data.
|
||||
|
||||
This function performs the following steps:
|
||||
1. Creates a checkpoint directory with the current optimization step
|
||||
@@ -572,26 +572,18 @@ def save_training_checkpoint(
|
||||
6. If an offline replay buffer exists, saves it as a separate dataset
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Training configuration, saved alongside the checkpoint.
|
||||
optimization_step (`int`): Current optimization step; used to name the checkpoint directory.
|
||||
online_steps (`int`): Total number of online steps; used to size the checkpoint directory's
|
||||
zero-padded step number.
|
||||
interaction_message (`dict | None`): Latest interaction message; its `"Interaction step"`
|
||||
entry is saved for resuming training.
|
||||
policy (`Module`): Policy model to save.
|
||||
optimizers (`dict`): Dictionary of optimizers whose states are saved.
|
||||
replay_buffer (`ReplayBuffer`): Replay buffer to save as a dataset.
|
||||
algorithm (`lerobot.rl.algorithms.base.RLAlgorithm | None`, *optional*): Algorithm whose state
|
||||
dict (critic ensembles, temperature, etc.) should also be saved.
|
||||
offline_replay_buffer (`lerobot.rl.buffer.ReplayBuffer | None`, *optional*): Optional offline
|
||||
replay buffer, saved as a separate dataset when provided.
|
||||
dataset_repo_id (`str | None`, *optional*): Repository id used when converting the replay
|
||||
buffer(s) to a dataset.
|
||||
fps (`int`, *optional*, defaults to 30): Frames per second recorded on the saved dataset(s).
|
||||
preprocessor (`PolicyProcessorPipeline | None`, *optional*): Optional preprocessor pipeline to
|
||||
save alongside the policy.
|
||||
postprocessor (`PolicyProcessorPipeline | None`, *optional*): Optional postprocessor pipeline
|
||||
to save alongside the policy.
|
||||
cfg: Training configuration
|
||||
optimization_step: Current optimization step
|
||||
online_steps: Total number of online steps
|
||||
interaction_message: Dictionary containing interaction information
|
||||
policy: Policy model to save
|
||||
optimizers: Dictionary of optimizers
|
||||
replay_buffer: Replay buffer to save as dataset
|
||||
offline_replay_buffer: Optional offline replay buffer to save
|
||||
dataset_repo_id: Repository ID for dataset
|
||||
fps: Frames per second for dataset
|
||||
preprocessor: Optional preprocessor pipeline to save
|
||||
postprocessor: Optional postprocessor pipeline to save
|
||||
"""
|
||||
logging.info(f"Checkpoint policy after step {optimization_step}")
|
||||
_num_digits = max(6, len(str(online_steps)))
|
||||
@@ -658,7 +650,8 @@ def save_training_checkpoint(
|
||||
|
||||
|
||||
def handle_resume_logic(cfg: TrainRLServerPipelineConfig) -> TrainRLServerPipelineConfig:
|
||||
"""Handle the resume logic for training.
|
||||
"""
|
||||
Handle the resume logic for training.
|
||||
|
||||
If resume is True:
|
||||
- Verifies that a checkpoint exists
|
||||
@@ -719,19 +712,19 @@ def load_training_state(
|
||||
algorithm: RLAlgorithm | None = None,
|
||||
device: str | torch.device = "cpu",
|
||||
):
|
||||
"""Loads the training state from the most recent checkpoint.
|
||||
|
||||
Restores optimizers, RNG state, the optimization/interaction step, and algorithm-owned tensors.
|
||||
"""
|
||||
Loads the training state (optimizers, RNG, step + interaction step, and
|
||||
algorithm-owned tensors) from the most recent checkpoint.
|
||||
|
||||
Args:
|
||||
cfg (`TrainRLServerPipelineConfig`): Training configuration; `cfg.resume` gates the load and
|
||||
cfg (TrainRLServerPipelineConfig): Training configuration; `cfg.resume` gates the load and
|
||||
`cfg.output_dir` locates the last checkpoint.
|
||||
optimizers (`Optimizer | dict[str, Optimizer]`): Optimizers to load state into.
|
||||
algorithm (`RLAlgorithm | None`, *optional*): Algorithm whose state dict should be restored.
|
||||
optimizers (Optimizer | dict[str, Optimizer]): Optimizers to load state into.
|
||||
algorithm (RLAlgorithm | None, optional): Algorithm whose state dict should be restored.
|
||||
Required for full main-equivalent resume; the policy itself is restored separately via
|
||||
`make_policy`.
|
||||
device (`str | torch.device`, *optional*, defaults to `"cpu"`): Device on which to place
|
||||
loaded algorithm tensors.
|
||||
`make_policy`. Defaults to None.
|
||||
device (str | torch.device, optional): Device on which to place loaded algorithm tensors.
|
||||
Defaults to "cpu".
|
||||
|
||||
Returns:
|
||||
tuple[int | None, int | None]: `(optimization_step, interaction_step)`, or `(None, None)`
|
||||
@@ -779,7 +772,8 @@ def load_training_state(
|
||||
|
||||
|
||||
def log_training_info(cfg: TrainRLServerPipelineConfig, policy: nn.Module) -> None:
|
||||
"""Log information about the training process.
|
||||
"""
|
||||
Log information about the training process.
|
||||
|
||||
Args:
|
||||
cfg (TrainRLServerPipelineConfig): Training configuration
|
||||
@@ -798,7 +792,8 @@ def log_training_info(cfg: TrainRLServerPipelineConfig, policy: nn.Module) -> No
|
||||
def initialize_replay_buffer(
|
||||
cfg: TrainRLServerPipelineConfig, device: str, storage_device: str
|
||||
) -> ReplayBuffer:
|
||||
"""Initialize a replay buffer, either empty or from a dataset if resuming.
|
||||
"""
|
||||
Initialize a replay buffer, either empty or from a dataset if resuming.
|
||||
|
||||
Args:
|
||||
cfg (TrainRLServerPipelineConfig): Training configuration
|
||||
@@ -842,7 +837,8 @@ def initialize_offline_replay_buffer(
|
||||
device: str,
|
||||
storage_device: str,
|
||||
) -> ReplayBuffer:
|
||||
"""Initialize an offline replay buffer from a dataset.
|
||||
"""
|
||||
Initialize an offline replay buffer from a dataset.
|
||||
|
||||
Args:
|
||||
cfg (TrainRLServerPipelineConfig): Training configuration
|
||||
@@ -879,7 +875,6 @@ def initialize_offline_replay_buffer(
|
||||
|
||||
|
||||
def use_threads(cfg: TrainRLServerPipelineConfig) -> bool:
|
||||
"""Whether the learner's background workers should run as threads instead of processes."""
|
||||
return cfg.policy.concurrency.learner == "threads"
|
||||
|
||||
|
||||
@@ -889,14 +884,14 @@ def check_nan_in_transition(
|
||||
next_state: torch.Tensor,
|
||||
raise_error: bool = False,
|
||||
) -> bool:
|
||||
"""Check for NaN values in transition data.
|
||||
"""
|
||||
Check for NaN values in transition data.
|
||||
|
||||
Args:
|
||||
observations (`Tensor`): Dictionary of observation tensors.
|
||||
actions (`Tensor`): Action tensor.
|
||||
next_state (`Tensor`): Dictionary of next-observation tensors.
|
||||
raise_error (`bool`, *optional*, defaults to `False`): Whether to raise a `ValueError` instead
|
||||
of just logging when a NaN is found.
|
||||
observations: Dictionary of observation tensors
|
||||
actions: Action tensor
|
||||
next_state: Dictionary of next state tensors
|
||||
raise_error: If True, raises ValueError when NaN is detected
|
||||
|
||||
Returns:
|
||||
bool: True if NaN values were detected, False otherwise
|
||||
@@ -930,12 +925,6 @@ def check_nan_in_transition(
|
||||
|
||||
|
||||
def push_actor_policy_to_queue(parameters_queue: Queue, algorithm: RLAlgorithm) -> None:
|
||||
"""Serialize `algorithm`'s current weights and enqueue them for the actor-facing gRPC stream.
|
||||
|
||||
Args:
|
||||
parameters_queue (`Queue`): Queue drained by `LearnerService.StreamParameters`.
|
||||
algorithm (`RLAlgorithm`): Source of the weights, via `get_weights`.
|
||||
"""
|
||||
logging.debug("[LEARNER] Pushing actor policy to the queue")
|
||||
|
||||
# Create a dictionary to hold all the state dicts
|
||||
@@ -969,13 +958,11 @@ def process_transitions(
|
||||
"""Process all available transitions from the queue.
|
||||
|
||||
Args:
|
||||
transition_queue (`Queue`): Queue filled by `LearnerService.SendTransitions`.
|
||||
replay_buffer (`ReplayBuffer`): Buffer every non-NaN transition is added to.
|
||||
offline_replay_buffer (`ReplayBuffer`): Buffer intervention transitions are additionally added
|
||||
to, when `dataset_repo_id` is set.
|
||||
dataset_repo_id (`str | None`): When set, transitions tagged as interventions are also added
|
||||
to `offline_replay_buffer`.
|
||||
shutdown_event (`Event`): Event that stops the loop when set.
|
||||
transition_queue: Queue for receiving transitions from the actor
|
||||
replay_buffer: Replay buffer to add transitions to
|
||||
offline_replay_buffer: Offline replay buffer to add transitions to
|
||||
dataset_repo_id: Repository ID for dataset
|
||||
shutdown_event: Event to signal shutdown
|
||||
"""
|
||||
while not transition_queue.empty() and not shutdown_event.is_set():
|
||||
transition_list = transition_queue.get()
|
||||
@@ -1009,12 +996,10 @@ def process_interaction_messages(
|
||||
"""Process all available interaction messages from the queue.
|
||||
|
||||
Args:
|
||||
interaction_message_queue (`Queue`): Queue filled by `LearnerService.SendInteractions`.
|
||||
interaction_step_shift (`int`): Offset added to each message's `"Interaction step"` so it
|
||||
stays consistent with checkpointed state after a resume.
|
||||
wandb_logger (`lerobot.common.wandb_utils.WandBLogger | None`): Logger the message is
|
||||
forwarded to, when set.
|
||||
shutdown_event (`Event`): Event that stops the loop when set.
|
||||
interaction_message_queue: Queue for receiving interaction messages
|
||||
interaction_step_shift: Amount to shift interaction step by
|
||||
wandb_logger: Logger for tracking progress
|
||||
shutdown_event: Event to signal shutdown
|
||||
|
||||
Returns:
|
||||
dict | None: The last interaction message processed, or None if none were processed
|
||||
|
||||
@@ -44,10 +44,10 @@ SHUTDOWN_TIMEOUT = 10
|
||||
|
||||
|
||||
class LearnerService(_ServicerBase):
|
||||
"""Implementation of the LearnerService gRPC service.
|
||||
|
||||
Sends policy parameters to the actor and receives transitions and interactions from it; see
|
||||
`transport.proto` for the gRPC service definition.
|
||||
"""
|
||||
Implementation of the LearnerService gRPC service
|
||||
This service is used to send parameters to the Actor and receive transitions and interactions from the Actor
|
||||
check transport.proto for the gRPC service definition
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -59,19 +59,6 @@ class LearnerService(_ServicerBase):
|
||||
interaction_message_queue: Queue,
|
||||
queue_get_timeout: float = 0.001,
|
||||
):
|
||||
"""Create the servicer.
|
||||
|
||||
Args:
|
||||
shutdown_event (`Event`): Set to stop `StreamParameters`'s push loop.
|
||||
parameters_queue (`Queue`): Queue of serialized policy weights, drained and streamed to
|
||||
the actor by `StreamParameters`.
|
||||
seconds_between_pushes (`float`): Minimum interval between successive parameter pushes.
|
||||
transition_queue (`Queue`): Queue filled by `SendTransitions` with received transitions.
|
||||
interaction_message_queue (`Queue`): Queue filled by `SendInteractions` with received
|
||||
interaction messages.
|
||||
queue_get_timeout (`float`, *optional*, defaults to 0.001): Timeout used when polling
|
||||
`parameters_queue`.
|
||||
"""
|
||||
self.shutdown_event = shutdown_event
|
||||
self.parameters_queue = parameters_queue
|
||||
self.seconds_between_pushes = seconds_between_pushes
|
||||
@@ -82,17 +69,6 @@ class LearnerService(_ServicerBase):
|
||||
def StreamParameters( # noqa: N802
|
||||
self, request: "services_pb2.Empty", context: "grpc.ServicerContext"
|
||||
):
|
||||
"""GRPC server-streaming RPC: push the latest policy parameters to the actor.
|
||||
|
||||
Runs until `shutdown_event` is set, pushing at most once every `seconds_between_pushes`.
|
||||
|
||||
Args:
|
||||
request (`services_pb2.Empty`): Unused; required by the gRPC service signature.
|
||||
context (`grpc.ServicerContext`): gRPC call context.
|
||||
|
||||
Yields:
|
||||
Chunks of a `services_pb2.Parameters` message, produced by `send_bytes_in_chunks`.
|
||||
"""
|
||||
# TODO: authorize the request
|
||||
logging.info("[LEARNER] Received request to stream parameters from the Actor")
|
||||
|
||||
@@ -128,16 +104,6 @@ class LearnerService(_ServicerBase):
|
||||
return services_pb2.Empty()
|
||||
|
||||
def SendTransitions(self, request_iterator, _context: "grpc.ServicerContext"): # noqa: N802
|
||||
"""GRPC client-streaming RPC: receive transition chunks from the actor into `transition_queue`.
|
||||
|
||||
Args:
|
||||
request_iterator: Stream of `services_pb2.Transition` chunks sent by the actor's
|
||||
`transitions_stream`.
|
||||
_context (`grpc.ServicerContext`): gRPC call context.
|
||||
|
||||
Returns:
|
||||
services_pb2.Empty: Acknowledgement sent once the actor closes the stream.
|
||||
"""
|
||||
# TODO: authorize the request
|
||||
logging.info("[LEARNER] Received request to receive transitions from the Actor")
|
||||
|
||||
@@ -152,16 +118,6 @@ class LearnerService(_ServicerBase):
|
||||
return services_pb2.Empty()
|
||||
|
||||
def SendInteractions(self, request_iterator, _context: "grpc.ServicerContext"): # noqa: N802
|
||||
"""GRPC client-streaming RPC: receive interaction-message chunks into `interaction_message_queue`.
|
||||
|
||||
Args:
|
||||
request_iterator: Stream of `services_pb2.InteractionMessage` chunks sent by the actor's
|
||||
`interactions_stream`.
|
||||
_context (`grpc.ServicerContext`): gRPC call context.
|
||||
|
||||
Returns:
|
||||
services_pb2.Empty: Acknowledgement sent once the actor closes the stream.
|
||||
"""
|
||||
# TODO: authorize the request
|
||||
logging.info("[LEARNER] Received request to receive interactions from the Actor")
|
||||
|
||||
@@ -176,5 +132,4 @@ class LearnerService(_ServicerBase):
|
||||
return services_pb2.Empty()
|
||||
|
||||
def Ready(self, request: "services_pb2.Empty", context: "grpc.ServicerContext"): # noqa: N802
|
||||
"""GRPC health check: returns immediately, confirming the learner server is up."""
|
||||
return services_pb2.Empty()
|
||||
|
||||
@@ -23,19 +23,6 @@ from torch.multiprocessing import Queue
|
||||
|
||||
|
||||
def get_last_item_from_queue(queue: Queue, block=True, timeout: float = 0.1) -> Any:
|
||||
"""Drain `queue` and return only the most recently enqueued item.
|
||||
|
||||
Args:
|
||||
queue (`Queue`): A `torch.multiprocessing.Queue` to drain.
|
||||
block (`bool`, *optional*, defaults to `True`): Whether to block for up to `timeout` seconds
|
||||
waiting for a first item before draining. When `False`, returns `None` if the queue is
|
||||
currently empty.
|
||||
timeout (`float`, *optional*, defaults to 0.1): Seconds to wait for a first item when `block`
|
||||
is `True`.
|
||||
|
||||
Returns:
|
||||
Any: The most recent item, or `None` if the queue was (and stayed) empty.
|
||||
"""
|
||||
if block:
|
||||
try:
|
||||
item = queue.get(timeout=timeout)
|
||||
|
||||
@@ -28,100 +28,6 @@ from .algorithms.sac import SACAlgorithmConfig # noqa: F401
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class TrainRLServerPipelineConfig(TrainPipelineConfig):
|
||||
"""Top-level config for the actor/learner distributed RL training server.
|
||||
|
||||
Extends [`~configs.train.TrainPipelineConfig`] with an optional (rather than required) offline
|
||||
`dataset` and the RL-specific algorithm/data-mixing fields below.
|
||||
|
||||
Args:
|
||||
env (`lerobot.envs.configs.EnvConfig | None`, *optional*):
|
||||
Simulation environment configuration, used for `env_eval_freq` evaluation rollouts.
|
||||
policy (`lerobot.configs.policies.PreTrainedConfig | None`, *optional*):
|
||||
The actor policy configuration.
|
||||
reward_model (`lerobot.configs.rewards.RewardModelConfig | None`, *optional*):
|
||||
Reward model configuration, when training a reward model instead of a policy.
|
||||
output_dir (`pathlib.Path | None`, *optional*):
|
||||
Directory to save run outputs to. Reusing the same value across runs overwrites its
|
||||
contents unless `resume` is `True`.
|
||||
job_name (`str | None`, *optional*):
|
||||
Name used for logging and checkpoint directory naming.
|
||||
resume (`bool`, *optional*, defaults to `False`):
|
||||
Whether to resume a previous run from `--config_path`'s checkpoint.
|
||||
seed (`int | None`, *optional*, defaults to 1000):
|
||||
Random seed for model initialization, dataset shuffling, and evaluation environments.
|
||||
cudnn_deterministic (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use deterministic cuDNN algorithms for reproducibility. Disables
|
||||
`cudnn.benchmark`, which may reduce training speed.
|
||||
num_workers (`int`, *optional*, defaults to 4):
|
||||
Number of dataloader worker processes.
|
||||
batch_size (`int`, *optional*, defaults to 8):
|
||||
Offline-dataset dataloader batch size.
|
||||
prefetch_factor (`int`, *optional*, defaults to 4):
|
||||
Number of batches prefetched per dataloader worker.
|
||||
persistent_workers (`bool`, *optional*, defaults to `True`):
|
||||
Whether dataloader workers stay alive between epochs.
|
||||
dataloader_multiprocessing_context (`str | None`, *optional*, defaults to `"spawn"`):
|
||||
DataLoader worker start method. `None` uses Python's platform default.
|
||||
steps (`int`, *optional*, defaults to 100000):
|
||||
Total number of training steps.
|
||||
env_eval_freq (`int`, *optional*, defaults to 20000):
|
||||
Run the policy in the simulation environment every N steps to measure reward/success.
|
||||
`0` disables environment evaluation.
|
||||
log_freq (`int`, *optional*, defaults to 200):
|
||||
Log training metrics every N steps.
|
||||
eval_steps (`int`, *optional*, defaults to 0):
|
||||
Compute eval loss on held-out episodes every N steps. `0` disables it.
|
||||
max_eval_samples (`int`, *optional*, defaults to 0):
|
||||
Cap on total eval samples, split uniformly across tasks. `0` uses all held-out data.
|
||||
tolerance_s (`float`, *optional*, defaults to 0.0001):
|
||||
Maximum timestamp tolerance, in seconds, when loading dataset frames.
|
||||
save_checkpoint (`bool`, *optional*, defaults to `True`):
|
||||
Whether to save training checkpoints at all.
|
||||
save_freq (`int`, *optional*, defaults to 20000):
|
||||
Save a checkpoint every N training steps, and after the last step. A non-positive value
|
||||
disables periodic saving, keeping only the final checkpoint.
|
||||
checkpoint_format (`CheckpointFormat`, *optional*, defaults to `CheckpointFormat.SAFETENSORS`):
|
||||
Model-artifact format inside checkpoints.
|
||||
use_policy_training_preset (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use the policy's own recommended optimizer/scheduler preset when `optimizer`/
|
||||
`scheduler` are unset.
|
||||
optimizer (`lerobot.optim.optimizers.OptimizerConfig | None`, *optional*):
|
||||
Optimizer configuration override.
|
||||
scheduler (`lerobot.optim.schedulers.LRSchedulerConfig | None`, *optional*):
|
||||
Learning-rate scheduler configuration override.
|
||||
parallelism (`ParallelismConfig`, *optional*):
|
||||
Process topology: `dp_replicate`/`dp_shard` (HSDP) and context-parallel degree.
|
||||
accelerator (`AcceleratorConfig`, *optional*):
|
||||
Execution runtime handed to the Accelerator: mixed precision, gradient accumulation,
|
||||
FSDP/DDP tuning knobs, compile & activation-checkpointing.
|
||||
eval (`EvalConfig`, *optional*):
|
||||
Simulation-environment evaluation configuration (number of episodes, batch size).
|
||||
wandb (`WandBConfig`, *optional*):
|
||||
Weights & Biases logging configuration.
|
||||
peft (`lerobot.configs.default.PeftConfig | None`, *optional*):
|
||||
PEFT (e.g. LoRA) adapter configuration for parameter-efficient fine-tuning.
|
||||
job (`JobConfig`, *optional*):
|
||||
Where to run training: local (default) or an HF Jobs flavor.
|
||||
save_checkpoint_to_hub (`bool`, *optional*, defaults to `False`):
|
||||
Whether to push each saved checkpoint to the Hub as it is written, not just the final
|
||||
model.
|
||||
sample_weighting (`lerobot.utils.sample_weighting.SampleWeightingConfig | None`, *optional*):
|
||||
Sample weighting configuration (e.g. for RA-BC training).
|
||||
rename_map (`dict`, *optional*):
|
||||
Mapping to override observation image/state key names.
|
||||
dataset (`DatasetConfig | None`, *optional*):
|
||||
Optional offline dataset config. Unlike imitation-learning training, RL doesn't require an
|
||||
offline dataset — data comes from the online replay buffer.
|
||||
algorithm (`RLAlgorithmConfig | None`, *optional*):
|
||||
RL algorithm configuration. Defaults to a SAC config (with `policy_config` populated from
|
||||
`self.policy`) in `validate` when unset.
|
||||
mixer (`str`, *optional*, defaults to `"online_offline"`):
|
||||
Data mixer strategy name. Currently only `"online_offline"` is supported.
|
||||
online_ratio (`float`, *optional*, defaults to 0.5):
|
||||
Fraction of each training batch sampled from the online replay buffer when using
|
||||
`OnlineOfflineMixer`; the remainder comes from the offline dataset.
|
||||
"""
|
||||
|
||||
# NOTE: In RL, we don't need an offline dataset
|
||||
# TODO: Make `TrainPipelineConfig.dataset` optional
|
||||
dataset: DatasetConfig | None = None # type: ignore[assignment] # because the parent class has made it's type non-optional
|
||||
@@ -135,11 +41,6 @@ class TrainRLServerPipelineConfig(TrainPipelineConfig):
|
||||
online_ratio: float = 0.5
|
||||
|
||||
def validate(self) -> None:
|
||||
"""See [`~configs.train.TrainPipelineConfig.validate`].
|
||||
|
||||
Additionally defaults `algorithm` to a SAC config and populates its `policy_config` from
|
||||
`self.policy` when unset.
|
||||
"""
|
||||
super().validate()
|
||||
|
||||
if self.algorithm is None:
|
||||
|
||||
@@ -38,16 +38,6 @@ class RLTrainer:
|
||||
*,
|
||||
preprocessor: Any | None = None,
|
||||
):
|
||||
"""Build the trainer and its optimizers.
|
||||
|
||||
Args:
|
||||
algorithm (`RLAlgorithm`): The RL algorithm to train. `make_optimizers_and_scheduler` is
|
||||
called on it immediately.
|
||||
data_mixer (`DataMixer`): Data source the training-batch iterator is built from.
|
||||
batch_size (`int`): Batch size requested from `data_mixer` on each training step.
|
||||
preprocessor (`Any | None`, *optional*): When set, each sampled batch is passed through
|
||||
`preprocess_rl_batch` before reaching the algorithm.
|
||||
"""
|
||||
self.algorithm = algorithm
|
||||
self.data_mixer = data_mixer
|
||||
self.batch_size = batch_size
|
||||
@@ -100,15 +90,12 @@ class _PreprocessedIterator:
|
||||
__slots__ = ("_raw", "_preprocessor")
|
||||
|
||||
def __init__(self, raw_iterator: Iterator[BatchType], preprocessor: Any) -> None:
|
||||
"""Wrap `raw_iterator`, applying `preprocessor` to each yielded batch."""
|
||||
self._raw = raw_iterator
|
||||
self._preprocessor = preprocessor
|
||||
|
||||
def __iter__(self) -> _PreprocessedIterator:
|
||||
"""Return `self` (this object is its own iterator)."""
|
||||
return self
|
||||
|
||||
def __next__(self) -> BatchType:
|
||||
"""Return the next preprocessed batch from the wrapped iterator."""
|
||||
batch = next(self._raw)
|
||||
return preprocess_rl_batch(self._preprocessor, batch)
|
||||
|
||||
@@ -29,18 +29,14 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BiOpenArmFollower(BimanualMixin, Robot):
|
||||
"""A bimanual pair of OpenArm follower arms driven as one robot."""
|
||||
"""
|
||||
Bimanual OpenArm Follower Arms
|
||||
"""
|
||||
|
||||
config_class = BiOpenArmFollowerConfig
|
||||
name = "bi_openarm_follower"
|
||||
|
||||
def __init__(self, config: BiOpenArmFollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`BiOpenArmFollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
@@ -118,43 +114,19 @@ class BiOpenArmFollower(BimanualMixin, Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Motor ID configuration is typically done via manufacturer tools for CAN motors."
|
||||
)
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
obs_dict: RobotObservation = {}
|
||||
|
||||
# Add "left_" prefix to per-arm keys; keep top-level camera keys unprefixed.
|
||||
@@ -174,23 +146,6 @@ class BiOpenArmFollower(BimanualMixin, Robot):
|
||||
custom_kp: dict[str, float] | None = None,
|
||||
custom_kd: dict[str, float] | None = None,
|
||||
) -> RobotAction:
|
||||
"""Command both arms to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`], i.e. prefixed `left_` and
|
||||
`right_`.
|
||||
custom_kp (`dict[str, float]`, *optional*):
|
||||
Per-motor proportional gains for this step only. Defaults to each arm's `position_kp`.
|
||||
custom_kd (`dict[str, float]`, *optional*):
|
||||
Per-motor derivative gains for this step only. Defaults to each arm's `position_kd`.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
# Remove "left_" prefix
|
||||
left_action = {
|
||||
key.removeprefix("left_"): value for key, value in action.items() if key.startswith("left_")
|
||||
|
||||
@@ -25,28 +25,7 @@ from ..openarm_follower import OpenArmFollowerConfigBase
|
||||
@RobotConfig.register_subclass("bi_openarm_follower")
|
||||
@dataclass(kw_only=True)
|
||||
class BiOpenArmFollowerConfig(RobotConfig):
|
||||
"""Configuration for a bimanual pair of OpenArm follower arms.
|
||||
|
||||
The two arms are configured independently, then driven as one robot: observation and action keys from
|
||||
each arm are prefixed with `left_` and `right_`.
|
||||
|
||||
Calibration is per arm, taken from each arm config's own settings.
|
||||
|
||||
Args:
|
||||
id (`str`, *optional*, defaults to `"bi_openarm_follower"`):
|
||||
Identifier for the pair as a whole.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused at this level; each arm calibrates through its own config.
|
||||
left_arm_config (`OpenArmFollowerConfigBase`):
|
||||
Configuration for the left arm, including its own CAN interface. Set its `side` to `"left"` so
|
||||
the correct joint limits apply.
|
||||
right_arm_config (`OpenArmFollowerConfigBase`):
|
||||
Configuration for the right arm, including its own CAN interface. Set its `side` to `"right"`.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras not attached to either arm, such as an overhead view. These keys appear in
|
||||
observations unchanged, whereas cameras declared on an arm config are prefixed with that
|
||||
arm's side.
|
||||
"""
|
||||
"""Configuration class for Bi OpenArm Follower robots."""
|
||||
|
||||
id: str | None = "bi_openarm_follower"
|
||||
|
||||
|
||||
@@ -39,12 +39,6 @@ class BiRebotB601Follower(BimanualMixin, Robot):
|
||||
name = "bi_rebot_b601_follower"
|
||||
|
||||
def __init__(self, config: BiRebotB601FollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`BiRebotB601FollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
@@ -126,33 +120,14 @@ class BiRebotB601Follower(BimanualMixin, Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
obs_dict: RobotObservation = {}
|
||||
for k, v in self.left_arm.get_observation().items():
|
||||
obs_dict[k if k in self._top_level_cam_keys else f"left_{k}"] = v
|
||||
@@ -162,18 +137,6 @@ class BiRebotB601Follower(BimanualMixin, Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Command the robot to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`].
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
left_action = {
|
||||
key.removeprefix("left_"): value for key, value in action.items() if key.startswith("left_")
|
||||
}
|
||||
|
||||
@@ -25,27 +25,7 @@ from ..rebot_b601_follower import RebotB601FollowerConfig
|
||||
@RobotConfig.register_subclass("bi_rebot_b601_follower")
|
||||
@dataclass
|
||||
class BiRebotB601FollowerConfig(RobotConfig):
|
||||
"""Configuration for a bimanual pair of reBot B601-DM follower arms.
|
||||
|
||||
The two arms are configured independently, then driven as one robot: observation and action keys from
|
||||
each arm are prefixed with `left_` and `right_`.
|
||||
|
||||
Calibration is per arm, taken from each arm config's own `id` and `calibration_dir`.
|
||||
|
||||
Args:
|
||||
left_arm_config (`RebotB601FollowerConfig`):
|
||||
Configuration for the left arm, including its own `port` and CAN settings.
|
||||
right_arm_config (`RebotB601FollowerConfig`):
|
||||
Configuration for the right arm, including its own `port` and CAN settings.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras not attached to either arm, such as an overhead view. These keys appear in
|
||||
observations unchanged, whereas cameras declared on an arm config are prefixed with that
|
||||
arm's side.
|
||||
id (`str`, *optional*):
|
||||
Identifier for the pair as a whole.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused at this level; each arm calibrates through its own config.
|
||||
"""
|
||||
"""Configuration class for the bimanual reBot B601-DM follower robot."""
|
||||
|
||||
left_arm_config: RebotB601FollowerConfig
|
||||
right_arm_config: RebotB601FollowerConfig
|
||||
|
||||
@@ -29,18 +29,14 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BiSOFollower(BimanualMixin, Robot):
|
||||
"""A bimanual pair of [SO follower arms](https://github.com/TheRobotStudio/SO-ARM100) by TheRobotStudio."""
|
||||
"""
|
||||
[Bimanual SO Follower Arms](https://github.com/TheRobotStudio/SO-ARM100) designed by TheRobotStudio
|
||||
"""
|
||||
|
||||
config_class = BiSOFollowerConfig
|
||||
name = "bi_so_follower"
|
||||
|
||||
def __init__(self, config: BiSOFollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`BiSOFollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
@@ -111,42 +107,18 @@ class BiSOFollower(BimanualMixin, Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
self.left_arm.setup_motors()
|
||||
self.right_arm.setup_motors()
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
obs_dict: RobotObservation = {}
|
||||
|
||||
# Add "left_" prefix to per-arm keys; keep top-level camera keys unprefixed.
|
||||
@@ -162,18 +134,6 @@ class BiSOFollower(BimanualMixin, Robot):
|
||||
@check_if_not_connected
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
# Remove "left_" prefix
|
||||
"""Command the robot to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`].
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
left_action = {
|
||||
key.removeprefix("left_"): value for key, value in action.items() if key.startswith("left_")
|
||||
}
|
||||
|
||||
@@ -25,27 +25,7 @@ from ..so_follower import SOFollowerConfig
|
||||
@RobotConfig.register_subclass("bi_so_follower")
|
||||
@dataclass
|
||||
class BiSOFollowerConfig(RobotConfig):
|
||||
"""Configuration for a bimanual pair of SO follower arms.
|
||||
|
||||
The two arms are configured independently, then driven as one robot: observation and action keys from
|
||||
each arm are prefixed with `left_` and `right_`.
|
||||
|
||||
Calibration is per arm, taken from each arm config's own `id` and `calibration_dir`.
|
||||
|
||||
Args:
|
||||
left_arm_config (`SOFollowerConfig`):
|
||||
Configuration for the left arm, including its own `port`.
|
||||
right_arm_config (`SOFollowerConfig`):
|
||||
Configuration for the right arm, including its own `port`.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras not attached to either arm, such as an overhead view. These keys appear in
|
||||
observations unchanged, whereas cameras declared on an arm config are prefixed with that
|
||||
arm's side.
|
||||
id (`str`, *optional*):
|
||||
Identifier for the pair as a whole.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused at this level; each arm calibrates through its own config.
|
||||
"""
|
||||
"""Configuration class for Bi SO Follower robots."""
|
||||
|
||||
left_arm_config: SOFollowerConfig
|
||||
right_arm_config: SOFollowerConfig
|
||||
|
||||
@@ -21,33 +21,12 @@ import draccus
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class RobotConfig(draccus.ChoiceRegistry, abc.ABC):
|
||||
"""Base configuration shared by every robot.
|
||||
|
||||
Concrete robots subclass this and register themselves with
|
||||
`@RobotConfig.register_subclass("name")`, which is what makes `--robot.type=name` work on the command
|
||||
line. Subclasses inherit the two fields below and must document them alongside their own.
|
||||
|
||||
Args:
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular unit, used to tell apart several robots of the same type. It
|
||||
also names the calibration file, so keep it stable for a given piece of hardware.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to a per-robot directory under the
|
||||
LeRobot calibration home.
|
||||
"""
|
||||
|
||||
# Allows to distinguish between different robots of the same type
|
||||
id: str | None = None
|
||||
# Directory to store calibration file
|
||||
calibration_dir: Path | None = None
|
||||
|
||||
def __post_init__(self):
|
||||
"""Validate that every configured camera specifies the fields a robot requires.
|
||||
|
||||
Raises:
|
||||
ValueError: If a camera does not set `width`, `height` and `fps`. A robot records frames at a
|
||||
fixed shape, so these cannot be left to the driver's defaults.
|
||||
"""
|
||||
if hasattr(self, "cameras") and self.cameras:
|
||||
for _, config in self.cameras.items():
|
||||
for attr in ["width", "height", "fps"]:
|
||||
@@ -58,10 +37,4 @@ class RobotConfig(draccus.ChoiceRegistry, abc.ABC):
|
||||
|
||||
@property
|
||||
def type(self) -> str:
|
||||
"""The registered name of this robot type.
|
||||
|
||||
Returns:
|
||||
`str`: The name passed to `@RobotConfig.register_subclass`, e.g. `"so101_follower"`. This is
|
||||
what `make_robot_from_config` dispatches on and what a user writes as `--robot.type=...`.
|
||||
"""
|
||||
return self.get_choice_name(self.__class__)
|
||||
|
||||
@@ -23,18 +23,13 @@ from ..config import RobotConfig
|
||||
@RobotConfig.register_subclass("earthrover_mini_plus")
|
||||
@dataclass
|
||||
class EarthRoverMiniPlusConfig(RobotConfig):
|
||||
"""Configuration for the EarthRover Mini Plus rover.
|
||||
"""Configuration for EarthRover Mini Plus robot using Frodobots SDK.
|
||||
|
||||
This robot is driven over the cloud through the Frodobots SDK's HTTP API rather than a local bus, so
|
||||
there is no serial port and no LeRobot calibration file. Camera frames come from SDK HTTP endpoints.
|
||||
This robot uses cloud-based control via the Frodobots SDK HTTP API.
|
||||
Camera frames are accessed directly through SDK HTTP endpoints.
|
||||
|
||||
Args:
|
||||
sdk_url (`str`, *optional*, defaults to `"http://localhost:8000"`):
|
||||
Base URL of the Frodobots SDK server. Commands and camera frames both go through it.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular rover.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused: the rover exposes no calibration.
|
||||
Attributes:
|
||||
sdk_url: URL of the Frodobots SDK server (default: http://localhost:8000)
|
||||
"""
|
||||
|
||||
sdk_url: str = "http://localhost:8000"
|
||||
|
||||
@@ -70,7 +70,8 @@ OBS_WHEEL_RPM_3 = "wheel_rpm_3"
|
||||
|
||||
|
||||
class EarthRoverMiniPlus(Robot):
|
||||
"""EarthRover Mini Plus robot controlled via Frodobots SDK HTTP API.
|
||||
"""
|
||||
EarthRover Mini Plus robot controlled via Frodobots SDK HTTP API.
|
||||
|
||||
This robot uses cloud-based control through the Frodobots SDK instead of direct
|
||||
hardware connection. Cameras stream via WebRTC through Agora cloud, and control
|
||||
@@ -81,9 +82,9 @@ class EarthRoverMiniPlus(Robot):
|
||||
- Linear and angular velocity control
|
||||
- Battery and orientation telemetry
|
||||
|
||||
**Attributes**:
|
||||
- **config** -- Robot configuration
|
||||
- **sdk_base_url** -- URL of the Frodobots SDK server (default: http://localhost:8000)
|
||||
Attributes:
|
||||
config: Robot configuration
|
||||
sdk_base_url: URL of the Frodobots SDK server (default: http://localhost:8000)
|
||||
"""
|
||||
|
||||
config_class = EarthRoverMiniPlusConfig
|
||||
@@ -129,6 +130,7 @@ class EarthRoverMiniPlus(Robot):
|
||||
DeviceAlreadyConnectedError: If robot is already connected
|
||||
DeviceNotConnectedError: If cannot connect to SDK server
|
||||
"""
|
||||
|
||||
# Verify SDK is running and accessible
|
||||
try:
|
||||
response = requests.get(f"{self.sdk_base_url}/data", timeout=10.0)
|
||||
@@ -278,6 +280,7 @@ class EarthRoverMiniPlus(Robot):
|
||||
Robot telemetry is retrieved from /data endpoint.
|
||||
All SDK values are normalized to appropriate ranges for dataset recording.
|
||||
"""
|
||||
|
||||
observation = {}
|
||||
|
||||
# Get camera images from SDK
|
||||
@@ -367,6 +370,7 @@ class EarthRoverMiniPlus(Robot):
|
||||
Raises:
|
||||
DeviceNotConnectedError: If robot is not connected
|
||||
"""
|
||||
|
||||
# Stop the robot before disconnecting
|
||||
try:
|
||||
self._send_command_to_sdk(0.0, 0.0)
|
||||
|
||||
@@ -24,27 +24,6 @@ from ..config import RobotConfig
|
||||
@RobotConfig.register_subclass("hope_jr_hand")
|
||||
@dataclass
|
||||
class HopeJrHandConfig(RobotConfig):
|
||||
"""Configuration for one Hope Jr hand.
|
||||
|
||||
Each hand is a separate robot, so a two-handed setup uses two of these with different `side` and
|
||||
`port` values.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the hand is connected to. Run `lerobot-find-port` to identify it.
|
||||
side (`str`):
|
||||
Which hand this is, `"left"` or `"right"`. Determines the motor layout, so it must match the
|
||||
hardware.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the hand's joint positions.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular hand; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
port: str # Port to connect to the hand
|
||||
side: str # "left" / "right"
|
||||
|
||||
@@ -53,12 +32,6 @@ class HopeJrHandConfig(RobotConfig):
|
||||
cameras: dict[str, CameraConfig] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Validate the camera settings and the hand side.
|
||||
|
||||
Raises:
|
||||
ValueError: If `side` is not `"left"` or `"right"`, or if a camera omits `width`, `height` or
|
||||
`fps`.
|
||||
"""
|
||||
super().__post_init__()
|
||||
if self.side not in ["right", "left"]:
|
||||
raise ValueError(self.side)
|
||||
@@ -67,26 +40,6 @@ class HopeJrHandConfig(RobotConfig):
|
||||
@RobotConfig.register_subclass("hope_jr_arm")
|
||||
@dataclass
|
||||
class HopeJrArmConfig(RobotConfig):
|
||||
"""Configuration for one Hope Jr arm.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to. Run `lerobot-find-port` to identify it.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
port: str # Port to connect to the hand
|
||||
disable_torque_on_disconnect: bool = True
|
||||
|
||||
|
||||
@@ -35,22 +35,10 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HopeJrArm(Robot):
|
||||
"""One arm of the Hope Jr humanoid.
|
||||
|
||||
The arm and the hand are separate robots; pair this with [`~robots.hope_jr.HopeJrHand`] for a full
|
||||
limb. See [`~robots.Robot`] for the contract every method here implements.
|
||||
"""
|
||||
|
||||
config_class = HopeJrArmConfig
|
||||
name = "hope_jr_arm"
|
||||
|
||||
def __init__(self, config: HopeJrArmConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`HopeJrArmConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
self.bus = FeetechMotorsBus(
|
||||
@@ -89,47 +77,23 @@ class HopeJrArm(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus.is_connected and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect the motor bus and cameras, calibrating and configuring the arm.
|
||||
|
||||
> [!WARNING]
|
||||
> The arm is assumed to be at rest when this is called, because torque is disabled to run
|
||||
> calibration.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to run calibration if the arm is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
We assume that at connection time, arm is in a rest position,
|
||||
and torque can be safely disabled to run calibration.
|
||||
"""
|
||||
|
||||
self.bus.connect(handshake=False)
|
||||
if not self.is_calibrated and calibrate:
|
||||
self.calibrate()
|
||||
@@ -143,18 +107,9 @@ class HopeJrArm(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
groups = {
|
||||
"all": list(self.bus.motors.keys()),
|
||||
"shoulder": ["shoulder_pitch", "shoulder_yaw", "shoulder_roll"],
|
||||
@@ -167,17 +122,11 @@ class HopeJrArm(Robot):
|
||||
print("Calibration saved to", self.calibration_fpath)
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
with self.bus.torque_disabled():
|
||||
self.bus.configure_motors(maximum_acceleration=30, acceleration=30)
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
# TODO: add docstring
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
for motor in reversed(self.bus.motors):
|
||||
input(f"Connect the controller board to the '{motor}' motor only and press enter.")
|
||||
self.bus.setup_motor(motor)
|
||||
@@ -186,14 +135,6 @@ class HopeJrArm(Robot):
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
# Read arm position
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
obs_dict = self.bus.sync_read("Present_Position", self.other_motors)
|
||||
obs_dict[self.shoulder_pitch] = self.bus.read("Present_Position", self.shoulder_pitch)
|
||||
@@ -219,18 +160,6 @@ class HopeJrArm(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Command the robot to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`].
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
goal_pos = {key.removesuffix(".pos"): val for key, val in action.items() if key.endswith(".pos")}
|
||||
|
||||
# Cap goal position when too far away from present position.
|
||||
@@ -245,11 +174,6 @@ class HopeJrArm(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
for cam in self.cameras.values():
|
||||
cam.disconnect()
|
||||
|
||||
@@ -59,22 +59,10 @@ LEFT_HAND_INVERSIONS = [
|
||||
|
||||
|
||||
class HopeJrHand(Robot):
|
||||
"""One hand of the Hope Jr humanoid.
|
||||
|
||||
Each hand is its own robot, so a two-handed setup uses two of these with different `side` values. See
|
||||
[`~robots.Robot`] for the contract every method here implements.
|
||||
"""
|
||||
|
||||
config_class = HopeJrHandConfig
|
||||
name = "hope_jr_hand"
|
||||
|
||||
def __init__(self, config: HopeJrHandConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`HopeJrHandConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
self.bus = FeetechMotorsBus(
|
||||
@@ -125,43 +113,18 @@ class HopeJrHand(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus.is_connected and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect to the robot and its cameras, then apply the configured settings.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to run calibration if the robot is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
self.bus.connect()
|
||||
if not self.is_calibrated and calibrate:
|
||||
self.calibrate()
|
||||
@@ -175,18 +138,9 @@ class HopeJrHand(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
fingers = {}
|
||||
for finger in ["thumb", "index", "middle", "ring", "pinky"]:
|
||||
fingers[finger] = [motor for motor in self.bus.motors if motor.startswith(finger)]
|
||||
@@ -198,17 +152,11 @@ class HopeJrHand(Robot):
|
||||
print("Calibration saved to", self.calibration_fpath)
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
with self.bus.torque_disabled():
|
||||
self.bus.configure_motors()
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
# TODO: add docstring
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
for motor in self.bus.motors:
|
||||
input(f"Connect the controller board to the '{motor}' motor only and press enter.")
|
||||
self.bus.setup_motor(motor)
|
||||
@@ -216,14 +164,6 @@ class HopeJrHand(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
obs_dict = {}
|
||||
|
||||
# Read hand position
|
||||
@@ -251,29 +191,12 @@ class HopeJrHand(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Command the robot to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`].
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
goal_pos = {key.removesuffix(".pos"): val for key, val in action.items() if key.endswith(".pos")}
|
||||
self.bus.sync_write("Goal_Position", goal_pos)
|
||||
return action
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
for cam in self.cameras.values():
|
||||
cam.disconnect()
|
||||
|
||||
@@ -22,30 +22,6 @@ from ..config import RobotConfig
|
||||
@RobotConfig.register_subclass("koch_follower")
|
||||
@dataclass
|
||||
class KochFollowerConfig(RobotConfig):
|
||||
"""Configuration for the Koch v1.1 follower arm.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to, e.g. `/dev/ttyACM0` on Linux or `COM3` on Windows. Run
|
||||
`lerobot-find-port` to identify it.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions, keyed by the name they appear under in
|
||||
observations. Each must specify `width`, `height` and `fps`.
|
||||
use_degrees (`bool`, *optional*, defaults to `False`):
|
||||
Whether to report and accept joint positions in degrees rather than as a normalised range.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
# Port to connect to the arm
|
||||
port: str
|
||||
|
||||
|
||||
@@ -35,24 +35,16 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class KochFollower(Robot):
|
||||
"""The Koch follower arm, in either of its two revisions.
|
||||
|
||||
- [Koch v1.0](https://github.com/AlexanderKoch-Koch/low_cost_robot), with and without the
|
||||
wrist-to-elbow expansion, developed by Alexander Koch from
|
||||
[Tau Robotics](https://tau-robotics.com).
|
||||
- [Koch v1.1](https://github.com/jess-moss/koch-v1-1), developed by Jess Moss.
|
||||
"""
|
||||
- [Koch v1.0](https://github.com/AlexanderKoch-Koch/low_cost_robot), with and without the wrist-to-elbow
|
||||
expansion, developed by Alexander Koch from [Tau Robotics](https://tau-robotics.com)
|
||||
- [Koch v1.1](https://github.com/jess-moss/koch-v1-1) developed by Jess Moss
|
||||
"""
|
||||
|
||||
config_class = KochFollowerConfig
|
||||
name = "koch_follower"
|
||||
|
||||
def __init__(self, config: KochFollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`KochFollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
norm_mode_body = MotorNormMode.DEGREES if config.use_degrees else MotorNormMode.RANGE_M100_100
|
||||
@@ -87,47 +79,23 @@ class KochFollower(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus.is_connected and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect the motor bus and cameras, calibrating and configuring the arm.
|
||||
|
||||
> [!WARNING]
|
||||
> The arm is assumed to be at rest when this is called, because torque is disabled to run
|
||||
> calibration.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to run calibration if the arm is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
We assume that at connection time, arm is in a rest position,
|
||||
and torque can be safely disabled to run calibration.
|
||||
"""
|
||||
|
||||
self.bus.connect()
|
||||
if not self.is_calibrated and calibrate:
|
||||
logger.info(
|
||||
@@ -143,18 +111,9 @@ class KochFollower(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
self.bus.disable_torque()
|
||||
if self.calibration:
|
||||
# Calibration file exists, ask user whether to use it or run new calibration
|
||||
@@ -198,7 +157,6 @@ class KochFollower(Robot):
|
||||
logger.info(f"Calibration saved to {self.calibration_fpath}")
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
with self.bus.torque_disabled():
|
||||
self.bus.configure_motors()
|
||||
# Use 'extended position mode' for all motors except gripper, because in joint mode the servos
|
||||
@@ -223,11 +181,6 @@ class KochFollower(Robot):
|
||||
self.bus.write("Position_D_Gain", "elbow_flex", 600)
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
for motor in reversed(self.bus.motors):
|
||||
input(f"Connect the controller board to the '{motor}' motor only and press enter.")
|
||||
self.bus.setup_motor(motor)
|
||||
@@ -236,14 +189,6 @@ class KochFollower(Robot):
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
# Read arm position
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
obs_dict = self.bus.sync_read("Present_Position")
|
||||
obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
|
||||
@@ -280,6 +225,7 @@ class KochFollower(Robot):
|
||||
Returns:
|
||||
RobotAction: The action sent to the motors, potentially clipped.
|
||||
"""
|
||||
|
||||
goal_pos = {key.removesuffix(".pos"): val for key, val in action.items() if key.endswith(".pos")}
|
||||
|
||||
# Cap goal position when too far away from present position.
|
||||
@@ -295,11 +241,6 @@ class KochFollower(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
for cam in self.cameras.values():
|
||||
cam.disconnect()
|
||||
|
||||
@@ -21,12 +21,6 @@ from ..config import RobotConfig
|
||||
|
||||
|
||||
def lekiwi_cameras_config() -> dict[str, CameraConfig]:
|
||||
"""Build the default camera set for a LeKiwi base.
|
||||
|
||||
Returns:
|
||||
`dict[str, CameraConfig]`: The `front` and `wrist` OpenCV cameras at the device paths and
|
||||
rotations of a standard LeKiwi build. Override the `cameras` field if yours is wired differently.
|
||||
"""
|
||||
return {
|
||||
"front": OpenCVCameraConfig(
|
||||
index_or_path="/dev/video0",
|
||||
@@ -50,34 +44,6 @@ def lekiwi_cameras_config() -> dict[str, CameraConfig]:
|
||||
@RobotConfig.register_subclass("lekiwi")
|
||||
@dataclass
|
||||
class LeKiwiConfig(RobotConfig):
|
||||
"""Configuration for LeKiwi, running on the robot itself.
|
||||
|
||||
This is the config used by the process on the LeKiwi's own computer. To drive one from another machine,
|
||||
use [`LeKiwiClientConfig`] instead.
|
||||
|
||||
Args:
|
||||
port (`str`, *optional*, defaults to `"/dev/ttyACM0"`):
|
||||
Serial port of the motor bus on the robot's computer.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the joint positions. Defaults to the standard `front` and `wrist`
|
||||
build; see [`lekiwi_cameras_config`].
|
||||
use_degrees (`bool`, *optional*, defaults to `True`):
|
||||
Whether to report and accept arm joint positions in degrees.
|
||||
num_read_retries (`int`, *optional*, defaults to 2):
|
||||
Extra attempts when a `sync_read` fails. Feetech buses occasionally return a corrupted status
|
||||
packet, which would otherwise abort the control loop.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular robot; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
port: str = "/dev/ttyACM0" # port to connect to the bus
|
||||
|
||||
disable_torque_on_disconnect: bool = True
|
||||
@@ -101,22 +67,6 @@ class LeKiwiConfig(RobotConfig):
|
||||
|
||||
@dataclass
|
||||
class LeKiwiHostConfig:
|
||||
"""Configuration for the host process that serves a LeKiwi over the network.
|
||||
|
||||
Args:
|
||||
port_zmq_cmd (`int`, *optional*, defaults to 5555):
|
||||
ZMQ port the host listens on for actions.
|
||||
port_zmq_observations (`int`, *optional*, defaults to 5556):
|
||||
ZMQ port the host publishes observations on.
|
||||
connection_time_s (`int`, *optional*, defaults to 30):
|
||||
How long the host stays up before shutting down.
|
||||
watchdog_timeout_ms (`int`, *optional*, defaults to 500):
|
||||
Stop the robot if no command arrives within this window. Guards against a dropped client
|
||||
leaving the base driving.
|
||||
max_loop_freq_hz (`int`, *optional*, defaults to 30):
|
||||
Control loop frequency. Lower it if the robot jitters, and watch CPU load with `top`.
|
||||
"""
|
||||
|
||||
# Network Configuration
|
||||
port_zmq_cmd: int = 5555
|
||||
port_zmq_observations: int = 5556
|
||||
@@ -134,32 +84,6 @@ class LeKiwiHostConfig:
|
||||
@RobotConfig.register_subclass("lekiwi_client")
|
||||
@dataclass
|
||||
class LeKiwiClientConfig(RobotConfig):
|
||||
"""Configuration for driving a LeKiwi from another machine.
|
||||
|
||||
Presents the same [`~robots.Robot`] interface as the robot-side [`LeKiwiConfig`], but every call goes
|
||||
over ZMQ to the host process. Calibration lives on the robot, so nothing here configures it.
|
||||
|
||||
Args:
|
||||
remote_ip (`str`):
|
||||
IP address of the LeKiwi's computer on the network.
|
||||
port_zmq_cmd (`int`, *optional*, defaults to 5555):
|
||||
ZMQ port to send actions to. Must match the host's `port_zmq_cmd`.
|
||||
port_zmq_observations (`int`, *optional*, defaults to 5556):
|
||||
ZMQ port to receive observations on. Must match the host's `port_zmq_observations`.
|
||||
teleop_keys (`dict[str, str]`, *optional*):
|
||||
Keyboard bindings for driving the base: movement, rotation, speed control and quit.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras expected in the observation stream. Defaults to the standard `front` and `wrist` build.
|
||||
polling_timeout_ms (`int`, *optional*, defaults to 15):
|
||||
How long to wait for an observation before giving up on that step.
|
||||
connect_timeout_s (`int`, *optional*, defaults to 5):
|
||||
How long to wait for the host to answer when connecting.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular robot.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused by the client: calibration is held on the robot.
|
||||
"""
|
||||
|
||||
# Network Configuration
|
||||
remote_ip: str
|
||||
port_zmq_cmd: int = 5555
|
||||
|
||||
@@ -39,25 +39,17 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LeKiwi(Robot):
|
||||
"""A three-omniwheel mobile base with a follower arm on top, running on the robot itself.
|
||||
|
||||
The leader arm is connected to the operator's laptop; its joint positions are recorded and forwarded
|
||||
to this follower arm after a safety clamp. In parallel, keyboard teleoperation generates raw velocity
|
||||
commands for the wheels.
|
||||
|
||||
To drive one of these from another machine, use [`~robots.lekiwi.LeKiwiClient`].
|
||||
"""
|
||||
The robot includes a three omniwheel mobile base and a remote follower arm.
|
||||
The leader arm is connected locally (on the laptop) and its joint positions are recorded and then
|
||||
forwarded to the remote follower arm (after applying a safety clamp).
|
||||
In parallel, keyboard teleoperation is used to generate raw velocity commands for the wheels.
|
||||
"""
|
||||
|
||||
config_class = LeKiwiConfig
|
||||
name = "lekiwi"
|
||||
|
||||
def __init__(self, config: LeKiwiConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`LeKiwiConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
norm_mode_body = MotorNormMode.DEGREES if config.use_degrees else MotorNormMode.RANGE_M100_100
|
||||
@@ -113,43 +105,18 @@ class LeKiwi(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._state_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._state_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus.is_connected and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect to the robot and its cameras, then apply the configured settings.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to run calibration if the robot is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
self.bus.connect()
|
||||
if not self.is_calibrated and calibrate:
|
||||
logger.info(
|
||||
@@ -165,18 +132,9 @@ class LeKiwi(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
if self.calibration:
|
||||
# Calibration file exists, ask user whether to use it or run new calibration
|
||||
user_input = input(
|
||||
@@ -231,7 +189,6 @@ class LeKiwi(Robot):
|
||||
# Set-up arm actuators (position mode)
|
||||
# We assume that at connection time, arm is in a rest position,
|
||||
# and torque can be safely disabled to run calibration.
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
self.bus.disable_torque()
|
||||
self.bus.configure_motors()
|
||||
for name in self.arm_motors:
|
||||
@@ -248,11 +205,6 @@ class LeKiwi(Robot):
|
||||
self.bus.enable_torque()
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
for motor in chain(reversed(self.arm_motors), reversed(self.base_motors)):
|
||||
input(f"Connect the controller board to the '{motor}' motor only and press enter.")
|
||||
self.bus.setup_motor(motor)
|
||||
@@ -286,7 +238,8 @@ class LeKiwi(Robot):
|
||||
base_radius: float = 0.125,
|
||||
max_raw: int = 3000,
|
||||
) -> dict:
|
||||
"""Convert desired body-frame velocities into wheel raw commands.
|
||||
"""
|
||||
Convert desired body-frame velocities into wheel raw commands.
|
||||
|
||||
Parameters:
|
||||
x_cmd : Linear velocity in x (m/s).
|
||||
@@ -349,7 +302,8 @@ class LeKiwi(Robot):
|
||||
wheel_radius: float = 0.05,
|
||||
base_radius: float = 0.125,
|
||||
) -> dict[str, Any]:
|
||||
"""Convert wheel raw command feedback back into body-frame velocities.
|
||||
"""
|
||||
Convert wheel raw command feedback back into body-frame velocities.
|
||||
|
||||
Parameters:
|
||||
wheel_raw : Vector with raw wheel commands ("base_left_wheel", "base_back_wheel", "base_right_wheel").
|
||||
@@ -359,6 +313,7 @@ class LeKiwi(Robot):
|
||||
Returns:
|
||||
A dict (x.vel, y.vel, theta.vel) all in m/s
|
||||
"""
|
||||
|
||||
# Convert each raw command back to an angular speed in deg/s.
|
||||
wheel_degps = np.array(
|
||||
[
|
||||
@@ -391,14 +346,6 @@ class LeKiwi(Robot):
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
# Read actuators position for arm and vel for base
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
arm_pos = self.bus.sync_read(
|
||||
"Present_Position", self.arm_motors, num_retry=self.config.num_read_retries
|
||||
@@ -443,6 +390,7 @@ class LeKiwi(Robot):
|
||||
Returns:
|
||||
RobotAction: the action sent to the motors, potentially clipped.
|
||||
"""
|
||||
|
||||
arm_goal_pos = {k: v for k, v in action.items() if k.endswith(".pos")}
|
||||
base_goal_vel = {k: v for k, v in action.items() if k.endswith(".vel")}
|
||||
|
||||
@@ -471,17 +419,11 @@ class LeKiwi(Robot):
|
||||
return {**arm_goal_pos, **base_goal_vel}
|
||||
|
||||
def stop_base(self):
|
||||
"""Bring the mobile base to a halt by commanding zero velocity on its wheels."""
|
||||
self.bus.sync_write("Goal_Velocity", dict.fromkeys(self.base_motors, 0), num_retry=5)
|
||||
logger.info("Base motors stopped")
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.stop_base()
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
for cam in self.cameras.values():
|
||||
|
||||
@@ -31,23 +31,10 @@ from .config_lekiwi import LeKiwiClientConfig
|
||||
|
||||
|
||||
class LeKiwiClient(Robot):
|
||||
"""Drives a LeKiwi over the network from another machine.
|
||||
|
||||
Presents the same [`~robots.Robot`] interface as [`~robots.lekiwi.LeKiwi`], but every observation and
|
||||
action crosses a ZMQ connection to the host process running on the robot. Calibration stays on the
|
||||
robot, so this class does not perform it.
|
||||
"""
|
||||
|
||||
config_class = LeKiwiClientConfig
|
||||
name = "lekiwi_client"
|
||||
|
||||
def __init__(self, config: LeKiwiClientConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`LeKiwiClientConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
import zmq
|
||||
|
||||
self._zmq = zmq
|
||||
@@ -118,50 +105,24 @@ class LeKiwiClient(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._state_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._state_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self._is_connected
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
pass
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self) -> None:
|
||||
"""Open the ZMQ command and observation sockets to the LeKiwi host.
|
||||
"""Establishes ZMQ sockets with the remote mobile robot"""
|
||||
|
||||
Takes no `calibrate` argument: calibration lives on the robot, not on the client.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the client is already connected.
|
||||
"""
|
||||
zmq = self._zmq
|
||||
self.zmq_context = zmq.Context()
|
||||
self.zmq_cmd_socket = self.zmq_context.socket(zmq.PUSH)
|
||||
@@ -185,10 +146,6 @@ class LeKiwiClient(Robot):
|
||||
self._is_connected = True
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _poll_and_get_latest_message(self) -> list[bytes] | None:
|
||||
@@ -246,6 +203,7 @@ class LeKiwiClient(Robot):
|
||||
self, observation: RobotObservation
|
||||
) -> tuple[dict[str, np.ndarray], RobotObservation]:
|
||||
"""Extracts frames, and state from the parsed observation."""
|
||||
|
||||
flat_state = {key: observation.get(key, 0.0) for key in self._state_order}
|
||||
|
||||
state_vec = np.array([flat_state[key] for key in self._state_order], dtype=np.float32)
|
||||
@@ -264,12 +222,14 @@ class LeKiwiClient(Robot):
|
||||
return current_frames, obs_dict
|
||||
|
||||
def _get_data(self) -> tuple[dict[str, np.ndarray], RobotObservation]:
|
||||
"""Polls the video socket for the latest observation data.
|
||||
"""
|
||||
Polls the video socket for the latest observation data.
|
||||
|
||||
Attempts to retrieve and decode the latest message within a short timeout.
|
||||
If successful, updates and returns the new frames, speed, and arm state.
|
||||
If no new data arrives or decoding fails, returns the last known values.
|
||||
"""
|
||||
|
||||
# 1. Get the latest message's frames from the socket
|
||||
latest_frames = self._poll_and_get_latest_message()
|
||||
|
||||
@@ -298,17 +258,12 @@ class LeKiwiClient(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Receive one observation from the remote robot over ZMQ.
|
||||
|
||||
Wheel speeds arrive as raw motor velocities and are converted here to body-frame `x`, `y` and
|
||||
`theta`.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Follower arm positions, body-frame base velocities and camera frames.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the client is not connected.
|
||||
"""
|
||||
Capture observations from the remote robot: current follower arm positions,
|
||||
present wheel speeds (converted to body-frame velocities: x, y, theta),
|
||||
and a camera frame. Receives over ZMQ, translate to body-frame vel
|
||||
"""
|
||||
|
||||
frames, obs_dict = self._get_data()
|
||||
|
||||
# Loop over each configured camera
|
||||
@@ -353,24 +308,21 @@ class LeKiwiClient(Robot):
|
||||
}
|
||||
|
||||
def configure(self):
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
pass
|
||||
|
||||
@check_if_not_connected
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Send a target configuration to the remote robot over ZMQ.
|
||||
|
||||
Body-frame base velocities are translated into wheel velocities before sending.
|
||||
"""Command lekiwi to move to a target joint configuration. Translates to motor space + sends over ZMQ
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`): Goal positions for the arm and body-frame velocities for the base.
|
||||
|
||||
action (RobotAction): array containing the goal positions for the motors.
|
||||
Raises:
|
||||
RobotDeviceNotConnectedError: if robot is not connected.
|
||||
|
||||
Returns:
|
||||
np.ndarray: the action sent to the motors, potentially clipped.
|
||||
"""
|
||||
|
||||
# Action values may be torch tensors (e.g. replayed from a dataset) or numpy
|
||||
# scalars; json.dumps only serializes Python primitives, so coerce each value to a
|
||||
# plain float before sending.
|
||||
@@ -386,11 +338,8 @@ class LeKiwiClient(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Close the ZMQ sockets and terminate the context.
|
||||
"""Cleans ZMQ comms"""
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the client is not connected.
|
||||
"""
|
||||
self.zmq_observation_socket.close()
|
||||
self.zmq_cmd_socket.close()
|
||||
self.zmq_context.term()
|
||||
|
||||
@@ -36,19 +36,7 @@ class LeKiwiServerConfig:
|
||||
|
||||
|
||||
class LeKiwiHost:
|
||||
"""Serves a [`~robots.lekiwi.LeKiwi`] over ZMQ so a client can drive it from another machine.
|
||||
|
||||
Runs on the robot's own computer, receiving actions on one socket and publishing observations on
|
||||
another.
|
||||
"""
|
||||
|
||||
def __init__(self, config: LeKiwiHostConfig):
|
||||
"""Bind the command and observation sockets.
|
||||
|
||||
Args:
|
||||
config (`LeKiwiHostConfig`):
|
||||
Ports, loop frequency and watchdog settings for the host.
|
||||
"""
|
||||
self.zmq_context = zmq.Context()
|
||||
self.zmq_cmd_socket = self.zmq_context.socket(zmq.PULL)
|
||||
self.zmq_cmd_socket.setsockopt(zmq.CONFLATE, 1)
|
||||
@@ -65,11 +53,6 @@ class LeKiwiHost:
|
||||
self.max_loop_freq_hz = config.max_loop_freq_hz
|
||||
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.zmq_observation_socket.close()
|
||||
self.zmq_cmd_socket.close()
|
||||
self.zmq_context.term()
|
||||
@@ -77,12 +60,6 @@ class LeKiwiHost:
|
||||
|
||||
@draccus.wrap()
|
||||
def main(cfg: LeKiwiServerConfig):
|
||||
"""Run the LeKiwi host loop until the configured connection time elapses.
|
||||
|
||||
Args:
|
||||
cfg (`LeKiwiServerConfig`):
|
||||
The robot and host configuration to serve.
|
||||
"""
|
||||
logging.info("Configuring LeKiwi")
|
||||
robot = LeKiwi(cfg.robot)
|
||||
|
||||
|
||||
@@ -22,30 +22,6 @@ from ..config import RobotConfig
|
||||
@RobotConfig.register_subclass("omx_follower")
|
||||
@dataclass
|
||||
class OmxFollowerConfig(RobotConfig):
|
||||
"""Configuration for the OpenMANIPULATOR-X follower arm.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to, e.g. `/dev/ttyUSB0`. Run `lerobot-find-port` to identify
|
||||
it.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions, keyed by the name they appear under in
|
||||
observations. Each must specify `width`, `height` and `fps`.
|
||||
use_degrees (`bool`, *optional*, defaults to `False`):
|
||||
Whether to report and accept joint positions in degrees rather than as a normalised range.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
# Port to connect to the arm
|
||||
port: str
|
||||
|
||||
|
||||
@@ -36,21 +36,15 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OmxFollower(Robot):
|
||||
"""The [OpenMANIPULATOR-X](https://github.com/ROBOTIS-GIT/open_manipulator) follower arm.
|
||||
|
||||
Developed by Woojin Wie and Junha Cha at [ROBOTIS](https://ai.robotis.com/).
|
||||
"""
|
||||
- [OMX](https://github.com/ROBOTIS-GIT/open_manipulator),
|
||||
expansion, developed by Woojin Wie and Junha Cha from [ROBOTIS](https://ai.robotis.com/)
|
||||
"""
|
||||
|
||||
config_class = OmxFollowerConfig
|
||||
name = "omx_follower"
|
||||
|
||||
def __init__(self, config: OmxFollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`OmxFollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
norm_mode_body = MotorNormMode.DEGREES if config.use_degrees else MotorNormMode.RANGE_M100_100
|
||||
@@ -85,50 +79,25 @@ class OmxFollower(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus.is_connected and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect the motor bus and cameras, handling the pre-calibrated case.
|
||||
|
||||
OMX arms ship calibrated, so this avoids asking for a manual calibration where possible:
|
||||
|
||||
- if the packaged default calibration does not match the motors, the motors' own values are read
|
||||
and saved;
|
||||
- if no calibration file exists, factory defaults are used (`homing_offset=0`, `range_min=0`,
|
||||
`range_max=4095`).
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to calibrate if the arm is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
For OMX robots that come pre-calibrated:
|
||||
- If default calibration from package doesn't match motors, read from motors and save
|
||||
- This allows using pre-calibrated robots without manual calibration
|
||||
- If no calibration file exists, use factory default values (homing_offset=0, range_min=0, range_max=4095)
|
||||
"""
|
||||
|
||||
self.bus.connect()
|
||||
if not self.is_calibrated and calibrate:
|
||||
logger.info(
|
||||
@@ -144,18 +113,9 @@ class OmxFollower(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
self.bus.disable_torque()
|
||||
logger.info(f"\nUsing factory default calibration values for {self}")
|
||||
logger.info(f"\nWriting default configuration of {self} to the motors")
|
||||
@@ -180,7 +140,6 @@ class OmxFollower(Robot):
|
||||
logger.info(f"Calibration saved to {self.calibration_fpath}")
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
with self.bus.torque_disabled():
|
||||
self.bus.configure_motors()
|
||||
# Use 'extended position mode' for all motors except gripper, because in joint mode the servos
|
||||
@@ -205,11 +164,6 @@ class OmxFollower(Robot):
|
||||
self.bus.write("Position_D_Gain", "elbow_flex", 600)
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
for motor in reversed(self.bus.motors):
|
||||
input(f"Connect the controller board to the '{motor}' motor only and press enter.")
|
||||
self.bus.setup_motor(motor)
|
||||
@@ -218,14 +172,6 @@ class OmxFollower(Robot):
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
# Read arm position
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
obs_dict = self.bus.sync_read("Present_Position")
|
||||
obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
|
||||
@@ -262,6 +208,7 @@ class OmxFollower(Robot):
|
||||
Returns:
|
||||
RobotAction: The action sent to the motors, potentially clipped.
|
||||
"""
|
||||
|
||||
goal_pos = {key.removesuffix(".pos"): val for key, val in action.items() if key.endswith(".pos")}
|
||||
|
||||
# Cap goal position when too far away from present position.
|
||||
@@ -277,11 +224,6 @@ class OmxFollower(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
for cam in self.cameras.values():
|
||||
cam.disconnect()
|
||||
|
||||
@@ -45,13 +45,7 @@ RIGHT_DEFAULT_JOINTS_LIMITS: dict[str, tuple[float, float]] = {
|
||||
|
||||
@dataclass
|
||||
class OpenArmFollowerConfigBase:
|
||||
"""Field definitions for the OpenArm follower, a 7-DOF arm plus gripper on Damiao CAN motors.
|
||||
|
||||
This class only carries the fields. The registered configuration users instantiate is
|
||||
[`OpenArmFollowerConfig`], which documents them all in one place — doc-builder renders only a class's
|
||||
own docstring, never its bases'. It is also used directly as the per-arm config of
|
||||
[`~robots.bi_openarm_follower.BiOpenArmFollowerConfig`].
|
||||
"""
|
||||
"""Base configuration for the OpenArms follower robot with Damiao motors."""
|
||||
|
||||
# CAN interfaces - one per arm
|
||||
# arm CAN interface (e.g., "can1")
|
||||
@@ -129,57 +123,4 @@ class OpenArmFollowerConfigBase:
|
||||
@RobotConfig.register_subclass("openarm_follower")
|
||||
@dataclass
|
||||
class OpenArmFollowerConfig(RobotConfig, OpenArmFollowerConfigBase):
|
||||
"""Configuration for a single OpenArm follower arm.
|
||||
|
||||
OpenArm is a 7-DOF arm plus gripper on Damiao CAN motors, so `port` names a CAN interface rather than a
|
||||
serial device. Calibration follows the usual LeRobot flow and is stored per `id`.
|
||||
|
||||
> [!WARNING]
|
||||
> `joint_limits` defaults to a deliberately tiny range so an uncalibrated arm cannot swing. Set `side`
|
||||
> to `"left"` or `"right"` to get the real limits for that arm, or pass your own.
|
||||
|
||||
The per-joint lists — `position_kp`, `position_kd` — hold 8 values in motor order: `joint_1` through
|
||||
`joint_7`, then `gripper`.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
CAN interface the arm is on, e.g. `"can0"` on Linux.
|
||||
side (`str`, *optional*):
|
||||
Which arm this is, `"left"` or `"right"`. Selects that side's joint limits. Leaving it `None`
|
||||
keeps the small safety defaults.
|
||||
can_interface (`str`, *optional*, defaults to `"socketcan"`):
|
||||
CAN backend: `"socketcan"` on Linux, `"slcan"` for a serial adapter, or `"auto"` to detect.
|
||||
use_can_fd (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use CAN FD. OpenArm uses it by default.
|
||||
can_bitrate (`int`, *optional*, defaults to 1000000):
|
||||
Nominal CAN bitrate, 1 Mbps.
|
||||
can_data_bitrate (`int`, *optional*, defaults to 5000000):
|
||||
CAN FD data bitrate, 5 Mbps. Only used when `use_can_fd` is `True`.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
use_velocity_and_torque (`bool`, *optional*, defaults to `False`):
|
||||
Whether to expose `.vel` and `.torque` per motor in the observation features. Kept `False` by
|
||||
default for compatibility with the position-only `openarm_mini` teleoperator.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions.
|
||||
motor_config (`dict[str, tuple[int, int, str]]`, *optional*):
|
||||
Maps motor name to `(send_can_id, recv_can_id, motor_type)`. Defaults to the stock OpenArm
|
||||
layout; change it only if you have rewired or re-addressed the motors.
|
||||
position_kp (`list[float]`, *optional*):
|
||||
MIT-mode proportional gains used by `send_action`, 8 values in motor order.
|
||||
position_kd (`list[float]`, *optional*):
|
||||
MIT-mode derivative gains used by `send_action`, 8 values in motor order.
|
||||
joint_limits (`dict[str, tuple[float, float]]`, *optional*):
|
||||
Soft `(min, max)` limits in degrees per joint, clipped against on every action.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
@@ -37,22 +37,15 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OpenArmFollower(Robot):
|
||||
"""The OpenArm follower: a 7-DOF arm plus gripper on a CAN bus.
|
||||
|
||||
Uses Damiao motors in MIT control mode. See [`~robots.Robot`] for the contract every method here
|
||||
implements.
|
||||
"""
|
||||
OpenArms Follower Robot which uses CAN bus communication to control 7 DOF arm with a gripper.
|
||||
The arm uses Damiao motors in MIT control mode.
|
||||
"""
|
||||
|
||||
config_class = OpenArmFollowerConfig
|
||||
name = "openarm_follower"
|
||||
|
||||
def __init__(self, config: OpenArmFollowerConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`OpenArmFollowerConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
@@ -134,11 +127,13 @@ class OpenArmFollower(Robot):
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect to the robot and optionally calibrate.
|
||||
"""
|
||||
Connect to the robot and optionally calibrate.
|
||||
|
||||
We assume that at connection time, the arms are in a safe rest position,
|
||||
and torque can be safely disabled to run calibration if needed.
|
||||
"""
|
||||
|
||||
# Connect to CAN bus
|
||||
logger.info(f"Connecting arm on {self.config.port}...")
|
||||
self.bus.connect()
|
||||
@@ -165,7 +160,8 @@ class OpenArmFollower(Robot):
|
||||
return self.bus.is_calibrated
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Run calibration procedure for OpenArms robot.
|
||||
"""
|
||||
Run calibration procedure for OpenArms robot.
|
||||
|
||||
The calibration procedure:
|
||||
1. Disable torque
|
||||
@@ -221,18 +217,14 @@ class OpenArmFollower(Robot):
|
||||
self.bus.configure_motors()
|
||||
|
||||
def setup_motors(self) -> None:
|
||||
"""Assign each motor its bus ID, one at a time.
|
||||
|
||||
Run this once when building the robot. Interactive: prompts you to connect the controller board to a
|
||||
single motor at a time.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Motor ID configuration is typically done via manufacturer tools for CAN motors."
|
||||
)
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Get current observation from robot including position, velocity, and torque.
|
||||
"""
|
||||
Get current observation from robot including position, velocity, and torque.
|
||||
|
||||
Reads all motor states (pos/vel/torque) in one CAN refresh cycle
|
||||
instead of 3 separate reads.
|
||||
@@ -276,7 +268,8 @@ class OpenArmFollower(Robot):
|
||||
custom_kp: dict[str, float] | None = None,
|
||||
custom_kd: dict[str, float] | None = None,
|
||||
) -> RobotAction:
|
||||
"""Send action command to robot.
|
||||
"""
|
||||
Send action command to robot.
|
||||
|
||||
The action magnitude may be clipped based on safety limits.
|
||||
|
||||
@@ -288,6 +281,7 @@ class OpenArmFollower(Robot):
|
||||
Returns:
|
||||
The action actually sent (potentially clipped)
|
||||
"""
|
||||
|
||||
goal_pos = {key.removesuffix(".pos"): val for key, val in action.items() if key.endswith(".pos")}
|
||||
|
||||
# Apply joint limit clipping to arm
|
||||
@@ -349,6 +343,7 @@ class OpenArmFollower(Robot):
|
||||
@check_if_not_connected
|
||||
def disconnect(self):
|
||||
"""Disconnect from robot."""
|
||||
|
||||
# Disconnect CAN bus
|
||||
self.bus.disconnect(self.config.disable_torque_on_disconnect)
|
||||
|
||||
|
||||
@@ -23,58 +23,6 @@ from ..config import RobotConfig
|
||||
@RobotConfig.register_subclass("reachy2")
|
||||
@dataclass
|
||||
class Reachy2RobotConfig(RobotConfig):
|
||||
"""Configuration for the Reachy 2 humanoid.
|
||||
|
||||
Reachy 2 is driven over the network rather than a serial bus, so `port` is a TCP port on the robot's
|
||||
gRPC service rather than a device path. Calibration is handled by the robot itself and there is no
|
||||
LeRobot calibration file.
|
||||
|
||||
Which joints appear in observations and actions is selected by the `with_*` flags: turning a part off
|
||||
removes its joints entirely. At least one part must stay enabled.
|
||||
|
||||
Args:
|
||||
max_relative_target (`float`, *optional*):
|
||||
Caps how far a single action may move a joint from its present position, as a safety limit.
|
||||
`None` disables clipping.
|
||||
ip_address (`str`, *optional*, defaults to `"localhost"`):
|
||||
Address of the Reachy 2 robot.
|
||||
port (`int`, *optional*, defaults to 50065):
|
||||
TCP port of the robot's service. Not a serial port.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `False`):
|
||||
Whether to call `turn_off_smoothly()` before disconnecting.
|
||||
use_external_commands (`bool`, *optional*, defaults to `False`):
|
||||
Set `True` when another system drives the robot, such as the official
|
||||
[teleoperation app](https://github.com/pollen-robotics/Reachy2Teleoperation). In that mode
|
||||
[`~robots.Robot.send_action`] does not send anything to the robot.
|
||||
with_mobile_base (`bool`, *optional*, defaults to `True`):
|
||||
Whether to include the mobile base's joints.
|
||||
with_l_arm (`bool`, *optional*, defaults to `True`):
|
||||
Whether to include the left arm's joints.
|
||||
with_r_arm (`bool`, *optional*, defaults to `True`):
|
||||
Whether to include the right arm's joints.
|
||||
with_neck (`bool`, *optional*, defaults to `True`):
|
||||
Whether to include the neck's joints.
|
||||
with_antennas (`bool`, *optional*, defaults to `True`):
|
||||
Whether to include the antennas' joints.
|
||||
with_left_teleop_camera (`bool`, *optional*, defaults to `False`):
|
||||
Whether to add the left teleoperation camera to observations.
|
||||
with_right_teleop_camera (`bool`, *optional*, defaults to `False`):
|
||||
Whether to add the right teleoperation camera to observations.
|
||||
with_torso_camera (`bool`, *optional*, defaults to `False`):
|
||||
Whether to add the torso RGB camera to observations.
|
||||
camera_width (`int`, *optional*, defaults to 640):
|
||||
Frame width for the built-in cameras. Their frame rate is fixed at 30 and is not configurable.
|
||||
camera_height (`int`, *optional*, defaults to 480):
|
||||
Frame height for the built-in cameras.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Additional cameras beyond the three built-in ones. The `with_*_camera` flags populate this
|
||||
field, so anything set here is merged with them.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular robot.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Unused: Reachy 2 manages its own calibration.
|
||||
"""
|
||||
|
||||
# `max_relative_target` limits the magnitude of the relative positional target vector for safety purposes.
|
||||
# Set this to a positive scalar to have the same value for all motors.
|
||||
max_relative_target: float | None = None
|
||||
@@ -117,11 +65,6 @@ class Reachy2RobotConfig(RobotConfig):
|
||||
cameras: dict[str, CameraConfig] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Add the built-in cameras selected by the `with_*_camera` flags and validate the part selection.
|
||||
|
||||
Raises:
|
||||
ValueError: If every robot part is disabled, which would leave no joints to control.
|
||||
"""
|
||||
# Add cameras with same ip_address as the robot
|
||||
if self.with_left_teleop_camera:
|
||||
self.cameras["teleop_left"] = Reachy2CameraConfig(
|
||||
|
||||
@@ -73,18 +73,14 @@ REACHY2_VEL = {
|
||||
|
||||
|
||||
class Reachy2Robot(Robot):
|
||||
"""[Reachy 2](https://www.pollen-robotics.com/reachy/), the humanoid by Pollen Robotics."""
|
||||
"""
|
||||
[Reachy 2](https://www.pollen-robotics.com/reachy/), by Pollen Robotics.
|
||||
"""
|
||||
|
||||
config_class = Reachy2RobotConfig
|
||||
name = "reachy2"
|
||||
|
||||
def __init__(self, config: Reachy2RobotConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`Reachy2RobotConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
require_package("reachy2_sdk", extra="reachy2")
|
||||
super().__init__(config)
|
||||
|
||||
@@ -101,41 +97,18 @@ class Reachy2Robot(Robot):
|
||||
|
||||
@property
|
||||
def observation_features(self) -> dict[str, Any]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self.motors_features, **self.camera_features}
|
||||
|
||||
@property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self.motors_features
|
||||
|
||||
@property
|
||||
def camera_features(self) -> dict[str, tuple[int | None, int | None, int]]:
|
||||
"""The shape of each configured camera's frames.
|
||||
|
||||
Returns:
|
||||
`dict[str, tuple[int | None, int | None, int]]`: Camera name mapped to
|
||||
`(height, width, channels)`.
|
||||
"""
|
||||
return {cam: (self.cameras[cam].height, self.cameras[cam].width, 3) for cam in self.cameras}
|
||||
|
||||
@property
|
||||
def motors_features(self) -> dict[str, type]:
|
||||
"""The joints this robot exposes, given which parts are enabled in the config.
|
||||
|
||||
Returns:
|
||||
`dict[str, type]`: Joint name mapped to `float`, including the mobile base's velocity
|
||||
components when `with_mobile_base` is set.
|
||||
"""
|
||||
if self.config.with_mobile_base:
|
||||
return {
|
||||
**dict.fromkeys(
|
||||
@@ -152,23 +125,9 @@ class Reachy2Robot(Robot):
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.reachy.is_connected() if self.reachy is not None else False
|
||||
|
||||
def connect(self, calibrate: bool = False) -> None:
|
||||
"""Connect to the robot and its cameras, then apply the configured settings.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `False`):
|
||||
Accepted for interface compatibility and ignored: Reachy 2 manages its own calibration.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
self.reachy = ReachySDK(self.config.ip_address)
|
||||
if not self.is_connected:
|
||||
raise ConnectionError()
|
||||
@@ -179,25 +138,15 @@ class Reachy2Robot(Robot):
|
||||
self.configure()
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
if self.reachy is not None:
|
||||
self.reachy.turn_on()
|
||||
self.reachy.reset_default_limits()
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return True
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _generate_joints_dict(self) -> dict[str, str]:
|
||||
@@ -223,14 +172,6 @@ class Reachy2Robot(Robot):
|
||||
return {}
|
||||
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
obs_dict: RobotObservation = {}
|
||||
|
||||
# Read Reachy 2 state
|
||||
@@ -245,18 +186,6 @@ class Reachy2Robot(Robot):
|
||||
return obs_dict
|
||||
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Command the robot to move towards a target configuration.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
Target values, keyed as in [`~robots.Robot.action_features`].
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent, which may be clipped by `max_relative_target`.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
if self.reachy is not None:
|
||||
if not self.is_connected:
|
||||
raise ConnectionError()
|
||||
@@ -299,11 +228,6 @@ class Reachy2Robot(Robot):
|
||||
return action
|
||||
|
||||
def disconnect(self) -> None:
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
if self.reachy is not None:
|
||||
for cam in self.cameras.values():
|
||||
cam.disconnect()
|
||||
|
||||
@@ -23,12 +23,10 @@ from ..config import RobotConfig
|
||||
|
||||
@dataclass
|
||||
class RebotB601FollowerConfig:
|
||||
"""Field definitions for the Seeed Studio reBot B601-DM follower arm.
|
||||
"""Base configuration class for the Seeed Studio reBot B601-DM follower arm.
|
||||
|
||||
This class only carries the fields. The registered configuration users instantiate is
|
||||
[`RebotB601FollowerRobotConfig`], which documents them all in one place — doc-builder renders only a
|
||||
class's own docstring, never its bases'. It is also used directly as the per-arm config of
|
||||
[`~robots.bi_rebot_b601_follower.BiRebotB601FollowerConfig`].
|
||||
The B601-DM is a 6-DOF arm plus gripper driven by Damiao CAN motors. Motor
|
||||
communication goes through the ``motorbridge`` package.
|
||||
"""
|
||||
|
||||
# Communication port. For ``can_adapter="damiao"`` this is the Damiao serial
|
||||
@@ -106,62 +104,6 @@ class RebotB601FollowerConfig:
|
||||
@RobotConfig.register_subclass("rebot_b601_follower")
|
||||
@dataclass
|
||||
class RebotB601FollowerRobotConfig(RobotConfig, RebotB601FollowerConfig):
|
||||
"""Configuration for the Seeed Studio reBot B601-DM follower arm.
|
||||
|
||||
The B601-DM is a 6-DOF arm plus gripper on Damiao CAN motors, driven through the `motorbridge`
|
||||
package. What `port` means depends on `can_adapter`. Calibration follows the usual LeRobot flow and is
|
||||
stored per `id`.
|
||||
|
||||
The arm and the gripper are controlled separately: `control_mode` governs the six arm joints and
|
||||
`gripper_control_mode` the gripper, and each mode uses a different subset of the gain fields.
|
||||
|
||||
Per-joint lists hold 7 values in motor order: `shoulder_pan`, `shoulder_lift`, `elbow_flex`,
|
||||
`wrist_flex`, `wrist_yaw`, `wrist_roll`, `gripper`.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Where the arm is reached. For `can_adapter="damiao"` this is the serial bridge device, e.g.
|
||||
`/dev/ttyACM0`; for `can_adapter="socketcan"` it is the CAN channel name, e.g. `can0`.
|
||||
can_adapter (`str`, *optional*, defaults to `"damiao"`):
|
||||
`"damiao"` for the dedicated Damiao serial bridge, or `"socketcan"` for SocketCAN adapters
|
||||
such as PCAN, slcan and embedded controllers.
|
||||
dm_serial_baud (`int`, *optional*, defaults to 921600):
|
||||
Baud rate of the Damiao serial bridge. Only used when `can_adapter="damiao"`.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, in degrees. A scalar
|
||||
applies to every motor; a dict maps motor name to a per-motor cap. `None` disables clipping.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions.
|
||||
motor_can_ids (`dict[str, tuple[int, int]]`, *optional*):
|
||||
Maps motor name to its `(send_can_id, recv_can_id)` pair. Change it only if you have
|
||||
re-addressed the motors.
|
||||
pos_vel_velocity (`float | list[float]`, *optional*):
|
||||
Maximum speed in deg/s per joint, used by the arm in `pos_vel` mode and by the gripper in
|
||||
`force_pos` mode.
|
||||
control_mode (`str`, *optional*, defaults to `"mit"`):
|
||||
How the six arm joints are driven: `"mit"` or `"pos_vel"`.
|
||||
mit_kp (`float | list[float]`, *optional*):
|
||||
MIT-mode proportional gains per arm joint. Unused when `control_mode="pos_vel"`.
|
||||
mit_kd (`float | list[float]`, *optional*):
|
||||
MIT-mode derivative gains per arm joint. Unused when `control_mode="pos_vel"`.
|
||||
gripper_control_mode (`str`, *optional*, defaults to `"force_pos"`):
|
||||
How the gripper is driven: `"force_pos"` or `"mit"`.
|
||||
gripper_torque_ratio (`float`, *optional*, defaults to 0.07):
|
||||
Maximum grip force as a fraction in `[0, 1]`. Only used when
|
||||
`gripper_control_mode="force_pos"`.
|
||||
gripper_mit_kp (`float`, *optional*, defaults to 8.0):
|
||||
Gripper MIT-mode proportional gain. Only used when `gripper_control_mode="mit"`.
|
||||
gripper_mit_kd (`float`, *optional*, defaults to 0.3):
|
||||
Gripper MIT-mode derivative gain. Only used when `gripper_control_mode="mit"`.
|
||||
joint_limits (`dict[str, tuple[float, float]]`, *optional*):
|
||||
Soft `(min, max)` limits in degrees per joint, clipped against on every action.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
"""
|
||||
"""Registered configuration for the reBot B601-DM follower robot."""
|
||||
|
||||
pass
|
||||
|
||||
@@ -66,12 +66,6 @@ class RebotB601Follower(Robot):
|
||||
name = "rebot_b601_follower"
|
||||
|
||||
def __init__(self, config: RebotB601FollowerRobotConfig):
|
||||
"""Build the robot from its configuration.
|
||||
|
||||
Args:
|
||||
config (`RebotB601FollowerRobotConfig`):
|
||||
The robot's configuration. Its `port` and `cameras` determine what is connected.
|
||||
"""
|
||||
require_package("motorbridge", extra="rebot")
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
@@ -97,43 +91,18 @@ class RebotB601Follower(Robot):
|
||||
|
||||
@cached_property
|
||||
def observation_features(self) -> dict[str, type | tuple]:
|
||||
"""The values this robot reports, and their types or shapes.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys as returned by [`~robots.Robot.get_observation`], mapped to a scalar type for
|
||||
proprioceptive values or to a `(height, width, channels)` shape for images.
|
||||
"""
|
||||
return {**self._motors_ft, **self._cameras_ft}
|
||||
|
||||
@cached_property
|
||||
def action_features(self) -> dict[str, type]:
|
||||
"""The values this robot accepts, and their types.
|
||||
|
||||
Returns:
|
||||
`dict`: Keys accepted by [`~robots.Robot.send_action`], mapped to their type.
|
||||
"""
|
||||
return self._motors_ft
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether every device this robot uses is connected.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` only when the robot and all its cameras are connected.
|
||||
"""
|
||||
return self.bus is not None and all(cam.is_connected for cam in self.cameras.values())
|
||||
|
||||
@check_if_already_connected
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Connect to the robot and its cameras, then apply the configured settings.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to run calibration if the robot is not already calibrated.
|
||||
|
||||
Raises:
|
||||
DeviceAlreadyConnectedError: If the robot is already connected.
|
||||
"""
|
||||
logger.info(f"Connecting {self} on {self.config.port} (adapter={self.config.can_adapter})...")
|
||||
if self.config.can_adapter == "damiao":
|
||||
self.bus = MotorBridgeController.from_dm_serial(
|
||||
@@ -164,18 +133,9 @@ class RebotB601Follower(Robot):
|
||||
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` when no calibration is needed before use.
|
||||
"""
|
||||
return bool(self.calibration)
|
||||
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot and store the result.
|
||||
|
||||
Interactive: prompts on stdin and asks you to move the robot through the required positions.
|
||||
"""
|
||||
if self.calibration:
|
||||
user_input = input(
|
||||
f"Press ENTER to use provided calibration file associated with the id {self.id}, "
|
||||
@@ -214,7 +174,6 @@ class RebotB601Follower(Robot):
|
||||
print(f"Calibration saved to {self.calibration_fpath}")
|
||||
|
||||
def configure(self) -> None:
|
||||
"""Apply the operating mode, gains and limits from the configuration to the robot."""
|
||||
if self.config.control_mode not in ("pos_vel", "mit"):
|
||||
raise ValueError(
|
||||
f"Unsupported control_mode '{self.config.control_mode}'. Use 'pos_vel' or 'mit'."
|
||||
@@ -267,14 +226,6 @@ class RebotB601Follower(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Read the robot's current state and a frame from each camera.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: Keys matching [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
obs_dict = {f"{motor}.pos": pos for motor, pos in self._present_pos().items()}
|
||||
dt_ms = (time.perf_counter() - start) * 1e3
|
||||
@@ -360,11 +311,6 @@ class RebotB601Follower(Robot):
|
||||
|
||||
@check_if_not_connected
|
||||
def disconnect(self) -> None:
|
||||
"""Disconnect from the robot and its cameras.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If the robot is not connected.
|
||||
"""
|
||||
for motor in self.motors.values():
|
||||
if self.config.disable_torque_on_disconnect:
|
||||
motor.disable()
|
||||
|
||||
+63
-95
@@ -28,22 +28,15 @@ from .config import RobotConfig
|
||||
# TODO(aliberts): action/obs typing such as Generic[ObsType, ActType] similar to gym.Env ?
|
||||
# https://github.com/Farama-Foundation/Gymnasium/blob/3287c869f9a48d99454306b0d4b4ec537f0f35e3/gymnasium/core.py#L23
|
||||
class Robot(abc.ABC):
|
||||
"""The base abstract class for all LeRobot-compatible robots.
|
||||
"""
|
||||
The base abstract class for all LeRobot-compatible robots.
|
||||
|
||||
This class provides a standardized interface for interacting with physical robots. Subclasses must
|
||||
implement all abstract methods and properties to be usable.
|
||||
This class provides a standardized interface for interacting with physical robots.
|
||||
Subclasses must implement all abstract methods and properties to be usable.
|
||||
|
||||
Used as a context manager, a robot connects on entry and disconnects on exit even if the body raises:
|
||||
|
||||
```python
|
||||
>>> with SO101Follower(config) as robot: # doctest: +SKIP
|
||||
... obs = robot.get_observation()
|
||||
... robot.send_action(action)
|
||||
```
|
||||
|
||||
**Attributes**:
|
||||
- **config_class** (`type[RobotConfig]`) -- The expected configuration class for this robot.
|
||||
- **name** (`str`) -- The unique robot name used to identify this robot type.
|
||||
Attributes:
|
||||
config_class (RobotConfig): The expected configuration class for this robot.
|
||||
name (str): The unique robot name used to identify this robot type.
|
||||
"""
|
||||
|
||||
# Set these in ALL subclasses
|
||||
@@ -51,13 +44,6 @@ class Robot(abc.ABC):
|
||||
name: str
|
||||
|
||||
def __init__(self, config: RobotConfig):
|
||||
"""Set up identity and calibration paths, loading an existing calibration file if there is one.
|
||||
|
||||
Args:
|
||||
config (`RobotConfig`):
|
||||
The robot's configuration. Its `id` and `calibration_dir` decide where calibration is
|
||||
read from and written to.
|
||||
"""
|
||||
self.robot_type = self.name
|
||||
self.id = config.id
|
||||
self.calibration_dir = (
|
||||
@@ -70,24 +56,28 @@ class Robot(abc.ABC):
|
||||
self._load_calibration()
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""Return this robot's id and class name, e.g. `"my_arm SO101Follower"`.
|
||||
|
||||
Returns:
|
||||
`str`: A short identifier used in log messages.
|
||||
"""
|
||||
return f"{self.id} {self.__class__.__name__}"
|
||||
|
||||
def __enter__(self):
|
||||
"""Context manager entry. Automatically connects to the robot."""
|
||||
"""
|
||||
Context manager entry.
|
||||
Automatically connects to the camera.
|
||||
"""
|
||||
self.connect()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback) -> None:
|
||||
"""Context manager exit. Disconnects, ensuring resources are released even on error."""
|
||||
"""
|
||||
Context manager exit.
|
||||
Automatically disconnects, ensuring resources are released even on error.
|
||||
"""
|
||||
self.disconnect()
|
||||
|
||||
def __del__(self) -> None:
|
||||
"""Destructor safety net. Disconnects if the object is garbage collected without cleanup."""
|
||||
"""
|
||||
Destructor safety net.
|
||||
Attempts to disconnect if the object is garbage collected without cleanup.
|
||||
"""
|
||||
try:
|
||||
if self.is_connected:
|
||||
self.disconnect()
|
||||
@@ -98,102 +88,83 @@ class Robot(abc.ABC):
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def observation_features(self) -> dict:
|
||||
"""A dictionary describing the structure and types of the observations produced by the robot.
|
||||
"""
|
||||
A dictionary describing the structure and types of the observations produced by the robot.
|
||||
Its structure (keys) should match the structure of what is returned by :pymeth:`get_observation`.
|
||||
Values for the dict should either be:
|
||||
- The type of the value if it's a simple value, e.g. `float` for single proprioceptive value (a joint's position/velocity)
|
||||
- A tuple representing the shape if it's an array-type value, e.g. `(height, width, channel)` for images
|
||||
|
||||
Its keys should match the structure of what is returned by [`~robots.Robot.get_observation`]. Values
|
||||
should either be:
|
||||
|
||||
- the type of the value if it's a simple value, e.g. `float` for a single proprioceptive value
|
||||
(a joint's position or velocity)
|
||||
- a tuple representing the shape if it's an array-type value, e.g. `(height, width, channel)` for
|
||||
images
|
||||
|
||||
> [!NOTE]
|
||||
> This property must be callable regardless of whether the robot is connected.
|
||||
|
||||
Returns:
|
||||
`dict`: Observation names mapped to their type or shape.
|
||||
Note: this property should be able to be called regardless of whether the robot is connected or not.
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def action_features(self) -> dict:
|
||||
"""A dictionary describing the structure and types of the actions expected by the robot.
|
||||
"""
|
||||
A dictionary describing the structure and types of the actions expected by the robot. Its structure
|
||||
(keys) should match the structure of what is passed to :pymeth:`send_action`. Values for the dict
|
||||
should be the type of the value if it's a simple value, e.g. `float` for single proprioceptive value
|
||||
(a joint's goal position/velocity)
|
||||
|
||||
Its keys should match the structure of what is passed to [`~robots.Robot.send_action`]. Values should
|
||||
be the type of the value if it's a simple value, e.g. `float` for a single proprioceptive value
|
||||
(a joint's goal position or velocity).
|
||||
|
||||
> [!NOTE]
|
||||
> This property must be callable regardless of whether the robot is connected.
|
||||
|
||||
Returns:
|
||||
`dict`: Action names mapped to their type or shape.
|
||||
Note: this property should be able to be called regardless of whether the robot is connected or not.
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def is_connected(self) -> bool:
|
||||
"""Whether the robot is currently connected.
|
||||
|
||||
If `False`, calling [`~robots.Robot.get_observation`] or [`~robots.Robot.send_action`] should raise
|
||||
an error.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` if communication with the robot is established.
|
||||
"""
|
||||
Whether the robot is currently connected or not. If `False`, calling :pymeth:`get_observation` or
|
||||
:pymeth:`send_action` should raise an error.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
"""Establish communication with the robot.
|
||||
"""
|
||||
Establish communication with the robot.
|
||||
|
||||
Args:
|
||||
calibrate (`bool`, *optional*, defaults to `True`):
|
||||
Whether to automatically calibrate the robot after connecting, if it is not calibrated or
|
||||
needs recalibration. Whether calibration is needed is hardware-dependent.
|
||||
calibrate (bool): If True, automatically calibrate the robot after connecting if it's not
|
||||
calibrated or needs calibration (this is hardware-dependant).
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def is_calibrated(self) -> bool:
|
||||
"""Whether the robot is currently calibrated.
|
||||
|
||||
Returns:
|
||||
`bool`: `True` if the robot is calibrated. Always `True` for robots where calibration does not
|
||||
apply.
|
||||
"""
|
||||
"""Whether the robot is currently calibrated or not. Should be always `True` if not applicable"""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def calibrate(self) -> None:
|
||||
"""Calibrate the robot if applicable. If not, this should be a no-op.
|
||||
"""
|
||||
Calibrate the robot if applicable. If not, this should be a no-op.
|
||||
|
||||
This method should collect any necessary data (e.g. motor offsets) and update the `calibration`
|
||||
dictionary accordingly.
|
||||
This method should collect any necessary data (e.g., motor offsets) and update the
|
||||
:pyattr:`calibration` dictionary accordingly.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _load_calibration(self, fpath: Path | None = None) -> None:
|
||||
"""Helper to load calibration data from the specified file.
|
||||
"""
|
||||
Helper to load calibration data from the specified file.
|
||||
|
||||
Args:
|
||||
fpath (`Path`, *optional*):
|
||||
Path to the calibration file. Defaults to `self.calibration_fpath`.
|
||||
fpath (Path | None): Optional path to the calibration file. Defaults to `self.calibration_fpath`.
|
||||
"""
|
||||
fpath = self.calibration_fpath if fpath is None else fpath
|
||||
with open(fpath) as f, draccus.config_type("json"):
|
||||
self.calibration = draccus.load(dict[str, MotorCalibration], f)
|
||||
|
||||
def _save_calibration(self, fpath: Path | None = None) -> None:
|
||||
"""Helper to save calibration data to the specified file.
|
||||
"""
|
||||
Helper to save calibration data to the specified file.
|
||||
|
||||
Args:
|
||||
fpath (`Path`, *optional*):
|
||||
Path to save the calibration file to. Defaults to `self.calibration_fpath`.
|
||||
fpath (Path | None): Optional path to save the calibration file. Defaults to `self.calibration_fpath`.
|
||||
"""
|
||||
fpath = self.calibration_fpath if fpath is None else fpath
|
||||
with open(fpath, "w") as f, draccus.config_type("json"):
|
||||
@@ -201,39 +172,36 @@ class Robot(abc.ABC):
|
||||
|
||||
@abc.abstractmethod
|
||||
def configure(self) -> None:
|
||||
"""Apply any one-time or runtime configuration to the robot.
|
||||
|
||||
"""
|
||||
Apply any one-time or runtime configuration to the robot.
|
||||
This may include setting motor parameters, control modes, or initial state.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_observation(self) -> RobotObservation:
|
||||
"""Retrieve the current observation from the robot.
|
||||
"""
|
||||
Retrieve the current observation from the robot.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: A flat dictionary representing the robot's current sensory state. Its structure
|
||||
should match [`~robots.Robot.observation_features`].
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If [`~robots.Robot.connect`] has not been called.
|
||||
RobotObservation: A flat dictionary representing the robot's current sensory state. Its structure
|
||||
should match :pymeth:`observation_features`.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
"""Send an action command to the robot.
|
||||
"""
|
||||
Send an action command to the robot.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The desired action. Its structure should match [`~robots.Robot.action_features`].
|
||||
action (RobotAction): Dictionary representing the desired action. Its structure should match
|
||||
:pymeth:`action_features`.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The action actually sent to the motors, potentially clipped or modified, e.g.
|
||||
by safety limits on velocity. Prefer this over the requested action when logging or recording.
|
||||
|
||||
Raises:
|
||||
DeviceNotConnectedError: If [`~robots.Robot.connect`] has not been called.
|
||||
RobotAction: The action actually sent to the motors potentially clipped or modified, e.g. by
|
||||
safety limits on velocity.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
@@ -23,12 +23,7 @@ from ..config import RobotConfig
|
||||
|
||||
@dataclass
|
||||
class SOFollowerConfig:
|
||||
"""Field definitions shared by the SO-family follower arms.
|
||||
|
||||
This class only carries the fields. The registered configuration users instantiate is
|
||||
[`SOFollowerRobotConfig`], which combines these with [`~robots.RobotConfig`] and documents them all in
|
||||
one place — doc-builder renders only a class's own docstring, never its bases'.
|
||||
"""
|
||||
"""Base configuration class for SO Follower robots."""
|
||||
|
||||
# Port to connect to the arm
|
||||
port: str
|
||||
@@ -62,51 +57,6 @@ class SOFollowerConfig:
|
||||
@RobotConfig.register_subclass("so100_follower")
|
||||
@dataclass
|
||||
class SOFollowerRobotConfig(RobotConfig, SOFollowerConfig):
|
||||
"""Configuration for the SO-100 and SO-101 follower arms.
|
||||
|
||||
Both arms share this class; `SO100FollowerConfig` and `SO101FollowerConfig` are aliases for it. They
|
||||
differ in their calibration and gearing, not in their control code.
|
||||
|
||||
Args:
|
||||
port (`str`):
|
||||
Serial port the arm is connected to, e.g. `/dev/ttyACM0` on Linux or `COM3` on Windows. Run
|
||||
`lerobot-find-port` to identify it.
|
||||
disable_torque_on_disconnect (`bool`, *optional*, defaults to `True`):
|
||||
Whether to release the motors on disconnect. Leave `True` unless the arm is holding a load it
|
||||
must not drop.
|
||||
max_relative_target (`float | dict[str, float]`, *optional*):
|
||||
Caps how far a single action may move the arm from its present position, as a safety limit. A
|
||||
scalar applies to every motor; a dict maps motor name to a per-motor cap. `None` disables
|
||||
clipping. Enabling this costs an extra read of the present position on every step.
|
||||
cameras (`dict[str, CameraConfig]`, *optional*):
|
||||
Cameras to read alongside the arm's joint positions, keyed by the name they appear under in
|
||||
observations. Each must specify `width`, `height` and `fps`.
|
||||
use_degrees (`bool`, *optional*, defaults to `True`):
|
||||
Whether to report and accept joint positions in degrees. Keep `True` for compatibility with
|
||||
existing policies and datasets.
|
||||
position_p_coefficient (`int`, *optional*, defaults to 16):
|
||||
Proportional gain written to the Feetech STS3215 motors at connect time.
|
||||
position_i_coefficient (`int`, *optional*, defaults to 0):
|
||||
Integral gain written to the motors at connect time.
|
||||
position_d_coefficient (`int`, *optional*, defaults to 32):
|
||||
Derivative gain written to the motors at connect time.
|
||||
num_read_retries (`int`, *optional*, defaults to 2):
|
||||
Extra attempts when a `sync_read` fails. Feetech buses occasionally return a corrupted status
|
||||
packet, especially when several joints move at once, which would otherwise abort the control
|
||||
loop. Retries are immediate and only happen on failure, so steady-state read cost is unchanged.
|
||||
id (`str`, *optional*):
|
||||
Identifier for this particular arm; also names its calibration file.
|
||||
calibration_dir (`Path`, *optional*):
|
||||
Where to read and write the calibration file. Defaults to the LeRobot calibration home.
|
||||
|
||||
Example:
|
||||
```python
|
||||
>>> from lerobot.robots.so_follower import SO101Follower, SO101FollowerConfig
|
||||
>>> config = SO101FollowerConfig(port="/dev/ttyACM0", max_relative_target=5.0) # doctest: +SKIP
|
||||
>>> robot = SO101Follower(config) # doctest: +SKIP
|
||||
```
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
|
||||
@@ -40,7 +40,8 @@ logger = logging.getLogger(__name__)
|
||||
@ProcessorStepRegistry.register("ee_reference_and_delta")
|
||||
@dataclass
|
||||
class EEReferenceAndDelta(RobotActionProcessorStep):
|
||||
"""Computes a target end-effector pose from a relative delta command.
|
||||
"""
|
||||
Computes a target end-effector pose from a relative delta command.
|
||||
|
||||
This step takes a desired change in position and orientation (`target_*`) and applies it to a
|
||||
reference end-effector pose to calculate an absolute target pose. The reference pose is derived
|
||||
@@ -52,16 +53,15 @@ class EEReferenceAndDelta(RobotActionProcessorStep):
|
||||
2. `use_latched_reference=False`: The reference pose is updated to the robot's current pose at
|
||||
every step.
|
||||
|
||||
**Attributes**:
|
||||
- **kinematics** (`RobotKinematics`) -- The robot's kinematic model for forward kinematics.
|
||||
- **end_effector_step_sizes** (`dict`) -- A dictionary scaling the input delta commands.
|
||||
- **motor_names** (`list[str]`) -- A list of motor names required for forward kinematics.
|
||||
- **use_latched_reference** (`bool`) -- If True, latch the reference pose on enable; otherwise, always
|
||||
use the current pose as the reference.
|
||||
- **reference_ee_pose** (`np.ndarray | None`) -- Internal state storing the latched reference pose.
|
||||
- **_prev_enabled** (`bool`) -- Internal state to detect the rising edge of the enable signal.
|
||||
- **_command_when_disabled** (`np.ndarray | None`) -- Internal state to hold the last command while
|
||||
disabled.
|
||||
Attributes:
|
||||
kinematics: The robot's kinematic model for forward kinematics.
|
||||
end_effector_step_sizes: A dictionary scaling the input delta commands.
|
||||
motor_names: A list of motor names required for forward kinematics.
|
||||
use_latched_reference: If True, latch the reference pose on enable; otherwise, always use the
|
||||
current pose as the reference.
|
||||
reference_ee_pose: Internal state storing the latched reference pose.
|
||||
_prev_enabled: Internal state to detect the rising edge of the enable signal.
|
||||
_command_when_disabled: Internal state to hold the last command while disabled.
|
||||
"""
|
||||
|
||||
kinematics: RobotKinematics
|
||||
@@ -77,15 +77,6 @@ class EEReferenceAndDelta(RobotActionProcessorStep):
|
||||
_command_when_disabled: np.ndarray | None = field(default=None, init=False, repr=False)
|
||||
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
"""Transform the action for this step.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The incoming robot action.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The transformed action.
|
||||
"""
|
||||
raw_observation = self.transition.get(TransitionKey.OBSERVATION)
|
||||
|
||||
if raw_observation is None:
|
||||
@@ -176,16 +167,6 @@ class EEReferenceAndDelta(RobotActionProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
for feat in [
|
||||
"enabled",
|
||||
"target_x",
|
||||
@@ -209,19 +190,21 @@ class EEReferenceAndDelta(RobotActionProcessorStep):
|
||||
@ProcessorStepRegistry.register("ee_bounds_and_safety")
|
||||
@dataclass
|
||||
class EEBoundsAndSafety(RobotActionProcessorStep):
|
||||
"""Clips the end-effector pose to predefined bounds and checks for unsafe jumps.
|
||||
"""
|
||||
Clips the end-effector pose to predefined bounds and checks for unsafe jumps.
|
||||
|
||||
This step ensures that the target end-effector pose remains within a safe operational workspace.
|
||||
It also moderates the command to prevent large, sudden movements between consecutive steps.
|
||||
|
||||
**Attributes**:
|
||||
- **end_effector_bounds** (`dict`) -- A dictionary with "min" and "max" keys for position clipping.
|
||||
- **max_ee_step_m** (`float`) -- The maximum allowed change in position (in meters) between steps.
|
||||
- **raise_on_jump** (`bool`) -- When ``True`` (default) an over-limit per-frame step raises
|
||||
``ValueError`` (aborting the control loop). When ``False`` the step is rate-limited to
|
||||
``max_ee_step_m`` and a warning is logged instead — the safer choice for live teleoperation, where a
|
||||
transient tracking glitch should not crash the loop and leave the robot uncontrolled.
|
||||
- **_last_pos** (`np.ndarray | None`) -- Internal state storing the last commanded position.
|
||||
Attributes:
|
||||
end_effector_bounds: A dictionary with "min" and "max" keys for position clipping.
|
||||
max_ee_step_m: The maximum allowed change in position (in meters) between steps.
|
||||
raise_on_jump: When ``True`` (default) an over-limit per-frame step raises
|
||||
``ValueError`` (aborting the control loop). When ``False`` the step is
|
||||
rate-limited to ``max_ee_step_m`` and a warning is logged instead — the
|
||||
safer choice for live teleoperation, where a transient tracking glitch
|
||||
should not crash the loop and leave the robot uncontrolled.
|
||||
_last_pos: Internal state storing the last commanded position.
|
||||
"""
|
||||
|
||||
end_effector_bounds: dict
|
||||
@@ -230,15 +213,6 @@ class EEBoundsAndSafety(RobotActionProcessorStep):
|
||||
_last_pos: np.ndarray | None = field(default=None, init=False, repr=False)
|
||||
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
"""Transform the action for this step.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The incoming robot action.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The transformed action.
|
||||
"""
|
||||
x = action["ee.x"]
|
||||
y = action["ee.y"]
|
||||
z = action["ee.z"]
|
||||
@@ -294,39 +268,29 @@ class EEBoundsAndSafety(RobotActionProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("inverse_kinematics_ee_to_joints")
|
||||
@dataclass
|
||||
class InverseKinematicsEEToJoints(RobotActionProcessorStep):
|
||||
"""Computes desired joint positions from a target end-effector pose using inverse kinematics (IK).
|
||||
"""
|
||||
Computes desired joint positions from a target end-effector pose using inverse kinematics (IK).
|
||||
|
||||
This step translates a Cartesian command (position and orientation of the end-effector) into
|
||||
the corresponding joint-space commands for each motor.
|
||||
|
||||
**Attributes**:
|
||||
- **kinematics** (`RobotKinematics`) -- The robot's kinematic model for inverse kinematics.
|
||||
- **motor_names** (`list[str]`) -- A list of motor names for which to compute joint positions.
|
||||
- **q_curr** (`np.ndarray | None`) -- Internal state storing the last joint positions, used as an
|
||||
initial guess for the IK solver.
|
||||
- **initial_guess_current_joints** (`bool`) -- If True, use the robot's current joint state as the IK
|
||||
guess. If False, use the solution from the previous step.
|
||||
- **orientation_weight** (`float`) -- Weight for the orientation constraint passed to
|
||||
``RobotKinematics.inverse_kinematics``. Defaults to ``0.01`` (matching the solver default, so
|
||||
existing callers are unchanged). Set to ``0.0`` for position-only IK on under-actuated arms; a small
|
||||
nonzero weight gives soft-orientation IK on the 5-DOF SO-101, where the wrist tracks orientation
|
||||
only partially (position dominates).
|
||||
Attributes:
|
||||
kinematics: The robot's kinematic model for inverse kinematics.
|
||||
motor_names: A list of motor names for which to compute joint positions.
|
||||
q_curr: Internal state storing the last joint positions, used as an initial guess for the IK solver.
|
||||
initial_guess_current_joints: If True, use the robot's current joint state as the IK guess.
|
||||
If False, use the solution from the previous step.
|
||||
orientation_weight: Weight for the orientation constraint passed to
|
||||
``RobotKinematics.inverse_kinematics``. Defaults to ``0.01`` (matching the solver
|
||||
default, so existing callers are unchanged). Set to ``0.0`` for position-only IK on
|
||||
under-actuated arms; a small nonzero weight gives soft-orientation IK on the 5-DOF
|
||||
SO-101, where the wrist tracks orientation only partially (position dominates).
|
||||
"""
|
||||
|
||||
kinematics: RobotKinematics
|
||||
@@ -336,15 +300,6 @@ class InverseKinematicsEEToJoints(RobotActionProcessorStep):
|
||||
orientation_weight: float = 0.01
|
||||
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
"""Transform the action for this step.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The incoming robot action.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The transformed action.
|
||||
"""
|
||||
x = action.pop("ee.x")
|
||||
y = action.pop("ee.y")
|
||||
z = action.pop("ee.z")
|
||||
@@ -400,16 +355,6 @@ class InverseKinematicsEEToJoints(RobotActionProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
for feat in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]:
|
||||
features[PipelineFeatureType.ACTION].pop(f"ee.{feat}", None)
|
||||
|
||||
@@ -428,20 +373,20 @@ class InverseKinematicsEEToJoints(RobotActionProcessorStep):
|
||||
@ProcessorStepRegistry.register("gripper_velocity_to_joint")
|
||||
@dataclass
|
||||
class GripperVelocityToJoint(RobotActionProcessorStep):
|
||||
"""Converts a gripper velocity command into a target gripper joint position.
|
||||
"""
|
||||
Converts a gripper velocity command into a target gripper joint position.
|
||||
|
||||
This step integrates a normalized velocity command over time to produce a position command,
|
||||
taking the current gripper position as a starting point. It also supports a discrete mode
|
||||
where integer actions map to open, close, or no-op.
|
||||
|
||||
**Attributes**:
|
||||
- **motor_names** -- A list of motor names, which must include 'gripper'.
|
||||
- **speed_factor** (`float`) -- A scaling factor to convert the normalized velocity command to a
|
||||
position change.
|
||||
- **clip_min** (`float`) -- The minimum allowed gripper joint position.
|
||||
- **clip_max** (`float`) -- The maximum allowed gripper joint position.
|
||||
- **discrete_gripper** (`bool`) -- If True, interpret the input as a discrete class index {0 = close,
|
||||
1 = stay, 2 = open}, matching `GamepadTeleop.GripperAction`.
|
||||
Attributes:
|
||||
motor_names: A list of motor names, which must include 'gripper'.
|
||||
speed_factor: A scaling factor to convert the normalized velocity command to a position change.
|
||||
clip_min: The minimum allowed gripper joint position.
|
||||
clip_max: The maximum allowed gripper joint position.
|
||||
discrete_gripper: If True, interpret the input as a discrete class index
|
||||
{0 = close, 1 = stay, 2 = open}, matching `GamepadTeleop.GripperAction`.
|
||||
"""
|
||||
|
||||
speed_factor: float = 20.0
|
||||
@@ -450,15 +395,6 @@ class GripperVelocityToJoint(RobotActionProcessorStep):
|
||||
discrete_gripper: bool = False
|
||||
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
"""Transform the action for this step.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The incoming robot action.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The transformed action.
|
||||
"""
|
||||
raw_observation = self.transition.get(TransitionKey.OBSERVATION)
|
||||
|
||||
gripper_vel = action.pop("ee.gripper_vel")
|
||||
@@ -492,16 +428,6 @@ class GripperVelocityToJoint(RobotActionProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
features[PipelineFeatureType.ACTION].pop("ee.gripper_vel", None)
|
||||
features[PipelineFeatureType.ACTION]["ee.gripper_pos"] = PolicyFeature(
|
||||
type=FeatureType.ACTION, shape=(1,)
|
||||
@@ -513,21 +439,6 @@ class GripperVelocityToJoint(RobotActionProcessorStep):
|
||||
def compute_forward_kinematics_joints_to_ee(
|
||||
joints: dict[str, Any], kinematics: RobotKinematics, motor_names: list[str]
|
||||
) -> dict[str, Any]:
|
||||
"""Replace joint positions with the end-effector pose they produce.
|
||||
|
||||
Args:
|
||||
joints (`dict[str, Any]`):
|
||||
Joint values keyed `"<motor>.pos"`, including `"gripper.pos"`. Modified in place: the joint
|
||||
keys named in `motor_names` are removed.
|
||||
kinematics (`RobotKinematics`):
|
||||
The arm's kinematic model.
|
||||
motor_names (`list[str]`):
|
||||
The motors, in the order the kinematic model expects them.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The same dict with `ee.x`, `ee.y`, `ee.z` for position, `ee.wx`, `ee.wy`,
|
||||
`ee.wz` for orientation as a rotation vector, and `ee.gripper_pos` carried through unchanged.
|
||||
"""
|
||||
motor_joint_values = [joints[f"{n}.pos"] for n in motor_names]
|
||||
|
||||
q = np.array(motor_joint_values, dtype=float)
|
||||
@@ -550,44 +461,26 @@ def compute_forward_kinematics_joints_to_ee(
|
||||
@ProcessorStepRegistry.register("forward_kinematics_joints_to_ee_observation")
|
||||
@dataclass
|
||||
class ForwardKinematicsJointsToEEObservation(ObservationProcessorStep):
|
||||
"""Computes the end-effector pose from joint positions using forward kinematics (FK).
|
||||
"""
|
||||
Computes the end-effector pose from joint positions using forward kinematics (FK).
|
||||
|
||||
This step is typically used to add the robot's Cartesian pose to the observation space,
|
||||
which can be useful for visualization or as an input to a policy.
|
||||
|
||||
**Attributes**:
|
||||
- **kinematics** (`RobotKinematics`) -- The robot's kinematic model.
|
||||
Attributes:
|
||||
kinematics: The robot's kinematic model.
|
||||
"""
|
||||
|
||||
kinematics: RobotKinematics
|
||||
motor_names: list[str]
|
||||
|
||||
def observation(self, observation: RobotObservation) -> RobotObservation:
|
||||
"""Replace the observation's joint positions with the end-effector pose.
|
||||
|
||||
Args:
|
||||
observation (`dict[str, Any]`):
|
||||
The incoming observation, containing `"<motor>.pos"` keys.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The observation with `ee.*` keys in place of the joint positions.
|
||||
"""
|
||||
return compute_forward_kinematics_joints_to_ee(observation, self.kinematics, self.motor_names)
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
# We only use the ee pose in the dataset, so we don't need the joint positions
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
for n in self.motor_names:
|
||||
features[PipelineFeatureType.OBSERVATION].pop(f"{n}.pos", None)
|
||||
# We specify the dataset features of this step that we want to be stored in the dataset
|
||||
@@ -601,44 +494,26 @@ class ForwardKinematicsJointsToEEObservation(ObservationProcessorStep):
|
||||
@ProcessorStepRegistry.register("forward_kinematics_joints_to_ee_action")
|
||||
@dataclass
|
||||
class ForwardKinematicsJointsToEEAction(RobotActionProcessorStep):
|
||||
"""Computes the end-effector pose from joint positions using forward kinematics (FK).
|
||||
"""
|
||||
Computes the end-effector pose from joint positions using forward kinematics (FK).
|
||||
|
||||
This step is typically used to add the robot's Cartesian pose to the observation space,
|
||||
which can be useful for visualization or as an input to a policy.
|
||||
|
||||
**Attributes**:
|
||||
- **kinematics** (`RobotKinematics`) -- The robot's kinematic model.
|
||||
Attributes:
|
||||
kinematics: The robot's kinematic model.
|
||||
"""
|
||||
|
||||
kinematics: RobotKinematics
|
||||
motor_names: list[str]
|
||||
|
||||
def action(self, action: RobotAction) -> RobotAction:
|
||||
"""Transform the action for this step.
|
||||
|
||||
Args:
|
||||
action (`dict[str, Any]`):
|
||||
The incoming robot action.
|
||||
|
||||
Returns:
|
||||
`dict[str, Any]`: The transformed action.
|
||||
"""
|
||||
return compute_forward_kinematics_joints_to_ee(action, self.kinematics, self.motor_names)
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
# We only use the ee pose in the dataset, so we don't need the joint positions
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
for n in self.motor_names:
|
||||
features[PipelineFeatureType.ACTION].pop(f"{n}.pos", None)
|
||||
# Store end-effector features as actions in the dataset schema
|
||||
@@ -652,21 +527,10 @@ class ForwardKinematicsJointsToEEAction(RobotActionProcessorStep):
|
||||
@ProcessorStepRegistry.register(name="forward_kinematics_joints_to_ee")
|
||||
@dataclass
|
||||
class ForwardKinematicsJointsToEE(ProcessorStep):
|
||||
"""Applies forward kinematics to whichever of the action and observation are present.
|
||||
|
||||
A convenience wrapper over [`ForwardKinematicsJointsToEEAction`] and
|
||||
[`ForwardKinematicsJointsToEEObservation`], so a pipeline needs one step instead of two.
|
||||
|
||||
**Attributes**:
|
||||
- **kinematics** (`RobotKinematics`) -- The arm's kinematic model.
|
||||
- **motor_names** (`list[str]`) -- The motors, in the order the kinematic model expects them.
|
||||
"""
|
||||
|
||||
kinematics: RobotKinematics
|
||||
motor_names: list[str]
|
||||
|
||||
def __post_init__(self):
|
||||
"""Build the action and observation sub-steps this step delegates to."""
|
||||
self.joints_to_ee_action_processor = ForwardKinematicsJointsToEEAction(
|
||||
kinematics=self.kinematics, motor_names=self.motor_names
|
||||
)
|
||||
@@ -675,15 +539,6 @@ class ForwardKinematicsJointsToEE(ProcessorStep):
|
||||
)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
"""Apply forward kinematics to whichever of the action and observation are present.
|
||||
|
||||
Args:
|
||||
transition (`EnvTransition`):
|
||||
The transition to transform.
|
||||
|
||||
Returns:
|
||||
`EnvTransition`: The transition with `ee.*` keys in place of joint positions.
|
||||
"""
|
||||
if transition.get(TransitionKey.ACTION) is not None:
|
||||
transition = self.joints_to_ee_action_processor(transition)
|
||||
if transition.get(TransitionKey.OBSERVATION) is not None:
|
||||
@@ -693,16 +548,6 @@ class ForwardKinematicsJointsToEE(ProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
if features[PipelineFeatureType.ACTION] is not None:
|
||||
features = self.joints_to_ee_action_processor.transform_features(features)
|
||||
if features[PipelineFeatureType.OBSERVATION] is not None:
|
||||
@@ -713,7 +558,8 @@ class ForwardKinematicsJointsToEE(ProcessorStep):
|
||||
@ProcessorStepRegistry.register("inverse_kinematics_rl_step")
|
||||
@dataclass
|
||||
class InverseKinematicsRLStep(ProcessorStep):
|
||||
"""Computes desired joint positions from a target end-effector pose using inverse kinematics (IK).
|
||||
"""
|
||||
Computes desired joint positions from a target end-effector pose using inverse kinematics (IK).
|
||||
|
||||
This is modified from the InverseKinematicsEEToJoints step to be used in the RL pipeline.
|
||||
"""
|
||||
@@ -724,15 +570,6 @@ class InverseKinematicsRLStep(ProcessorStep):
|
||||
initial_guess_current_joints: bool = True
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
"""Solve inverse kinematics for the transition's end-effector action.
|
||||
|
||||
Args:
|
||||
transition (`EnvTransition`):
|
||||
The transition to transform.
|
||||
|
||||
Returns:
|
||||
`EnvTransition`: The transition with joint targets in place of the `ee.*` action.
|
||||
"""
|
||||
new_transition = dict(transition)
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
@@ -796,16 +633,6 @@ class InverseKinematicsRLStep(ProcessorStep):
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Update the feature contract to match what this step does to the data.
|
||||
|
||||
Args:
|
||||
features (`dict[PipelineFeatureType, dict[str, PolicyFeature]]`):
|
||||
The pipeline's feature contract so far.
|
||||
|
||||
Returns:
|
||||
`dict[PipelineFeatureType, dict[str, PolicyFeature]]`: The contract with this step's key changes
|
||||
applied.
|
||||
"""
|
||||
for feat in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]:
|
||||
features[PipelineFeatureType.ACTION].pop(f"ee.{feat}", None)
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user