Compare commits

...

9 Commits

Author SHA1 Message Date
dependabot[bot] c2ac99abf2 chore(deps): bump the actions group across 1 directory with 6 updates
Bumps the actions group with 6 updates in the / directory:

| Package | From | To |
| --- | --- | --- |
| [docker/login-action](https://github.com/docker/login-action) | `4.4.0` | `4.6.0` |
| [anthropics/claude-code-action](https://github.com/anthropics/claude-code-action) | `1.0.179` | `1.0.183` |
| [astral-sh/setup-uv](https://github.com/astral-sh/setup-uv) | `8.3.2` | `9.0.0` |
| [pypa/gh-action-pypi-publish](https://github.com/pypa/gh-action-pypi-publish) | `1.14.1` | `1.14.2` |
| [trufflesecurity/trufflehog](https://github.com/trufflesecurity/trufflehog) | `3.95.9` | `3.96.0` |
| [actions/stale](https://github.com/actions/stale) | `10` | `11` |



Updates `docker/login-action` from 4.4.0 to 4.6.0
- [Release notes](https://github.com/docker/login-action/releases)
- [Commits](https://github.com/docker/login-action/compare/v4.4.0...v4.6.0)

Updates `anthropics/claude-code-action` from 1.0.179 to 1.0.183
- [Release notes](https://github.com/anthropics/claude-code-action/releases)
- [Commits](https://github.com/anthropics/claude-code-action/compare/b76a0776ae74036e77cd11018083743453d7ad35...be7b93b1907a4abad570368f3c74b6fe3807510b)

Updates `astral-sh/setup-uv` from 8.3.2 to 9.0.0
- [Release notes](https://github.com/astral-sh/setup-uv/releases)
- [Commits](https://github.com/astral-sh/setup-uv/compare/v8.3.2...v9.0.0)

Updates `pypa/gh-action-pypi-publish` from 1.14.1 to 1.14.2
- [Release notes](https://github.com/pypa/gh-action-pypi-publish/releases)
- [Commits](https://github.com/pypa/gh-action-pypi-publish/compare/ba38be9e461d3875417946c167d0b5f3d385a247...dc37677b2e1c63e2034f94d8a5b11f265b73ba33)

Updates `trufflesecurity/trufflehog` from 3.95.9 to 3.96.0
- [Release notes](https://github.com/trufflesecurity/trufflehog/releases)
- [Commits](https://github.com/trufflesecurity/trufflehog/compare/27b0417c16317ca9a472a9a8092acce143b49c55...6f3c981e7b77f235fd2702dd74af25fc4b72bf11)

Updates `actions/stale` from 10 to 11
- [Release notes](https://github.com/actions/stale/releases)
- [Changelog](https://github.com/actions/stale/blob/main/CHANGELOG.md)
- [Commits](https://github.com/actions/stale/compare/v10...v11)

---
updated-dependencies:
- dependency-name: docker/login-action
  dependency-version: 4.6.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: actions
- dependency-name: anthropics/claude-code-action
  dependency-version: 1.0.183
  dependency-type: direct:production
  update-type: version-update:semver-patch
  dependency-group: actions
- dependency-name: astral-sh/setup-uv
  dependency-version: 9.0.0
  dependency-type: direct:production
  update-type: version-update:semver-major
  dependency-group: actions
- dependency-name: pypa/gh-action-pypi-publish
  dependency-version: 1.14.2
  dependency-type: direct:production
  update-type: version-update:semver-patch
  dependency-group: actions
- dependency-name: trufflesecurity/trufflehog
  dependency-version: 3.96.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: actions
- dependency-name: actions/stale
  dependency-version: '11'
  dependency-type: direct:production
  update-type: version-update:semver-major
  dependency-group: actions
...

Signed-off-by: dependabot[bot] <support@github.com>
2026-08-07 14:00:00 +00:00
Pepijn 6c73c413eb docs: add API documentation infrastructure (#4348)
* docs: add API documentation infrastructure

LeRobot's documentation build passes `--not_python_module`, which tells
doc-builder there is no importable Python package and disables `[[autodoc]]`
entirely. The result is that all 90+ pages are hand-written guides and there is
no generated API reference at all.

This is the machinery to change that. It deliberately contains no docstring
changes of its own — every docstring edit lives in the follow-up PR, so this
one can be reviewed as tooling and configuration alone.

**The standard.** `docs/source/writing_docstrings.mdx` is the contract: Google
section headers with Hugging Face type formatting, the machine-checked argument
line, `**Attributes**:`, doc-builder cross-references, fenced doctest examples.
It also records three behaviours that are not discoverable from the source and
were verified against a local build: `[[autodoc]]` silently skips members with
no docstring; doc-builder does not inherit docstrings from base classes, so a
registered config shim whose body is `pass` renders every field with no
description; and module-level aliases resolve to the canonical class.

**Autodoc turned on**, with two changes that are not obvious:

- `--version main` on the main-docs job. Without `--not_python_module`,
  doc-builder resolves the version from `lerobot.__version__` and only maps it
  to the default branch when it contains "dev". transformers relies on that;
  our main carries 0.6.2. Verified by building both ways — dropping the flag
  alone would publish the main docs to /lerobot/v0.6.2/ instead of
  /lerobot/main/ and disable notebook building.
- `pre_command` on both jobs. doc-builder ships a mock-deps registry entry for
  lerobot, so the reusable workflow takes its light-install path, which cannot
  import the package. The heavy dependencies cannot be mocked either: draccus
  runs `register_subclass` at import time and `processor/converters.py` calls
  `functools.singledispatch.register(torch.Tensor)`, which needs a real class.
  `[dataset]` is the only extra required.

Workflow triggers gain `src/**`, since the reference is now generated from
docstrings. `docs/source/api/` is excluded from the prettier hook, which reads
`[[autodoc]]` member lists as lazy paragraph continuations and joins a ten-entry
list onto one line.

Nine API reference pages, scaffolded with each module's base class.

**Doctests.** `LeRobotDocTestParser` is mandatory rather than optional here:
ruff's `docstring-code-format = true` drops the blank line before a closing
fence, after which stdlib's `_EXAMPLE_RE` reads the fence as expected output and
every example with output fails. It is written against the installed pytest
rather than copied from transformers, whose version predates pytest 9's
`import_path` signature and its own fix for the `@property` line-number bug.
`preprocess_string` also diverges: the upstream fenced-block split puts a
single-line example's code in a chunk with no `>>>` in it, so neither the CUDA
skip nor the `+IGNORE_RESULT` injection fires for it.

**Checkers.** `utils/check_docstrings.py` is the ~300-line core of the
2203-line transformers original; the `@auto_docstring` system, modular
propagation, GitPython and `checkers.py` are not ported.
`utils/check_config_docstrings.py` checks that every registered robot config
documents its port and calibration semantics.

**Gates**, all set to values that pass today: ruff `D` with per-file-ignores
per unconverted module, `interrogate` at `fail-under = 52` against a measured
52.1%, and Makefile targets wired into the quality workflow. The doctest
allowlist ships empty and the `doctest` target handles that, because the files
carrying runnable examples arrive with the docstring PR.

* ci: build the docs on Python 3.12

The shared doc-builder workflows create their virtualenv with the runner's
system Python, which is 3.10.12 on ubuntu-22.04. lerobot requires >=3.12, so
the build died during "Setup environment":

    × No solution found when resolving dependencies:
    ╰─▶ Because the current Python version (3.10.12) does not satisfy
        Python>=3.12 and lerobot==0.6.2 depends on Python>=3.12 ...

That step runs before `pre_command`, so the real install this workflow already
performs never got the chance to run. There was no fix available on the caller
side either: `env:` does not propagate into a reusable workflow, so `UV_PYTHON`
is unavailable, and `uv venv` runs in the runner workspace root rather than the
checkout, so a `.python-version` file cannot reach it. The non-light fallback
(`uv pip install "./pkg[dev]"`) fails identically, so this is not specific to
the mock-deps path — it blocks any package requiring 3.12+.

huggingface/doc-builder#808 adds a `python_version` input to both build
workflows, which this passes. Pins move to that merge commit, picking up three
unrelated fixes in the same range (#810, #811, #812); the upload workflow is
unchanged there and is bumped only to keep all three pins on one SHA.

* chore: sort imports in vla_jepa tests

Enabling pydocstyle in the previous commit changes how ruff determines where a
module's import block ends, which makes I001 fire on three vla_jepa tests that
were clean before. The blank line between the `conftest` and `lerobot` imports
is the trigger: both are first-party, so isort wants them in one contiguous
block, and the docstring-aware analysis is what makes it notice.

These files are unrelated to the API reference, so the fix is only to satisfy
the new gate.
2026-08-07 15:55:25 +02:00
Pepijn 3aabd135d3 fix(tests): port the multi-GPU FSDP test off the superseded accelerate YAML flow (#4370)
The FSDP multi-GPU test still generated an `accelerate launch --config_file`
FSDP1 YAML, which exports ACCELERATE_USE_FSDP into the workers. Since the
FSDP2/parallelism rewrite, `guard_against_env_interference()` hard-errors on
exactly that variable, so the test failed on every rank. Its assertions were
stale too: sharded runs now write DCP optimizer shards, not a gathered
`optimizer_state.safetensors`.

Drop the YAML generation entirely and use `accelerate launch` as the plain
launcher the docs describe, with the topology coming from `--parallelism.*`.

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-08-07 20:54:14 +08:00
Pepijn 2c1adc378e test(ema): Rename test that trips TruffleHog's Lob detector (#4371)
TruffleHog's Lob API-key detector matches `\b((live|test)_[a-zA-Z0-9_]{35})\b`.
One EMA config test was named `test_` plus exactly 35 word characters, so it
matched, and Lob's verifier treats the 403/422 that api.lob.com returns for junk
keys as proof of a live key. That failed the Security workflow on main after
huggingface/lerobot#4323 with "Found verified Lob result" and exit code 183.

Add one word to the name so the suffix is 39 characters rather than 35. No
behavior change. Note that the offending string is deliberately not spelled out
here: TruffleHog scans commit messages as well as added lines, so quoting it
would retrigger the very detector this commit is working around.

Refs #4323

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-08-07 14:22:41 +02:00
alejodosr 266be2bd17 feat(train): add opt-in EMA of the policy weights (--ema.enable=true) (#4323)
* feat(train): add opt-in EMA of the policy weights (--ema.enable=true)

Maintain an EMA shadow via diffusers' EMAModel (lazy import, no new
dependency) with the reference Diffusion Policy schedule. Saves the
shadow for exact resume plus a loadable pretrained_model_ema/ per
checkpoint, evaluates the EMA weights during env eval, and pushes them
to a sibling <repo_id>-ema repo. Fixes huggingface/lerobot#4259.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

* docs(diffusion): document the --ema.enable training flag

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

* fix(tests): skip EMA training tests when accelerate/diffusers are missing

* feat(train): support constant EMA decay (--ema.decay) for openpi-style policies

* fix(train): gate EMA step on sync_gradients; use parallel_dims.is_sharded guard

---------

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-08-07 13:52:23 +02:00
Haoming Song ff7cc3de1d fix(train): narrow the accelerate env guard, and fix a device-bound assert (#4347)
Two post-merge CI failures on main, both from #4010.

Benchmark Integration Tests (Libero) — `accelerate launch` exports whole groups
of variables unconditionally (the five ACCELERATE_DYNAMO_* it writes default the
backend to "no"), so matching on prefixes refused launches that configure
nothing, contradicting the documented flow where accelerate is supported as a
plain launcher. The guard now watches only the three switches that hand a
subsystem to the environment.

GPU Tests — `test_metrics_tracker_reduce_across_ranks_invokes_all_reduce`
compared the captured reduction buffer against a CPU tensor, so the assert
raised "Expected all tensors to be on the same device" wherever CUDA is
available. The expected tensor is built on the buffer's device instead.
2026-08-06 18:39:24 +02:00
Hiroaki.Ishikawa 31fedfd9dd fix(dataset): use conservative bounds for quantile aggregation instead of incorrect weighted mean (#3804)
* fix(stats): use conservative bounds for quantile aggregation instead of incorrect weighted mean

* docs: add --overwrite/--skip-images/--root options to augment_dataset_quantile_stats usage

* fix(dataset): clarify quantile aggregation semantics

* fix(augment): handle quantile stats edge cases
2026-08-06 18:39:01 +02:00
Pepijn b1bf24f565 feat(train): split phase timing metrics (#4344) 2026-08-06 17:20:35 +02:00
Haoming Song ef88d4e52b feat(train): parallel training framework — FSDP2, HSDP, gradient accumulation, and DCP checkpoints (#4010)
* feat(train): parallel training engine with FSDP2, HSDP, and DCP checkpoints

Replace the FSDP1 training path with a config-owned parallel-training
engine:

- Topology and runtime configs (--parallelism.*, --accelerator.*):
  dp_replicate x dp_shard degrees select single-process, DDP (unchanged
  default), FSDP2, or HSDP; mixed precision, first-class gradient
  accumulation, and FSDP/DDP tuning knobs are mirrored as plain
  dataclasses that build the accelerate objects at runtime, so every
  run is reproducible from its train_config.json alone. Accelerate env
  vars are guarded against configuring the engine behind the config
  system's back.
- Declarative policy surface: policies declare FSDP2 wrap units
  (_fsdp_wrap_modules) and non-forward entry points
  (_fsdp_forward_methods); a shared engine resolves them around
  accelerator.prepare(). Context-parallel fields are reserved and
  validated to 1.
- Checkpoints: selectable --checkpoint_format (safetensors | dcp |
  safetensors_dcp); the sharded optimizer channel is always DCP;
  two-phase resume (step+RNG before prepare, DCP model/optimizer after)
  reshards across GPU-topology changes; lerobot-convert-dcp merges DCP
  shards into a distributable model.safetensors offline.
- Publishing: PreTrainedPolicy.push_model_to_hub is replaced by the
  free publish_trained_model (model + processors + card + train config,
  all-ranks gather with main-rank writes);
  PreTrainedPolicy._save_pretrained gathers state dicts internally,
  removing the state_dict= threading from save_pretrained.
- lerobot_train is restructured around the engine: optimizer built
  before the single prepare() call, deferred weight load on DCP
  resumes, collective save_checkpoint with no call-site rank branches,
  dp-world-size-based sample accounting.

Breaking changes: FSDP checkpoints from lerobot <= 0.6.x are not
resumable (weights stay loadable via from_pretrained; pin
lerobot==0.6.x to finish old runs); the `accelerate launch
--config_file` yaml flow is superseded by the config flags; training
autocast is owned exclusively by --accelerator.mixed_precision
(policy.dtype only casts parameters).

Also fixes: reward-model hub publishing crash (TypeError on extra
kwargs).

Verified by ~200 new CPU tests (config round-trips, checkpoint
round-trips per format, two-phase resume, publisher contracts,
converter equivalence, accelerate canaries), a 5-test 4-GPU suite
(FSDP2 save/resume bit-exactness, HSDP/DDP loss parity,
changed-topology resume, all-ranks save_pretrained, grad-accum
equivalence), and end-to-end ACT (1/4/8 GPUs) + FastWAM 6B
(FSDP2 + HSDP) training runs.
2026-08-06 19:16:41 +08:00
95 changed files with 7825 additions and 1365 deletions
+8 -8
View File
@@ -84,7 +84,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -242,7 +242,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -344,7 +344,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -451,7 +451,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -552,7 +552,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -660,7 +660,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -766,7 +766,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -870,7 +870,7 @@ jobs:
- name: Login to Docker Hub - name: Login to Docker Hub
if: ${{ env.DOCKERHUB_USERNAME != '' }} if: ${{ env.DOCKERHUB_USERNAME != '' }}
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
+1 -1
View File
@@ -53,7 +53,7 @@ jobs:
- name: Run Claude Code - name: Run Claude Code
id: claude id: claude
uses: anthropics/claude-code-action@b76a0776ae74036e77cd11018083743453d7ad35 # v1.0.179 uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1.0.183
with: with:
anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }} anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }}
additional_permissions: | additional_permissions: |
+2 -2
View File
@@ -61,7 +61,7 @@ jobs:
with: with:
cache-binary: false cache-binary: false
- name: Login to Docker Hub - name: Login to Docker Hub
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -96,7 +96,7 @@ jobs:
with: with:
cache-binary: false cache-binary: false
- name: Login to Docker Hub - name: Login to Docker Hub
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
@@ -33,7 +33,7 @@ jobs:
github.event.workflow_run.event == 'pull_request' && github.event.workflow_run.event == 'pull_request' &&
github.event.workflow_run.conclusion == 'success' && github.event.workflow_run.conclusion == 'success' &&
github.repository == 'huggingface/lerobot' github.repository == 'huggingface/lerobot'
uses: huggingface/doc-builder/.github/workflows/upload_pr_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main uses: huggingface/doc-builder/.github/workflows/upload_pr_documentation.yml@931031bf2b54aabb134ceb54980a6a2860a00f11 # main
with: with:
package_name: lerobot package_name: lerobot
secrets: secrets:
+28 -7
View File
@@ -24,19 +24,24 @@ on:
required: false required: false
type: string type: string
# Triggers the workflow on push events to main for the docs folder # Triggers on pushes to main that touch the docs or the sources the API reference is generated from.
# `src/**` is included because the API reference is built from docstrings via `[[autodoc]]`: without it,
# published API pages would go stale as soon as a docstring changed.
push: push:
branches: branches:
- main - main
paths: paths:
- "docs/**" - "docs/**"
- "src/**"
# Triggers the workflow on pull request events targeting main for the docs folder # Same for pull requests, so a docstring change gets a preview build and a broken `[[autodoc]]` path
# fails the PR rather than main.
pull_request: pull_request:
branches: branches:
- main - main
paths: paths:
- "docs/**" - "docs/**"
- "src/**"
release: release:
types: [published] types: [published]
@@ -55,16 +60,29 @@ jobs:
github.repository == 'huggingface/lerobot' github.repository == 'huggingface/lerobot'
permissions: permissions:
contents: read contents: read
uses: huggingface/doc-builder/.github/workflows/build_main_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main uses: huggingface/doc-builder/.github/workflows/build_main_documentation.yml@931031bf2b54aabb134ceb54980a6a2860a00f11 # main
with: with:
commit_sha: ${{ github.sha }} commit_sha: ${{ github.sha }}
package: lerobot package: lerobot
# The shared workflow builds its venv with the runner's system Python, which is 3.10 on
# ubuntu-22.04. lerobot requires >=3.12, so without this the install fails during setup —
# before `pre_command` below ever runs. Added upstream in huggingface/doc-builder#808.
python_version: "3.12"
# 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: >- additional_args: >-
--not_python_module
${{ ${{
(github.event_name == 'release' && format('--version {0}', github.event.release.tag_name)) || (github.event_name == 'release' && format('--version {0}', github.event.release.tag_name)) ||
(inputs.version != '' && format('--version {0}', inputs.version)) || (inputs.version != '' && format('--version {0}', inputs.version)) ||
'' '--version main'
}} }}
secrets: secrets:
token: ${{ secrets.HUGGINGFACE_PUSH }} token: ${{ secrets.HUGGINGFACE_PUSH }}
@@ -78,9 +96,12 @@ jobs:
permissions: permissions:
contents: read contents: read
pull-requests: write pull-requests: write
uses: huggingface/doc-builder/.github/workflows/build_pr_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main uses: huggingface/doc-builder/.github/workflows/build_pr_documentation.yml@931031bf2b54aabb134ceb54980a6a2860a00f11 # main
with: with:
commit_sha: ${{ github.event.pull_request.head.sha }} commit_sha: ${{ github.event.pull_request.head.sha }}
pr_number: ${{ github.event.number }} pr_number: ${{ github.event.number }}
package: lerobot package: lerobot
additional_args: --not_python_module # See the comment on build_main_docs. The PR workflow passes its own `--version pr_<n>`, so no
# additional_args are needed here.
python_version: "3.12"
pre_command: uv pip install "./lerobot[dataset]"
+1 -1
View File
@@ -87,7 +87,7 @@ jobs:
libusb-1.0-0-dev speech-dispatcher libgeos-dev portaudio19-dev libusb-1.0-0-dev speech-dispatcher libgeos-dev portaudio19-dev
- name: Setup uv and Python - name: Setup uv and Python
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
with: with:
enable-cache: true enable-cache: true
version: ${{ env.UV_VERSION }} version: ${{ env.UV_VERSION }}
+2 -2
View File
@@ -80,7 +80,7 @@ jobs:
speech-dispatcher libgeos-dev portaudio19-dev speech-dispatcher libgeos-dev portaudio19-dev
- name: Setup uv and Python - name: Setup uv and Python
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
with: with:
enable-cache: true enable-cache: true
version: ${{ env.UV_VERSION }} version: ${{ env.UV_VERSION }}
@@ -146,7 +146,7 @@ jobs:
with: with:
cache-binary: false cache-binary: false
- name: Login to Docker Hub - name: Login to Docker Hub
uses: docker/login-action@af1e73f918a031802d376d3c8bbc3fe56130a9b0 # v4.4.0 uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
+3 -3
View File
@@ -53,7 +53,7 @@ jobs:
persist-credentials: false persist-credentials: false
- name: Setup uv and Python - name: Setup uv and Python
uses: astral-sh/setup-uv@v8.3.2 # zizmor: ignore[unpinned-uses] uses: astral-sh/setup-uv@v9.0.0 # zizmor: ignore[unpinned-uses]
with: with:
version: ${{ env.UV_VERSION }} version: ${{ env.UV_VERSION }}
python-version: ${{ env.PYTHON_VERSION }} python-version: ${{ env.PYTHON_VERSION }}
@@ -115,7 +115,7 @@ jobs:
speech-dispatcher libgeos-dev portaudio19-dev speech-dispatcher libgeos-dev portaudio19-dev
- name: Setup uv and Python - name: Setup uv and Python
uses: astral-sh/setup-uv@v8.3.2 # zizmor: ignore[unpinned-uses] uses: astral-sh/setup-uv@v9.0.0 # zizmor: ignore[unpinned-uses]
with: with:
enable-cache: true enable-cache: true
version: ${{ env.UV_VERSION }} version: ${{ env.UV_VERSION }}
@@ -168,7 +168,7 @@ jobs:
with: with:
cache-binary: false cache-binary: false
- name: Login to Docker Hub - name: Login to Docker Hub
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses] uses: docker/login-action@v4.6.0 # zizmor: ignore[unpinned-uses]
with: with:
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }} username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }} password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
+38
View File
@@ -56,3 +56,41 @@ jobs:
uses: pre-commit/action@2c7b3805fd2a0fd8c1884dcaebf91fc102a13ecd # v3.0.1 uses: pre-commit/action@2c7b3805fd2a0fd8c1884dcaebf91fc102a13ecd # v3.0.1
with: with:
extra_args: --all-files --show-diff-on-failure --color=always 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@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
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
+3 -3
View File
@@ -104,7 +104,7 @@ jobs:
- name: Publish to TestPyPI for pre-releases - name: Publish to TestPyPI for pre-releases
# True for tags like 'v0.2.0-rc1' # True for tags like 'v0.2.0-rc1'
if: startsWith(github.ref, 'refs/tags/v') && contains(github.ref, '-') if: startsWith(github.ref, 'refs/tags/v') && contains(github.ref, '-')
uses: pypa/gh-action-pypi-publish@ba38be9e461d3875417946c167d0b5f3d385a247 # v1.14.1 uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
with: with:
repository-url: https://test.pypi.org/legacy/ repository-url: https://test.pypi.org/legacy/
verbose: true verbose: true
@@ -112,7 +112,7 @@ jobs:
- name: Publish to PyPI - name: Publish to PyPI
if: startsWith(github.ref, 'refs/tags/v') && !contains(github.ref, '-') if: startsWith(github.ref, 'refs/tags/v') && !contains(github.ref, '-')
uses: pypa/gh-action-pypi-publish@ba38be9e461d3875417946c167d0b5f3d385a247 # v1.14.1 uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
with: with:
verbose: true verbose: true
print-hash: true print-hash: true
@@ -137,7 +137,7 @@ jobs:
git curl libglib2.0-0 libegl1-mesa-dev ffmpeg libusb-1.0-0-dev \ git curl libglib2.0-0 libegl1-mesa-dev ffmpeg libusb-1.0-0-dev \
speech-dispatcher libgeos-dev portaudio19-dev speech-dispatcher libgeos-dev portaudio19-dev
- name: Setup uv and Python - name: Setup uv and Python
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
with: with:
enable-cache: true # zizmor: ignore[cache-poisoning] enable-cache: true # zizmor: ignore[cache-poisoning]
version: ${{ env.UV_VERSION }} version: ${{ env.UV_VERSION }}
+1 -1
View File
@@ -49,6 +49,6 @@ jobs:
persist-credentials: false persist-credentials: false
- name: Secret Scanning - name: Secret Scanning
uses: trufflesecurity/trufflehog@27b0417c16317ca9a472a9a8092acce143b49c55 # v3.95.9 uses: trufflesecurity/trufflehog@6f3c981e7b77f235fd2702dd74af25fc4b72bf11 # v3.96.0
with: with:
extra_args: --only-verified extra_args: --only-verified
+1 -1
View File
@@ -52,7 +52,7 @@ jobs:
issues: write issues: write
pull-requests: write pull-requests: write
steps: steps:
- uses: actions/stale@v10 - uses: actions/stale@v11
with: with:
repo-token: ${{ secrets.GITHUB_TOKEN }} repo-token: ${{ secrets.GITHUB_TOKEN }}
stale-issue-label: stale stale-issue-label: stale
+11 -2
View File
@@ -67,7 +67,11 @@ repos:
args: [--prose-wrap=preserve] args: [--prose-wrap=preserve]
# Jinja2 model-card templates use a .md extension but contain {% ... %} / # Jinja2 model-card templates use a .md extension but contain {% ... %} /
# {{ ... }} tags that prettier's Markdown formatter mangles (e.g. table loops). # {{ ... }} tags that prettier's Markdown formatter mangles (e.g. table loops).
exclude: ^src/lerobot/templates/.*\.md$ #
# docs/source/api/ holds the generated API reference. Its `[[autodoc]]` blocks restrict output
# to an indented `- member` list, which prettier reads as a lazy paragraph continuation and
# joins onto one line — silently turning a member list into part of the directive.
exclude: ^(src/lerobot/templates/.*\.md|docs/source/api/.*\.mdx)$
##### Security ##### ##### Security #####
- repo: https://github.com/gitleaks/gitleaks - repo: https://github.com/gitleaks/gitleaks
@@ -104,8 +108,13 @@ repos:
# args: ["--docstring-style", "google", "-v", "2"] # args: ["--docstring-style", "google", "-v", "2"]
# exclude: ^tests/.*$ # 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 # - repo: https://github.com/econchick/interrogate
# rev: 1.7.0 # rev: 1.7.0
# hooks: # hooks:
# - id: interrogate # - id: interrogate
# args: ["-vv", "--config=pyproject.toml"] # args: ["--config=pyproject.toml"]
# pass_filenames: false
+4
View File
@@ -50,6 +50,10 @@ To run checks manually on all files:
pre-commit run --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 ### Running Tests
We use `pytest`. First, ensure you have test artifacts by installing **git-lfs**: We use `pytest`. First, ensure you have test artifacts by installing **git-lfs**:
+26
View File
@@ -184,3 +184,29 @@ test-smolvla-ete-eval:
# backend, so it does not require a real model checkpoint or GPU. # backend, so it does not require a real model checkpoint or GPU.
annotation-e2e: annotation-e2e:
uv run python -m tests.annotations.run_e2e_smoke 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
View File
@@ -0,0 +1,60 @@
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Root conftest: makes doctest collection use LeRobot's parser.
This only affects `--doctest-modules` runs (see `make doctest`). The test suite itself is configured by
`tests/conftest.py`.
"""
import doctest
import _pytest.doctest
from lerobot.utils.doctest_utils import LeRobotDoctestModule, LeRobotDocTestParser
# Lets an example opt out of output comparison with `# doctest: +IGNORE_RESULT`, for calls whose output is
# a progress bar or otherwise not reproducible.
IGNORE_RESULT = doctest.register_optionflag("IGNORE_RESULT")
OutputChecker = doctest.OutputChecker
class CustomOutputChecker(OutputChecker):
"""An output checker that honours the `IGNORE_RESULT` flag."""
def check_output(self, want, got, optionflags):
"""Return `True` when `IGNORE_RESULT` is set, otherwise defer to stdlib.
Args:
want (`str`):
The expected output.
got (`str`):
The actual output.
optionflags (`int`):
Bitmask of active doctest option flags.
Returns:
`bool`: Whether the output is considered a match.
"""
if IGNORE_RESULT & optionflags:
return True
return OutputChecker.check_output(self, want, got, optionflags)
# Reassigning these module attributes is how doctest behaviour is customised; mypy sees it as assigning to
# a type, which is exactly what is intended here.
doctest.OutputChecker = CustomOutputChecker # type: ignore[misc]
_pytest.doctest.DoctestModule = LeRobotDoctestModule
doctest.DocTestParser = LeRobotDocTestParser # type: ignore[misc]
+22
View File
@@ -191,6 +191,28 @@
- sections: - sections:
- local: contributing - local: contributing
title: Contribute to LeRobot title: Contribute to LeRobot
- local: writing_docstrings
title: Writing docstrings
- local: backwardcomp - local: backwardcomp
title: Backward compatibility title: Backward compatibility
title: "About" title: "About"
- sections:
- local: api/robots
title: Robots
- local: api/teleoperators
title: Teleoperators
- local: api/cameras
title: Cameras
- local: api/motors
title: Motors
- local: api/datasets
title: Datasets
- local: api/policies
title: Policies
- local: api/processor
title: Processors
- local: api/envs
title: Environments
- local: api/configs
title: Configuration
title: "API Reference"
+24
View File
@@ -0,0 +1,24 @@
# Cameras
Cameras supply the image observations a policy sees. Every backend — OpenCV, Intel RealSense, Reachy 2 —
implements the [`Camera`] interface, so swapping hardware does not change the code that reads frames.
See the [Cameras guide](../cameras) for choosing and configuring a camera, and
[Third-Party Cameras & Sensors](../third_party_sensors) for devices outside the core set.
## Camera
[[autodoc]] lerobot.cameras.Camera
- connect
- disconnect
- read
- async_read
- find_cameras
## CameraConfig
[[autodoc]] lerobot.cameras.CameraConfig
## make_cameras_from_configs
[[autodoc]] lerobot.cameras.make_cameras_from_configs
+27
View File
@@ -0,0 +1,27 @@
# Configuration
LeRobot configuration is plain dataclasses parsed by [draccus](https://github.com/dlwh/draccus), so every
field is settable from the CLI. [`TrainPipelineConfig`] is the top-level object for `lerobot-train`.
Polymorphic configs (policies, robots, environments) use `draccus.ChoiceRegistry`: a subclass registers
itself with `@register_subclass("name")` and is then selectable by that name on the command line.
## TrainPipelineConfig
[[autodoc]] lerobot.configs.train.TrainPipelineConfig
## PreTrainedConfig
[[autodoc]] lerobot.configs.PreTrainedConfig
## DatasetConfig
[[autodoc]] lerobot.configs.DatasetConfig
## EvalConfig
[[autodoc]] lerobot.configs.EvalConfig
## WandBConfig
[[autodoc]] lerobot.configs.WandBConfig
+23
View File
@@ -0,0 +1,23 @@
# Datasets
[`LeRobotDataset`] is the format every LeRobot script reads and writes. It is episode-aware, decodes video
observations on the fly, and round-trips to the Hugging Face Hub.
See [Using LeRobotDataset](../lerobot-dataset-v3) for the format and the common operations,
[Porting Large Datasets](../porting_datasets_v3) for migration, and [Tools](../tools) for the CLI.
## LeRobotDataset
[[autodoc]] lerobot.datasets.LeRobotDataset
## LeRobotDatasetMetadata
[[autodoc]] lerobot.datasets.LeRobotDatasetMetadata
## MultiLeRobotDataset
[[autodoc]] lerobot.datasets.MultiLeRobotDataset
## StreamingLeRobotDataset
[[autodoc]] lerobot.datasets.StreamingLeRobotDataset
+19
View File
@@ -0,0 +1,19 @@
# Environments
Simulation environments are configured through [`EnvConfig`] and built by [`make_env`]. Each subclass
declares its `gym_kwargs` and how to construct the vectorised environments.
See [Environments from the Hub](../envhub) for using published environments and
[Adding a New Benchmark](../adding_benchmarks) for contributing one.
## EnvConfig
[[autodoc]] lerobot.envs.EnvConfig
## make_env
[[autodoc]] lerobot.envs.make_env
## make_env_config
[[autodoc]] lerobot.envs.make_env_config
+23
View File
@@ -0,0 +1,23 @@
# Motors
`MotorsBus` is the low-level interface to a chain of servos on a serial bus. Robots use it to read positions
and write goal positions; you rarely touch it directly unless you are adding hardware.
See [Bring Your Own Hardware](../integrate_hardware) for adding a new bus, and
[Updating Feetech Firmware](../feetech) and [Damiao Motors and CAN Bus](../damiao) for device-specific notes.
## MotorsBus
[[autodoc]] lerobot.motors.motors_bus.MotorsBus
## Motor
[[autodoc]] lerobot.motors.Motor
## MotorCalibration
[[autodoc]] lerobot.motors.MotorCalibration
## MotorNormMode
[[autodoc]] lerobot.motors.MotorNormMode
+20
View File
@@ -0,0 +1,20 @@
# 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
+20
View File
@@ -0,0 +1,20 @@
# Processors
Processors are the data transformation layer between a robot, a dataset and a policy. A pipeline is a chain
of [`ProcessorStep`]s; each step declares how it transforms both the data and the feature contract.
See [Introduction to Robot Processors](../introduction_processors) for the concepts,
[Implement your own processor](../implement_your_own_processor) to write a step, and
[Debug your processor pipeline](../debug_processor_pipeline) when a pipeline misbehaves.
## ProcessorStep
[[autodoc]] lerobot.processor.pipeline.ProcessorStep
## DataProcessorPipeline
[[autodoc]] lerobot.processor.pipeline.DataProcessorPipeline
## PolicyProcessorPipeline
[[autodoc]] lerobot.processor.pipeline.PolicyProcessorPipeline
+147
View File
@@ -0,0 +1,147 @@
# Robots
Every robot in LeRobot implements the [`Robot`] interface: connect, read an observation, send an action,
disconnect. Writing a policy or a recording script against that interface means it works with any supported
arm without change.
This page is the generated reference. For wiring, calibration and first-run instructions, start with the
hardware guides — [SO-101](../so101), [LeKiwi](../lekiwi), [Hope Jr](../hope_jr), [Reachy 2](../reachy2),
[OpenArm](../openarm) — or [Imitation Learning for Robots](../il_robots) for the end-to-end workflow. To add
a robot of your own, see [Bring Your Own Hardware](../integrate_hardware).
## Robot
The abstract base class. Subclasses implement every method below; the contract described here is what a
policy or recording loop can rely on.
[[autodoc]] lerobot.robots.Robot
- connect
- disconnect
- configure
- calibrate
- get_observation
- send_action
- observation_features
- action_features
- is_connected
- is_calibrated
## RobotConfig
[[autodoc]] lerobot.robots.RobotConfig
## make_robot_from_config
[[autodoc]] lerobot.robots.make_robot_from_config
## SO-100 and SO-101 followers
`SO100Follower` and `SO101Follower` are aliases of the same `SOFollower` class; the two arms differ in their
configuration, not their control code. `SO100FollowerConfig` and `SO101FollowerConfig` are likewise aliases
of `SOFollowerRobotConfig`.
[[autodoc]] lerobot.robots.so_follower.SOFollower
- all
[[autodoc]] lerobot.robots.so_follower.SOFollowerRobotConfig
## BiSOFollower
Two SO followers driven as one bimanual robot.
[[autodoc]] lerobot.robots.bi_so_follower.BiSOFollower
- all
[[autodoc]] lerobot.robots.bi_so_follower.BiSOFollowerConfig
## KochFollower
[[autodoc]] lerobot.robots.koch_follower.KochFollower
- all
[[autodoc]] lerobot.robots.koch_follower.KochFollowerConfig
## LeKiwi
`LeKiwi` runs on the robot itself. `LeKiwiClient` is the host-side proxy that talks to it over the network
and presents the same [`Robot`] interface.
[[autodoc]] lerobot.robots.lekiwi.LeKiwi
- all
[[autodoc]] lerobot.robots.lekiwi.LeKiwiConfig
[[autodoc]] lerobot.robots.lekiwi.LeKiwiClient
- all
[[autodoc]] lerobot.robots.lekiwi.LeKiwiClientConfig
## OpenArmFollower
[[autodoc]] lerobot.robots.openarm_follower.OpenArmFollower
- all
[[autodoc]] lerobot.robots.openarm_follower.OpenArmFollowerConfig
## BiOpenArmFollower
[[autodoc]] lerobot.robots.bi_openarm_follower.BiOpenArmFollower
- all
[[autodoc]] lerobot.robots.bi_openarm_follower.BiOpenArmFollowerConfig
## OmxFollower
[[autodoc]] lerobot.robots.omx_follower.OmxFollower
- all
[[autodoc]] lerobot.robots.omx_follower.OmxFollowerConfig
## Reachy2Robot
[[autodoc]] lerobot.robots.reachy2.Reachy2Robot
- all
[[autodoc]] lerobot.robots.reachy2.Reachy2RobotConfig
## UnitreeG1
[[autodoc]] lerobot.robots.unitree_g1.UnitreeG1
- all
[[autodoc]] lerobot.robots.unitree_g1.UnitreeG1Config
## Hope Jr
The Hope Jr humanoid is exposed as two independent robots, an arm and a hand.
[[autodoc]] lerobot.robots.hope_jr.HopeJrArm
- all
[[autodoc]] lerobot.robots.hope_jr.HopeJrArmConfig
[[autodoc]] lerobot.robots.hope_jr.HopeJrHand
- all
[[autodoc]] lerobot.robots.hope_jr.HopeJrHandConfig
## RebotB601Follower
[[autodoc]] lerobot.robots.rebot_b601_follower.RebotB601Follower
- all
[[autodoc]] lerobot.robots.rebot_b601_follower.RebotB601FollowerRobotConfig
## BiRebotB601Follower
[[autodoc]] lerobot.robots.bi_rebot_b601_follower.BiRebotB601Follower
- all
[[autodoc]] lerobot.robots.bi_rebot_b601_follower.BiRebotB601FollowerConfig
## EarthRoverMiniPlus
[[autodoc]] lerobot.robots.earthrover_mini_plus.EarthRoverMiniPlus
- all
[[autodoc]] lerobot.robots.earthrover_mini_plus.EarthRoverMiniPlusConfig
+30
View File
@@ -0,0 +1,30 @@
# Teleoperators
A teleoperator produces actions for a robot to follow — a leader arm, a gamepad, a keyboard, a phone. All of
them implement the [`Teleoperator`] interface, so a recording script written against it works with any input
device.
See [Phone teleoperation](../phone_teleop) and [Isaac Teleop](../isaac_teleop) for setup guides, and
[Imitation Learning for Robots](../il_robots) for the recording workflow.
## Teleoperator
[[autodoc]] lerobot.teleoperators.Teleoperator
- connect
- disconnect
- configure
- calibrate
- get_action
- send_feedback
- action_features
- feedback_features
- is_connected
- is_calibrated
## TeleoperatorConfig
[[autodoc]] lerobot.teleoperators.TeleoperatorConfig
## make_teleoperator_from_config
[[autodoc]] lerobot.teleoperators.make_teleoperator_from_config
+12 -2
View File
@@ -161,6 +161,16 @@ The methods called by the train/eval loops:
Batches are flat dictionaries keyed by the constants in [`lerobot.utils.constants`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/utils/constants.py): `OBS_STATE` (`observation.state.<motor>`), `OBS_IMAGES` (`observation.images.<camera>`), `OBS_LANGUAGE`, `ACTION`, etc. Reuse the constants — don't invent new prefixes. Batches are flat dictionaries keyed by the constants in [`lerobot.utils.constants`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/utils/constants.py): `OBS_STATE` (`observation.state.<motor>`), `OBS_IMAGES` (`observation.images.<camera>`), `OBS_LANGUAGE`, `ACTION`, etc. Reuse the constants — don't invent new prefixes.
If your model is large enough to warrant [sharded multi-GPU training](./multi_gpu_training#sharded-training-fsdp), also declare its FSDP wrap units — the repeated block classes sharding operates on:
```python
class MyPolicy(PreTrainedPolicy):
...
_fsdp_wrap_modules = ["MyTransformerBlock"]
```
With this one declaration, `--parallelism.dp_shard=N` works out of the box for your policy (users can still override it with `--accelerator.fsdp.wrap_modules`). Without any wrap source, sharded runs fail at startup by design.
### Processor functions ### Processor functions
LeRobot uses `PolicyProcessorPipeline`s to normalize inputs and de-normalize outputs around your policy. For a concrete reference, see [`processor_act.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/act/processor_act.py) or [`processor_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/processor_diffusion.py). LeRobot uses `PolicyProcessorPipeline`s to normalize inputs and de-normalize outputs around your policy. For a concrete reference, see [`processor_act.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/act/processor_act.py) or [`processor_diffusion.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/diffusion/processor_diffusion.py).
@@ -300,7 +310,7 @@ The file names are load-bearing: the factory does lazy imports by name, and the
Two places need to know about your policy. All by name. Two places need to know about your policy. All by name.
1. **`policies/__init__.py`** — re-export `MyPolicyConfig` and add it to `__all__`. This import is what registers your policy: `@PreTrainedConfig.register_subclass("my_policy")` runs, and from then on the factory resolves everything by convention. **Don't** re-export the modeling class; it loads lazily through the factory (so `import lerobot` stays fast). 1. **`policies/__init__.py`** — re-export `MyPolicyConfig` and add it to `__all__`. This import is what registers your policy: `@PreTrainedConfig.register_subclass("my_policy")` runs, and from then on the factory resolves everything by convention. **Don't** re-export the modeling class; it loads lazily through the factory (so `import lerobot` stays fast).
2. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what `push_model_to_hub` renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page. 2. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what the end-of-training publisher renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page.
Mirror an existing policy that's structurally similar to yours; the diff is small. Mirror an existing policy that's structurally similar to yours; the diff is small.
@@ -344,7 +354,7 @@ A new policy is much easier to review — and far more useful — when it ships
**Pick at least one in-tree benchmark.** LeRobot ships sim benchmarks with per-benchmark Docker images (LIBERO, LIBERO-plus, Meta-World, RoboTwin 2.0, RoboCasa365, RoboCerebra, RoboMME, VLABench and more). Pick the one that matches your policy's modality — VLAs usually go to LIBERO or VLABench; image-only BC to LIBERO or Meta-World. The full list lives under [Benchmarks](./libero) in the docs sidebar. **Pick at least one in-tree benchmark.** LeRobot ships sim benchmarks with per-benchmark Docker images (LIBERO, LIBERO-plus, Meta-World, RoboTwin 2.0, RoboCasa365, RoboCerebra, RoboMME, VLABench and more). Pick the one that matches your policy's modality — VLAs usually go to LIBERO or VLABench; image-only BC to LIBERO or Meta-World. The full list lives under [Benchmarks](./libero) in the docs sidebar.
**Push the checkpoint & processors** to the Hub under `lerobot/<policy>_<benchmark>` (or your namespace if you don't have write access; a maintainer can mirror it). Use `PreTrainedPolicy.push_model_to_hub` so the repo gets `config.json`, `model.safetensors`, and a model card. **Push the checkpoint & processors** to the Hub under `lerobot/<policy>_<benchmark>` (or your namespace if you don't have write access; a maintainer can mirror it). The easiest way is training with `--policy.repo_id=<namespace>/<repo>` and `--policy.push_to_hub=true`: `lerobot-train` publishes the model, both processors, and a model card at the end of the run. To publish an existing checkpoint after the fact, upload its `pretrained_model/` directory (e.g. `huggingface-cli upload`), or use `lerobot-convert-dcp --push_to_hub=...` for sharded-format checkpoints.
**Report results in your policy's MDX**, with the exact `lerobot-eval` command and hardware so anyone can re-run: **Report results in your policy's MDX**, with the exact `lerobot-eval` command and hardware so anyone can re-run:
+4 -1
View File
@@ -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` | | `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 | | `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 ### Schedule and checkpoints
+11
View File
@@ -242,6 +242,17 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
--repo-id=your_dataset --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: Alternatively, train MolmoAct2 with mean/std normalization:
```bash ```bash
+114 -118
View File
@@ -1,28 +1,29 @@
# Multi-GPU Training # Multi-GPU Training
This guide shows you how to train policies on multiple GPUs using [Hugging Face Accelerate](https://huggingface.co/docs/accelerate). LeRobot trains on multiple GPUs through [Hugging Face Accelerate](https://huggingface.co/docs/accelerate). Three data-parallel layouts are supported:
| Layout | What it does | Config |
| -------- | ------------------------------------------------------------- | ------------------------------------------------------- |
| **DDP** | Replicates the full model on every GPU | default on any multi-GPU launch |
| **FSDP** | Shards parameters, gradients, and optimizer state across GPUs | `--parallelism.dp_shard=N` |
| **HSDP** | Shards within groups of GPUs, replicates across groups | `--parallelism.dp_replicate=R --parallelism.dp_shard=S` |
## Installation ## Installation
`accelerate` is included in the `training` extra. Install it with: `accelerate` is included in the `training` extra:
```bash ```bash
pip install 'lerobot[training]' pip install 'lerobot[training]'
``` ```
## Training with Multiple GPUs ## Launching
You can launch training in two ways: Distributed training can be launched through both `torchrun` and `accelerate launch`. Accelerate is used as a plain launcher: it does not manage the training configuration, and every distributed training setting lives in LeRobot's own config system.
### Option 1: Without config (specify parameters directly) With `torchrun`:
You can specify all parameters directly in the command without running `accelerate config`:
```bash ```bash
accelerate launch \ torchrun --nproc-per-node=2 $(which lerobot-train) \
--multi_gpu \
--num_processes=2 \
$(which lerobot-train) \
--dataset.repo_id=${HF_USER}/my_dataset \ --dataset.repo_id=${HF_USER}/my_dataset \
--policy.type=act \ --policy.type=act \
--policy.repo_id=${HF_USER}/my_trained_policy \ --policy.repo_id=${HF_USER}/my_trained_policy \
@@ -31,32 +32,10 @@ accelerate launch \
--wandb.enable=true --wandb.enable=true
``` ```
**Key accelerate parameters:** With `accelerate launch` (as a plain launcher):
- `--multi_gpu`: Enable multi-GPU training
- `--num_processes=2`: Number of GPUs to use
- `--mixed_precision=fp16`: Use fp16 mixed precision (or `bf16` if supported)
### Option 2: Using accelerate config
If you prefer to save your configuration, you can optionally configure accelerate for your hardware setup by running:
```bash ```bash
accelerate config accelerate launch --num_processes=2 $(which lerobot-train) \
```
This interactive setup will ask you questions about your training environment (number of GPUs, mixed precision settings, etc.) and saves the configuration for future use. For a simple multi-GPU setup on a single machine, you can use these recommended settings:
- Compute environment: This machine
- Number of machines: 1
- Number of processes: (number of GPUs you want to use)
- GPU ids to use: (leave empty to use all)
- Mixed precision: fp16 or bf16 (recommended for faster training)
Then launch training with:
```bash
accelerate launch $(which lerobot-train) \
--dataset.repo_id=${HF_USER}/my_dataset \ --dataset.repo_id=${HF_USER}/my_dataset \
--policy.type=act \ --policy.type=act \
--policy.repo_id=${HF_USER}/my_trained_policy \ --policy.repo_id=${HF_USER}/my_trained_policy \
@@ -65,116 +44,133 @@ accelerate launch $(which lerobot-train) \
--wandb.enable=true --wandb.enable=true
``` ```
## How It Works With no `--parallelism.*` flags, a multi-process launch runs plain DDP. Multi-node runs use the standard `torchrun --nnodes/--node-rank/--rdzv-endpoint` flags (or `accelerate launch --num_machines/--machine_rank/--main_process_ip`).
When you launch training with accelerate: > [!WARNING]
> Accelerate's YAML config files (`accelerate launch --config_file some.yaml`, `accelerate config`) are not supported. They configure the engine through environment variables, bypassing LeRobot's configuration system, so `train_config.json` would no longer describe the settings a run actually used. `lerobot-train` therefore refuses to start when [accelerate environment variables](https://huggingface.co/docs/accelerate/usage_guides/fsdp) are set. Put the settings in `--parallelism.*` / `--accelerator.*` flags instead, or set `LEROBOT_ALLOW_ACCELERATE_ENV=1` to acknowledge the override and proceed anyway.
1. **Automatic detection**: LeRobot automatically detects if it's running under accelerate ## Batch semantics, learning rate, and steps
2. **Data distribution**: Your batch is automatically split across GPUs
3. **Gradient synchronization**: Gradients are synchronized across GPUs during backpropagation
4. **Single process logging**: Only the main process logs to wandb and saves checkpoints
## Learning Rate and Training Steps Scaling Each of the `dp_replicate × dp_shard` data-parallel workers loads its own `--batch_size` micro-batch every step, so one training step consumes `batch_size × dp_world_size` samples, and `× gradient_accumulation_steps` of those go into each optimizer update:
**Important:** LeRobot does **NOT** automatically scale learning rates or training steps based on the number of GPUs. This gives you full control over your training hyperparameters. ```
effective_batch_size = batch_size × dp_world_size × gradient_accumulation_steps
### Why No Automatic Scaling?
Many distributed training frameworks automatically scale the learning rate by the number of GPUs (e.g., `lr = base_lr × num_gpus`).
However, LeRobot keeps the learning rate exactly as you specify it.
### When and How to Scale
If you want to scale your hyperparameters when using multiple GPUs, you should do it manually:
**Learning Rate Scaling:**
```bash
# Example: 2 GPUs with linear LR scaling
# Base LR: 1e-4, with 2 GPUs -> 2e-4
accelerate launch --num_processes=2 $(which lerobot-train) \
--optimizer.lr=2e-4 \
--dataset.repo_id=lerobot/pusht \
--policy.type=act
``` ```
**Training Steps Scaling:** The training banner prints this factorization at startup. `--steps` counts loop steps (micro-batches per worker), not optimizer updates.
Since the effective batch size `bs` increases with multiple GPUs (batch_size × num_gpus), you may want to reduce the number of training steps proportionally: Gradient accumulation is a first-class flag:
```bash ```bash
# Example: 2 GPUs with effective batch size 2x larger torchrun --nproc-per-node=2 $(which lerobot-train) \
# Original: batch_size=8, steps=100000 --batch_size=8 --accelerator.gradient_accumulation.steps=4 ...
# With 2 GPUs: batch_size=8 (16 in total), steps=50000
accelerate launch --num_processes=2 $(which lerobot-train) \
--batch_size=8 \
--steps=50000 \
--dataset.repo_id=lerobot/pusht \
--policy.type=act
``` ```
## Training Large Models with FSDP **LeRobot does not auto-scale the learning rate or the number of steps** when the effective batch size grows. If you scale out and want equivalent training, please adjust manually, e.g. with 2 GPUs: double `--optimizer.lr` (linear scaling), or halve `--steps`.
DDP replicates the full model on every GPU, so a model that doesn't fit on one GPU won't fit under ## Sharded training (FSDP)
DDP either. For large models, use **FSDP** (Fully Sharded Data Parallel), which shards parameters,
gradients, and optimizer state across GPUs. See the [accelerate FSDP guide](https://huggingface.co/docs/accelerate/usage_guides/fsdp) for background.
An example on how to launch LeRobot training with FSDP across 4 GPUs (1 machine): If a model is too large to train with DDP, shard it with FSDP2:
```bash ```bash
accelerate launch --config_file fsdp.yaml --num_processes=4 $(which lerobot-train) \ torchrun --nproc-per-node=4 $(which lerobot-train) \
--dataset.repo_id=${HF_USER}/my_dataset \ --dataset.repo_id=${HF_USER}/my_dataset \
--policy.type=<your_policy> \ --policy.type=<your_policy> \
--parallelism.dp_shard=4 \
--accelerator.mixed_precision=bf16 \
--output_dir=outputs/train/my_policy_fsdp --output_dir=outputs/train/my_policy_fsdp
``` ```
A minimal `fsdp.yaml` (FSDP1; shards params/grads/optimizer — ZeRO-3-equivalent): `--parallelism.dp_shard=-1` shards over however many processes the launcher started.
```yaml ### Wrap units
compute_environment: LOCAL_MACHINE
distributed_type: FSDP FSDP shards the model in units (typically the repeated transformer block) and gathers one unit at a time during forward/backward. Policies declare their wrap units via `_fsdp_wrap_modules` on the policy class. For example, ACT declares `["ACTEncoderLayer", "ACTDecoderLayer"]` and FastWAM declares `["MoTLayer"]`. For a policy without a `_fsdp_wrap_modules` declaration, pass one of the flags below. You can specify the module class name explicitly, or use a size-based policy instead:
mixed_precision: bf16
num_machines: 1 ```bash
num_processes: 4 --accelerator.fsdp.wrap_modules='["MyTransformerBlock"]' # explicit class names
fsdp_config: --accelerator.fsdp.min_num_params=1000000 # or: wrap every submodule above 1M params
fsdp_version: 1
fsdp_sharding_strategy: FULL_SHARD # params + grads + optimizer (ZeRO-3)
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_transformer_layer_cls_to_wrap: <YourTransformerBlock> # repeated block class to shard
fsdp_use_orig_params: true # required: optimizer is built pre-prepare
fsdp_state_dict_type: FULL_STATE_DICT
``` ```
Set `fsdp_transformer_layer_cls_to_wrap` to your model's repeated transformer-block class so each If a policy doesn't declare `_fsdp_wrap_modules` and no `--accelerator.fsdp.wrap_modules` or `--accelerator.fsdp.min_num_params` is passed, the run fails at startup rather than silently wrapping only the root module (which would forfeit all sharding memory savings).
block is sharded as its own unit. `fsdp_use_orig_params: true` is required because LeRobot builds the
optimizer before `accelerator.prepare()`.
### FSDP checkpoints Other sharding settings:
LeRobot gathers the full state dict across all ranks and the main process writes it as a single - `--accelerator.fsdp.reshard_after_forward`: whether to keep each unit's parameters resident after forward.
`model.safetensors`, loadable as usual with `Policy.from_pretrained(...)`. Two things to look out for: - `--accelerator.fsdp.cpu_offload`: keeps parameters, gradients and optimizer states on CPU.
- `--accelerator.fsdp.ignored_modules`: a regex of module paths to keep unsharded.
- **Checkpoints store fp32 weights.** Under mixed precision (`bf16`/`fp16`) FSDP keeps an fp32 master ### HSDP
copy, and the checkpoint saves it (~2× the bf16 size on disk) so training can resume consistently
with the fp32 optimizer state; `from_pretrained` casts back to the policy dtype on load. FSDP-specific Hybrid Sharded Data Parallel: parameters, gradients and optimizer states are sharded across `dp_shard` ranks, and that sharding is replicated `dp_replicate` times. Parameter all-gathers and gradient reduce-scatters stay inside a shard group; only the all-reduce that synchronizes the replicas crosses between groups. The two degrees must multiply to the world size:
caveat: an fp32 checkpoint is materialized in full precision on the target device _before_ casting,
so loading it for inference on a tight GPU can OOM even when the bf16 model would fit — load on CPU ```bash
first, or cast `model.safetensors` to the deployment dtype offline. # 16 GPUs = 2 nodes × 8: shard within each node, replicate across nodes
- The sharded optimizer state is gathered into a full (world-size-independent) state dict and saved torchrun --nnodes=2 --nproc-per-node=8 ... $(which lerobot-train) \
alongside the model in the same `optimizer_state.safetensors` / `optimizer_param_groups.json` --parallelism.dp_replicate=2 --parallelism.dp_shard=8 ...
format as single-GPU training, so **resume-from-checkpoint is supported** with `--resume=true`. ```
Resume reshards both the model and the optimizer state to the _current_ FSDP topology, so you can
resume an FSDP checkpoint on a different number of GPUs. Note that the data sampler is only ## Checkpoints
sample-exact when the world size and batch size match the original run (a warning is logged
otherwise); the optimizer/model state itself is unaffected. Every checkpoint contains a `pretrained_model/` directory and a `training_state/` directory:
```text
005000/ # the training step at that checkpoint
├── pretrained_model/
│ ├── config.json # policy config
│ ├── train_config.json # the full training config
│ ├── model.safetensors # full weights (checkpoint_format ∈ {safetensors, safetensors_dcp}, or any non-sharded run)
│ ├── pytorch_model_fsdp_0/ # DCP weight shards (checkpoint_format ∈ {dcp, safetensors_dcp})
│ ├── policy_preprocessor.json # preprocessor config (when the run has a preprocessor)
│ ├── policy_preprocessor_step_*.safetensors # state of the stateful preprocessor steps
│ ├── policy_postprocessor.json # postprocessor config (when the run has a postprocessor)
│ └── policy_postprocessor_step_*.safetensors # state of the stateful postprocessor steps
└── training_state/
├── training_step.json # step counter, topology, and batch semantics
├── rng_state.safetensors # rng states
├── scheduler_state.json # scheduler state (when the run has a scheduler)
├── optimizer_state.safetensors # full optimizer state (non-sharded runs)
├── optimizer_param_groups.json # optimizer param groups (non-sharded runs)
└── optimizer_0/ # DCP optimizer shards (sharded runs)
```
During single-GPU or DDP training, the pipeline serializes each state dict into a single file: `model.safetensors` for the model and `optimizer_state.safetensors` for the optimizer.
During sharded training, the optimizer state is saved as DCP shards under `training_state/optimizer_0/`, and the layout of the model under `pretrained_model/` can be configured through `--checkpoint_format`:
| `--checkpoint_format` | Weights artifact | Use when |
| ------------------------- | -------------------------------------------- | --------------------------------------------------------------------- |
| `safetensors` _(default)_ | single `model.safetensors` only | you want every checkpoint immediately loadable with `from_pretrained` |
| `dcp` | `pytorch_model_fsdp_0/` shard directory only | gathering the full weights makes saves and resumes too slow |
| `safetensors_dcp` | both | you want fast resume _and_ immediately loadable checkpoints |
Two things to know about gathered (`safetensors`) checkpoints from sharded runs:
- **They store fp32 weights.** Under mixed precision training, FSDP keeps an fp32 master copy, and the checkpoint saves the master copy to make sure training resumes consistently.
- The gather is collective (all ranks participate) but only the main process writes.
### Converting DCP checkpoints
`lerobot-convert-dcp` merges a DCP shard directory into a regular `model.safetensors`, offline and without GPUs:
```bash
lerobot-convert-dcp --checkpoint_dir=outputs/train/run/checkpoints/005000
lerobot-convert-dcp --checkpoint_dir=... --delete_dcp=true --push_to_hub=${HF_USER}/my_policy
```
`--push_to_hub` publishes the converted directory as a model repo.
### Resuming
Resume with `--resume=true --config_path=.../checkpoints/last/pretrained_model/train_config.json`. Resuming from a DCP checkpoint supports resharding the model and optimizer state to the _current_ topology, which means you can resume with a different `dp_replicate/dp_shard` split. The data sampler can always resume at the right epoch and offset, but is only _sample-exact_ when the world size and batch size match the original run (a warning is logged otherwise).
> [!NOTE]
> FSDP checkpoints written by LeRobot 0.6.x and earlier used a different on-disk layout (a gathered full optimizer state) and **cannot be resumed**.
## Notes ## Notes
- The `--policy.use_amp` flag in `lerobot-train` is only used when **not** running with accelerate. When using accelerate, mixed precision is controlled by accelerate's configuration. - Checkpoint saves and end-of-training publishes are collective (every rank enters them). Gathered weights, sidecar files and Hub uploads are written by the main process alone.
- Training logs, checkpoints, and hub uploads are only done by the main process to avoid conflicts. Non-main processes have console logging disabled to prevent duplicate output. - Metrics are reduced across ranks before logging: losses are averaged, and `samples/s` reports cluster-wide throughput.
- The effective batch size is `batch_size × num_gpus`. If you use 4 GPUs with `--batch_size=8`, your effective batch size is 32. - Learning-rate scheduling is stepped once per training step regardless of the number of processes (`step_scheduler_with_optimizer=False` is baked in).
- Learning rate scheduling is handled correctly across multiple processes—LeRobot sets `step_scheduler_with_optimizer=False` to prevent accelerate from adjusting scheduler steps based on the number of processes.
- When saving or pushing models, LeRobot automatically unwraps the model from accelerate's distributed wrapper to ensure compatibility.
- WandB integration automatically initializes only on the main process, preventing multiple runs from being created.
For more advanced configurations and troubleshooting, see the [Accelerate documentation](https://huggingface.co/docs/accelerate). If you want to learn more about how to train on a large number of GPUs, checkout this awesome guide: [Ultrascale Playbook](https://huggingface.co/spaces/nanotron/ultrascale-playbook). For background on the underlying machinery, see the [Accelerate FSDP guide](https://huggingface.co/docs/accelerate/usage_guides/fsdp). To go deeper on large-scale training, check out the [Ultrascale Playbook](https://huggingface.co/spaces/nanotron/ultrascale-playbook).
+11
View File
@@ -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"}'`. 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 ### 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. The same finetune with the VLM frozen: less memory, at some cost in success rate. Swap `--dataset.repo_id` for your own dataset.
+19
View File
@@ -2,6 +2,25 @@
https://diffusion-policy.cs.columbia.edu 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 ## Citation
```bibtex ```bibtex
+16
View File
@@ -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 ## Citation
If you use this work, please cite both **OpenPI** and the π₀.₅ paper: If you use this work, please cite both **OpenPI** and the π₀.₅ paper:
+12
View File
@@ -40,3 +40,15 @@ lerobot-eval \
``` ```
However, in most cases, presence of an accelerator is detected automatically and `policy.device` parameter can be omitted from CLI commands. However, in most cases, presence of an accelerator is detected automatically and `policy.device` parameter can be omitted from CLI commands.
## Mixed precision
Training precision is owned by `--accelerator.mixed_precision`, which accepts `no` (default) and `bf16`:
```bash
lerobot-train \
--policy.type=act \
--accelerator.mixed_precision=bf16 ...
```
`bf16` requires an accelerator that supports it.
+287
View File
@@ -0,0 +1,287 @@
# Writing docstrings
LeRobot's API reference is generated directly from the docstrings in `src/lerobot/`. A docstring is not a
comment — it is the published documentation for that object, and the format below is what the renderer and
the CI checks parse.
This page is the contract. If you are adding or editing anything public in `src/lerobot/`, follow it.
> [!IMPORTANT]
> **An undocumented public method is an invisible one.** `[[autodoc]]` silently skips members that have no
> docstring — no warning, no error, it simply does not appear on the rendered page. Coverage and
> API-reference completeness are the same problem.
## The format in one example
Google section headers, Hugging Face type formatting. Both, not one or the other.
````python
def send_action(self, action: RobotAction, rate_hz: float = 30.0) -> RobotAction:
"""Command the robot to move to a target joint configuration.
Values are clipped by the configured maximum relative target before reaching the motors, so the
returned action may differ from the requested one.
Args:
action (`dict[str, float]`):
Target values keyed by motor name, e.g. `{"shoulder_pan.pos": 0.0}`. Keys must match the
robot's action features.
rate_hz (`float`, *optional*, defaults to `30.0`):
Control loop frequency.
Returns:
`dict[str, float]`: The action actually written to the motors after safety clipping.
Raises:
DeviceNotConnectedError: If the robot has not been connected.
Example:
```python
>>> from lerobot.robots.so_follower import SO101Follower, SO101FollowerConfig
>>> robot = SO101Follower(SO101FollowerConfig(port="/dev/ttyACM0")) # doctest: +SKIP
>>> robot.connect() # doctest: +SKIP
>>> robot.send_action({"shoulder_pan.pos": 0.0}) # doctest: +SKIP
```
"""
````
Cross-references are omitted from the examples on this page — see [Cross-references](#cross-references) for
their syntax and why they cannot be shown inside a code block.
## Rules
### Sections
`Args:` · `Returns:` · `Raises:` · `Yields:` · `Example:` · `Note:`
In that order. No other section headers. A one-line summary comes first, then an optional free-form
description, then the sections.
### The `Args:` line is machine-parsed
```
name (`type`, *optional*, defaults to `X`):
Description, indented on its own line.
```
The `*optional*, defaults to` clause is **checked against the real signature default** by
`make check-docstrings`. It is not decorative — if you write a default that has drifted from the code, CI
fails. Omit the clause entirely for required parameters:
```python
Args:
port (`str`):
Serial port the arm is connected to, e.g. `/dev/ttyACM0`.
max_relative_target (`float | dict[str, float]`, *optional*):
Caps the magnitude of the relative positional target vector. `None` disables clipping.
use_degrees (`bool`, *optional*, defaults to `True`):
Keep `True` for backward compatibility with existing policies and datasets.
```
Types go in backticks. Use `*optional*` with no `defaults to` when the default is `None` or is otherwise not
worth restating.
### `Returns:` is type-first
One indented line, type first, then a colon, then the description:
```python
Returns:
`dict[str, float]`: The action actually written to the motors after safety clipping.
```
`Yields:` takes the same shape.
### `**Attributes**:`, never `Attributes:`
doc-builder parses a bare `Attributes:` as a **synonym for `Parameters:`**, so your attributes get rendered
as constructor arguments. This is silent and wrong. Whenever the attributes differ from the constructor
parameters, use the bold form with a `--` separator:
```python
class Robot(abc.ABC):
"""The base abstract class for all LeRobot-compatible robots.
**Attributes**:
- **config_class** (`type[RobotConfig]`) -- The expected configuration class for this robot.
- **name** (`str`) -- The unique robot name used to identify this robot type.
"""
```
Note `--`, not `:`.
### Cross-references
Use doc-builder's bracket syntax: a square-bracketed backtick-quoted path. **Sphinx roles (`:pymeth:`,
`:pyattr:`) are not supported** and render as literal text on the page.
| Want | Write |
| ---------------------------- | ----------------------------------- |
| Class in the main package | &#91;`Robot`&#93; |
| Method, show the full path | &#91;`Robot.connect`&#93; |
| Method, show the bare name | &#91;`~Robot.connect`&#93; |
| Nested path | &#91;`~robots.Robot.connect`&#93; |
| Object in another HF library | &#91;`~accelerate.Accelerator`&#93; |
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 &#91;`~module.Class.method`&#93;.
- [ ] 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.
+80 -17
View File
@@ -346,6 +346,7 @@ lerobot-record="lerobot.scripts.lerobot_record:main"
lerobot-replay="lerobot.scripts.lerobot_replay:main" lerobot-replay="lerobot.scripts.lerobot_replay:main"
lerobot-setup-motors="lerobot.scripts.lerobot_setup_motors:main" lerobot-setup-motors="lerobot.scripts.lerobot_setup_motors:main"
lerobot-teleoperate="lerobot.scripts.lerobot_teleoperate:main" lerobot-teleoperate="lerobot.scripts.lerobot_teleoperate:main"
lerobot-convert-dcp="lerobot.scripts.lerobot_convert_dcp:main"
lerobot-eval="lerobot.scripts.lerobot_eval:main" lerobot-eval="lerobot.scripts.lerobot_eval:main"
lerobot-train="lerobot.scripts.lerobot_train:main" lerobot-train="lerobot.scripts.lerobot_train:main"
lerobot-train-tokenizer="lerobot.scripts.lerobot_train_tokenizer:main" lerobot-train-tokenizer="lerobot.scripts.lerobot_train_tokenizer:main"
@@ -400,7 +401,7 @@ exclude = ["tests/artifacts/**/*.safetensors", "*_pb2.py", "*_pb2_grpc.py"]
# N: pep8-naming # N: pep8-naming
# TODO: Uncomment rules when ready to use # TODO: Uncomment rules when ready to use
select = [ select = [
"E", "W", "F", "I", "B", "C4", "T20", "N", "UP", "SIM" #, "A", "S", "D", "RUF" "E", "W", "F", "I", "B", "C4", "T20", "N", "UP", "SIM", "D" #, "A", "S", "RUF"
] ]
ignore = [ ignore = [
"E501", # Line too long "E501", # Line too long
@@ -410,9 +411,53 @@ ignore = [
] ]
[tool.ruff.lint.per-file-ignores] [tool.ruff.lint.per-file-ignores]
"__init__.py" = ["F401", "F403", "E402"] "__init__.py" = ["F401", "F403", "E402", "D104"]
# E402: conditional-import guards (TYPE_CHECKING / is_package_available) must precede the imports they protect # 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"] "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/rl/**" = ["D"]
"src/lerobot/robots/**" = ["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"]
# Package root: two one-line docstring fixes land with the docstring PR.
"src/lerobot/__init__.py" = ["D"]
"src/lerobot/__version__.py" = ["D"]
[tool.ruff.lint.isort] [tool.ruff.lint.isort]
combine-as-imports = true combine-as-imports = true
known-first-party = ["lerobot"] known-first-party = ["lerobot"]
@@ -456,25 +501,34 @@ default.extend-ignore-identifiers-re = [
"seperated_timestep", "seperated_timestep",
] ]
# TODO: Uncomment when ready to use # Docstring coverage gate. `fail-under` is a RATCHET, not a target: it is set just below the currently
# [tool.interrogate] # measured coverage so it passes today, and is raised in the same PR that documents a module. Never set it
# ignore-init-module = true # to a value that fails on main. The destination is 100; see docs/source/writing_docstrings.mdx.
# ignore-init-method = true [tool.interrogate]
# ignore-nested-functions = false ignore-init-module = true
# ignore-magic = false ignore-init-method = true
# ignore-semiprivate = false ignore-nested-functions = false
# ignore-private = false ignore-magic = false
# ignore-property-decorators = false ignore-semiprivate = false
# ignore-module = false ignore-private = false
# ignore-setters = false ignore-property-decorators = false
# fail-under = 80 ignore-module = false
# output-format = "term-missing" ignore-setters = false
# color = true fail-under = 52
# paths = ["src/lerobot"] output-format = "term-missing"
color = true
paths = ["src/lerobot"]
exclude = ["src/lerobot/policies/molmoact2/molmoact2_hf_model"]
# TODO: Enable mypy gradually module by module across multiple PRs # 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 # Uncomment [tool.mypy] first, then uncomment individual module overrides as they get proper type annotations
[tool.pytest.ini_options]
markers = [
"multigpu: distributed tests needing 2-4 GPUs (CI: docker_publish.yml lane)",
"multigpu_heavy: 8-GPU sweeps and soak tests; never run in CI",
]
[tool.mypy] [tool.mypy]
python_version = "3.12" python_version = "3.12"
ignore_missing_imports = true ignore_missing_imports = true
@@ -521,6 +575,15 @@ disallow_untyped_defs = true
disallow_incomplete_defs = true disallow_incomplete_defs = true
check_untyped_defs = true check_untyped_defs = true
[[tool.mypy.overrides]]
module = "lerobot.distributed.*"
ignore_errors = false
# extra strictness for the distributed engine
disallow_untyped_defs = true
disallow_incomplete_defs = true
check_untyped_defs = true
[[tool.mypy.overrides]] [[tool.mypy.overrides]]
module = "lerobot.optim.*" module = "lerobot.optim.*"
ignore_errors = false ignore_errors = false
+604 -174
View File
@@ -13,16 +13,41 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from pathlib import Path """Training-output persistence: checkpoints, two-phase resume, and hub publishing.
from huggingface_hub import HfApi, snapshot_download Rank discipline: every function here that can
contain a collective is documented as such and must run on ALL ranks; rank-0-only file writes
sit under one grouped ``is_main_process()`` gate per contiguous region, placed below all
collectives. The leaf save/load helpers carry no rank gates of their own — the exception is
``PreTrainedPolicy._save_pretrained``, whose gate is internal because its collective gather and
its writes live in the same method.
"""
import logging
from importlib.resources import files
from pathlib import Path
from tempfile import TemporaryDirectory
from typing import TYPE_CHECKING, Any
import torch.distributed as dist
from huggingface_hub import HfApi, ModelCard, ModelCardData, snapshot_download
from torch.optim import Optimizer from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler from torch.optim.lr_scheduler import LRScheduler
from lerobot.__version__ import __version__
from lerobot.configs.policies import PreTrainedConfig
from lerobot.configs.rewards import RewardModelConfig
from lerobot.configs.train import TrainPipelineConfig from lerobot.configs.train import TrainPipelineConfig
from lerobot.distributed.checkpoint import (
is_sharded_module,
load_sharded_model,
load_sharded_optimizer,
save_sharded_model,
save_sharded_optimizer,
)
from lerobot.distributed.utils import is_main_process
from lerobot.optim import ( from lerobot.optim import (
load_optimizer_state, load_optimizer_state,
load_optimizer_state_dict,
load_scheduler_state, load_scheduler_state,
save_optimizer_state, save_optimizer_state,
save_scheduler_state, save_scheduler_state,
@@ -40,14 +65,39 @@ from lerobot.utils.hub import find_latest_hub_checkpoint
from lerobot.utils.io_utils import load_json, write_json from lerobot.utils.io_utils import load_json, write_json
from lerobot.utils.random_utils import load_rng_state, save_rng_state from lerobot.utils.random_utils import load_rng_state, save_rng_state
if TYPE_CHECKING:
from accelerate import Accelerator
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
from lerobot.rewards.pretrained import PreTrainedRewardModel
def get_step_identifier(step: int, total_steps: int) -> str: def get_step_identifier(step: int, total_steps: int) -> str:
"""Format a step number as the zero-padded identifier used for checkpoint directory names.
Args:
step (int): The training step to format.
total_steps (int): The total number of training steps; sets the padding width
(minimum 6 digits).
Returns:
str: The zero-padded step identifier, e.g. `"005000"`.
"""
num_digits = max(6, len(str(total_steps))) num_digits = max(6, len(str(total_steps)))
return f"{step:0{num_digits}d}" return f"{step:0{num_digits}d}"
def get_step_checkpoint_dir(output_dir: Path, total_steps: int, step: int) -> Path: def get_step_checkpoint_dir(output_dir: Path, total_steps: int, step: int) -> Path:
"""Returns the checkpoint sub-directory corresponding to the step number.""" """Returns the checkpoint sub-directory corresponding to the step number.
Args:
output_dir (Path): The training run's output directory.
total_steps (int): The total number of training steps; sets the identifier padding.
step (int): The training step of the checkpoint.
Returns:
Path: The checkpoint step directory, `output_dir/checkpoints/<step-identifier>`.
"""
step_identifier = get_step_identifier(step, total_steps) step_identifier = get_step_identifier(step, total_steps)
return output_dir / CHECKPOINTS_DIR / step_identifier return output_dir / CHECKPOINTS_DIR / step_identifier
@@ -63,37 +113,15 @@ def should_save_checkpoint(step: int, save_freq: int, total_steps: int) -> bool:
return (save_freq > 0 and step % save_freq == 0) or step == total_steps return (save_freq > 0 and step % save_freq == 0) or step == total_steps
def save_training_step( def update_last_checkpoint(checkpoint_dir: Path) -> None:
step: int, save_dir: Path, num_processes: int | None = None, batch_size: int | None = None """Point the `last` symlink in the checkpoints directory at the given checkpoint.
) -> None:
state: dict = {"step": step}
# num_processes and batch_size are recorded so a resumed run can detect a changed world size or
# batch size: the sampler's resume offset is computed from the (num_processes, batch_size) that
# produced `step`, since both scale how many sampler positions a step consumes (see
# compute_sampler_state).
if num_processes is not None:
state["num_processes"] = num_processes
if batch_size is not None:
state["batch_size"] = batch_size
write_json(state, save_dir / TRAINING_STEP)
Any existing `last` symlink is replaced. The link target is relative to the checkpoints
directory, so the tree stays valid when the run directory is moved.
def load_training_step(save_dir: Path) -> int: Args:
training_step = load_json(save_dir / TRAINING_STEP) checkpoint_dir (Path): The checkpoint step directory the `last` link should target.
return training_step["step"] """
def load_training_num_processes(checkpoint_dir: Path) -> int | None:
"""World size recorded at checkpoint time, or None for checkpoints written before it was stored."""
return load_json(checkpoint_dir / TRAINING_STATE_DIR / TRAINING_STEP).get("num_processes")
def load_training_batch_size(checkpoint_dir: Path) -> int | None:
"""Per-process batch size recorded at checkpoint time, or None for older checkpoints."""
return load_json(checkpoint_dir / TRAINING_STATE_DIR / TRAINING_STEP).get("batch_size")
def update_last_checkpoint(checkpoint_dir: Path) -> Path:
last_checkpoint_dir = checkpoint_dir.parent / LAST_CHECKPOINT_LINK last_checkpoint_dir = checkpoint_dir.parent / LAST_CHECKPOINT_LINK
if last_checkpoint_dir.is_symlink(): if last_checkpoint_dir.is_symlink():
last_checkpoint_dir.unlink() last_checkpoint_dir.unlink()
@@ -101,6 +129,68 @@ def update_last_checkpoint(checkpoint_dir: Path) -> Path:
last_checkpoint_dir.symlink_to(relative_target) last_checkpoint_dir.symlink_to(relative_target)
# ---------------------------------------------------------------------------------------------
# training_step.json
# ---------------------------------------------------------------------------------------------
def save_training_metadata(step: int, save_dir: Path, cfg: TrainPipelineConfig) -> None:
"""Record the step counter plus everything a resume needs to reason about topology changes.
`step` counts loop iterations (= micro-batches), so
the sampler resume offset is `step x batch_size x dp_world_size` with no grad-accum factor.
`grad_accum_steps` and the parallelism snapshot are recorded so a resume can warn precisely
when the optimizer-update cadence or the sharding topology changed.
Args:
step (int): The training step (micro-batch counter) to record.
save_dir (Path): The `training_state/` directory to write `training_step.json` into.
cfg (TrainPipelineConfig): The training config whose batch size, gradient-accumulation,
and parallelism settings are snapshotted alongside the step.
"""
state: dict[str, Any] = {
"step": step,
"dp_world_size": cfg.parallelism.dp_world_size,
"batch_size": cfg.batch_size,
"grad_accum_steps": cfg.accelerator.gradient_accumulation.steps,
"parallelism": {
"dp_replicate": cfg.parallelism.dp_replicate,
"dp_shard": cfg.parallelism.dp_shard,
"ring_degree": cfg.parallelism.context_parallel.ring_degree,
"ulysses_degree": cfg.parallelism.context_parallel.ulysses_degree,
},
}
write_json(state, save_dir / TRAINING_STEP)
def load_training_metadata(training_state_dir: Path) -> dict[str, Any]:
"""Read everything `save_training_metadata` recorded, in a single pass.
Every key is always present: fields a checkpoint predates come back as None, so a caller
reading `metadata["batch_size"]` gets a KeyError on a typo rather than a silent None.
Args:
training_state_dir (Path): The checkpoint's `training_state/` directory.
Returns:
dict[str, Any]: `step` plus the `dp_world_size`, `batch_size`, `grad_accum_steps` and
`parallelism` snapshot recorded alongside it (None where not recorded).
"""
state = load_json(training_state_dir / TRAINING_STEP)
return {
"step": int(state["step"]),
"dp_world_size": state.get("dp_world_size", state.get("num_processes")),
"batch_size": state.get("batch_size"),
"grad_accum_steps": state.get("grad_accum_steps"),
"parallelism": state.get("parallelism"),
}
# ---------------------------------------------------------------------------------------------
# Checkpoint save
# ---------------------------------------------------------------------------------------------
def save_checkpoint( def save_checkpoint(
checkpoint_dir: Path, checkpoint_dir: Path,
step: int, step: int,
@@ -110,192 +200,301 @@ def save_checkpoint(
scheduler: LRScheduler | None = None, scheduler: LRScheduler | None = None,
preprocessor: PolicyProcessorPipeline | None = None, preprocessor: PolicyProcessorPipeline | None = None,
postprocessor: PolicyProcessorPipeline | None = None, postprocessor: PolicyProcessorPipeline | None = None,
num_processes: int | None = None, accelerator: "Accelerator | None" = None,
batch_size: int | None = None,
model_state_dict: dict | None = None,
optim_state_dict: dict | None = None,
) -> None: ) -> None:
"""This function creates the following directory structure: """This function creates the following directory structure:
005000/ # training step at checkpoint 005000/ # training step at checkpoint
├── pretrained_model/ ├── pretrained_model/
│ ├── config.json # policy config │ ├── config.json # policy config
│ ├── model.safetensors # policy weights │ ├── model.safetensors # policy weights (checkpoint_format ∈ {safetensors, safetensors_dcp}, or any non-sharded run)
│ ├── pytorch_model_fsdp_0/ # DCP model shards (checkpoint_format ∈ {dcp, safetensors_dcp})
│ ├── train_config.json # train config │ ├── train_config.json # train config
│ ├── processor.json # processor config (if preprocessor provided) │ ├── policy_preprocessor.json # preprocessor config (if preprocessor provided)
── step_*.safetensors # processor state files (if any) ── policy_preprocessor_step_*.safetensors # state of the stateful preprocessor steps
│ ├── policy_postprocessor.json # postprocessor config (if postprocessor provided)
│ └── policy_postprocessor_step_*.safetensors # state of the stateful postprocessor steps
└── training_state/ └── training_state/
├── optimizer_param_groups.json # optimizer param groups ├── optimizer_param_groups.json # optimizer param groups (non-sharded runs)
├── optimizer_state.safetensors # optimizer state ├── optimizer_state.safetensors # optimizer state (non-sharded runs)
├── optimizer_0/ # DCP optimizer shards (sharded runs)
├── rng_state.safetensors # rng states ├── rng_state.safetensors # rng states
├── scheduler_state.json # scheduler state ├── scheduler_state.json # scheduler state (if scheduler provided)
└── training_step.json # training step └── training_step.json # training step + dp_world_size/batch_size/grad_accum + topology
Collective: MUST be called on every rank. Rank-0-only writes are gated internally, so the
call site needs no rank branches.
Args: Args:
cfg (TrainPipelineConfig): The training config used for this run. checkpoint_dir (Path): The checkpoint step directory to write (e.g. `.../checkpoints/005000`).
step (int): The training step at that checkpoint. step (int): The training step at that checkpoint.
cfg (TrainPipelineConfig): The training config used for this run.
policy (PreTrainedPolicy): The policy to save. policy (PreTrainedPolicy): The policy to save.
optimizer (Optimizer | None, optional): The optimizer to save the state from. Defaults to None. optimizer (Optimizer): The optimizer to save the state from.
scheduler (LRScheduler | None, optional): The scheduler to save the state from. Defaults to None. scheduler (LRScheduler | None, optional): The scheduler to save the state from. Defaults to None.
preprocessor: The preprocessor/pipeline to save. Defaults to None. preprocessor (PolicyProcessorPipeline | None, optional): The preprocessor/pipeline to save.
postprocessor: The postprocessor/pipeline to save. Defaults to None.
num_processes (int | None, optional): Distributed world size to record for sample-exact
resume. Defaults to None (not recorded).
batch_size (int | None, optional): Per-process batch size to record for sample-exact
resume. Defaults to None (not recorded).
model_state_dict: Pre-gathered full (unsharded) model state dict. Required under FSDP,
where `policy.state_dict()` would return sharded tensors; the caller gathers it via a
cross-rank collective and passes it here so rank 0 can write it directly. It holds
FSDP's fp32 master weights and is saved as-is (the loader casts to the policy dtype on
read). When None (DDP / single-GPU), the model is saved the normal way. Defaults to None.
optim_state_dict: Pre-gathered full (unsharded) optimizer state dict. Required under FSDP
(gathered alongside `model_state_dict` via `gather_fsdp_state_dicts`); saved in the same
safetensors format as the single-GPU path. When None, `optimizer.state_dict()` is used.
Defaults to None. Defaults to None.
postprocessor (PolicyProcessorPipeline | None, optional): The postprocessor/pipeline to save.
Defaults to None.
accelerator (Accelerator | None, optional): The accelerator the policy was prepared with;
used to unwrap the model and required on sharded runs, where it owns the DCP save
channels. Defaults to None (plain single-process saves).
""" """
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
policy.save_pretrained(pretrained_dir, state_dict=model_state_dict) fmt = cfg.checkpoint_format
cfg.save_pretrained(pretrained_dir) policy_to_save = accelerator.unwrap_model(policy) if accelerator is not None else policy
sharded = is_sharded_module(policy_to_save)
# -- model artifact(s): the two collective-capable calls ----------------------------------
if cfg.peft is not None: if cfg.peft is not None:
# When using PEFT, policy.save_pretrained will only write the adapter weights + config, not the # PeftModel.save_pretrained is an external API with no internal rank gate, and the
# policy config which we need for loading the model. In this case we'll write it ourselves. # adapters are replicated (PEFT x sharded is rejected at validation): main rank writes.
policy.config.save_pretrained(pretrained_dir) if is_main_process():
if preprocessor is not None: policy_to_save.save_pretrained(pretrained_dir)
preprocessor.save_pretrained(pretrained_dir) elif fmt.wants_safetensors or not sharded:
if postprocessor is not None: # Collective when sharded (full gather); writes happen on the main process only in all
postprocessor.save_pretrained(pretrained_dir) # multi-rank layouts (the gate lives inside _save_pretrained, next to its collective gather).
policy_to_save.save_pretrained(pretrained_dir)
if fmt.wants_dcp and sharded:
save_sharded_model(accelerator, policy_to_save, pretrained_dir)
# -- sidecar configs: ONE gate for the whole contiguous rank-0-only region ----------------
if is_main_process():
if fmt.wants_dcp and not fmt.wants_safetensors:
# save_pretrained did not run: keep the DCP-only checkpoint self-describing.
policy_to_save.config.save_pretrained(pretrained_dir)
cfg.save_pretrained(pretrained_dir)
if cfg.peft is not None:
# PEFT's save_pretrained writes only adapter weights + config; the policy config
# needed to reload the base model is written explicitly.
policy_to_save.config.save_pretrained(pretrained_dir)
if preprocessor is not None:
preprocessor.save_pretrained(pretrained_dir)
if postprocessor is not None:
postprocessor.save_pretrained(pretrained_dir)
save_training_state( save_training_state(
checkpoint_dir, checkpoint_dir, step, cfg, optimizer, scheduler, accelerator, sharded=sharded, model=policy_to_save
step,
optimizer,
scheduler,
num_processes=num_processes,
batch_size=batch_size,
optim_state_dict=optim_state_dict,
) )
if accelerator is not None:
accelerator.wait_for_everyone()
def save_training_state( def save_training_state(
checkpoint_dir: Path, checkpoint_dir: Path,
train_step: int, step: int,
optimizer: Optimizer | None = None, cfg: TrainPipelineConfig,
optimizer: Optimizer | dict[str, Optimizer] | None = None,
scheduler: LRScheduler | None = None, scheduler: LRScheduler | None = None,
num_processes: int | None = None, accelerator: "Accelerator | None" = None,
batch_size: int | None = None, *,
optim_state_dict: dict | None = None, sharded: bool = False,
model: PreTrainedPolicy | None = None,
) -> None: ) -> None:
""" """Write training_state/. Collective under sharding: call on every rank.
Saves the training step, optimizer state, scheduler state, and rng state.
Args: Args:
save_dir (Path): The directory to save artifacts to. checkpoint_dir (Path): The checkpoint step directory; `training_state/` is created inside it.
train_step (int): Current training step. step (int): The training step at that checkpoint.
optimizer (Optimizer | None, optional): The optimizer from which to save the state_dict. cfg (TrainPipelineConfig): The training config used for this run (its topology and
accumulation settings are recorded in `training_step.json`).
optimizer (Optimizer | dict[str, Optimizer] | None, optional): The optimizer(s) to save
the state from. Defaults to None.
scheduler (LRScheduler | None, optional): The scheduler to save the state from.
Defaults to None. Defaults to None.
scheduler (LRScheduler | None, optional): The scheduler from which to save the state_dict. accelerator (Accelerator | None, optional): Required when `sharded` is True — it owns
Defaults to None. the DCP optimizer save channel. Defaults to None.
num_processes (int | None, optional): Distributed world size to record. Defaults to None. sharded (bool): The model's sharding state, computed once in `save_checkpoint` and
batch_size (int | None, optional): Per-process batch size to record. Defaults to None. threaded here so the two sites cannot disagree. Defaults to False.
optim_state_dict: Pre-gathered full optimizer state dict (for FSDP). Saved instead of model (PreTrainedPolicy | None, optional): Required only for the sharded optimizer
`optimizer.state_dict()` when provided. Defaults to None. channel: torch's optimizer DCP APIs are model-coupled (the state dict is keyed by
model FQNs), so accelerate's `save_fsdp_optimizer` needs the sharded module
alongside the optimizer. Defaults to None.
""" """
save_dir = checkpoint_dir / TRAINING_STATE_DIR save_dir = checkpoint_dir / TRAINING_STATE_DIR
# All ranks: the directory must exist before the DCP optimizer collective writes into it
# (exist_ok makes the concurrent mkdir race-free on shared filesystems).
save_dir.mkdir(parents=True, exist_ok=True) save_dir.mkdir(parents=True, exist_ok=True)
save_training_step(train_step, save_dir, num_processes=num_processes, batch_size=batch_size)
save_rng_state(save_dir) if optimizer is not None and sharded:
if optimizer is not None: if accelerator is None or model is None:
save_optimizer_state(optimizer, save_dir, optim_state_dict=optim_state_dict) raise ValueError("Saving a sharded optimizer state requires the accelerator and model.")
if scheduler is not None: # Collective — all ranks write their DCP shards into optimizer_0/.
save_scheduler_state(scheduler, save_dir) save_sharded_optimizer(accelerator, optimizer, model, save_dir)
if is_main_process(): # ONE grouped gate for the whole rank-0-only region
save_training_metadata(step, save_dir, cfg)
save_rng_state(save_dir)
if scheduler is not None:
save_scheduler_state(scheduler, save_dir)
if optimizer is not None and not sharded:
save_optimizer_state(optimizer, save_dir)
def load_training_state( # ---------------------------------------------------------------------------------------------
checkpoint_dir: Path, optimizer: Optimizer, scheduler: LRScheduler | None, load_optimizer: bool = True # Two-phase resume
) -> tuple[int, Optimizer, LRScheduler | None]: # ---------------------------------------------------------------------------------------------
"""
Loads the training step, optimizer state, scheduler state, and rng state.
This is used to resume a training run. def resume_before_prepare(cfg: TrainPipelineConfig) -> int:
"""Phase 1 — before `accelerator.prepare()`: restore RNG and return the step counter.
Pure loaders only. The sampler resume offset is *derived* from the returned step inside the
dataloader factory, and everything bound to sharded objects (model DCP shards, optimizer,
scheduler) loads in `resume_after_prepare`.
Args: Args:
checkpoint_dir (Path): The checkpoint directory. Should contain a 'training_state' dir. cfg (TrainPipelineConfig): The resumed training config; `cfg.checkpoint_path` locates
optimizer (Optimizer): The optimizer to load the state_dict to. the checkpoint to restore from.
scheduler (LRScheduler | None): The scheduler to load the state_dict to (can be None).
load_optimizer (bool, optional): Whether to load the optimizer state from disk. Defaults to Returns:
True. Set to False under FSDP, where the sharded optimizer state must be loaded after int: The training step recorded in the checkpoint (micro-batch counter).
`accelerator.prepare()` via `load_fsdp_optimizer_state` (the optimizer is returned
untouched here).
Raises: Raises:
NotADirectoryError: If 'checkpoint_dir' doesn't contain a 'training_state' dir NotADirectoryError: If the checkpoint has no `training_state/` directory.
ValueError: If the resumed topology crosses the sharded/non-sharded boundary relative
Returns: to the one recorded in the checkpoint.
tuple[int, Optimizer, LRScheduler | None]: training step, optimizer and scheduler with their
state_dict loaded.
""" """
training_state_dir = checkpoint_dir / TRAINING_STATE_DIR training_state_dir = cfg.checkpoint_path / TRAINING_STATE_DIR
if not training_state_dir.is_dir(): if not training_state_dir.is_dir():
raise NotADirectoryError(training_state_dir) raise NotADirectoryError(training_state_dir)
metadata = load_training_metadata(training_state_dir)
_guard_resume_changes(cfg, metadata)
load_rng_state(training_state_dir) load_rng_state(training_state_dir)
step = load_training_step(training_state_dir) return metadata["step"]
if load_optimizer:
optimizer = load_optimizer_state(optimizer, training_state_dir)
def _guard_resume_changes(cfg: TrainPipelineConfig, metadata: dict[str, Any]) -> None:
"""Check the resumed run settings against the ones recorded in the checkpoint.
Two tiers, both driven by the checkpoint's recorded parallelism snapshot:
- **Hard error** when the resume crosses the sharded/non-sharded boundary in either
direction: the checkpoint's training-state artifacts only support resuming on the same
kind of topology (resharding works across sizes, not across kinds). Checkpoints without
a recorded snapshot skip this check.
- **One warning** naming every other recorded setting that differs — those changes are
legal (DCP reshards weights and optimizer state across topologies and the sampler offset
adapts), but a changed ``grad_accum_steps`` shifts the optimizer-update cadence, so the
resume says precisely what differs. The sampler-exactness warnings
(``dp_world_size``/``batch_size``) live with the sampler math in the dataloader factory.
Args:
cfg (TrainPipelineConfig): The resumed training config, compared against the settings
recorded in the checkpoint.
metadata (dict[str, Any]): The checkpoint's recorded training metadata, as returned by
`load_training_metadata`.
Raises:
ValueError: If the checkpoint records a sharded topology and the resumed run is
non-sharded, or vice versa.
"""
snapshot = metadata["parallelism"]
if snapshot is not None:
recorded_sharded = (
snapshot.get("dp_shard", 1) != 1
or snapshot.get("ring_degree", 1) * snapshot.get("ulysses_degree", 1) > 1
)
if recorded_sharded != cfg.parallelism.is_sharded:
raise ValueError(
f"Cannot resume: the checkpoint was written with a "
f"{'sharded' if recorded_sharded else 'non-sharded'} topology "
f"(dp_replicate={snapshot.get('dp_replicate')}, dp_shard={snapshot.get('dp_shard')}) "
f"but this run is {'sharded' if cfg.parallelism.is_sharded else 'non-sharded'} "
f"(dp_replicate={cfg.parallelism.dp_replicate}, dp_shard={cfg.parallelism.dp_shard})."
)
recorded = {
"grad_accum_steps": (
metadata["grad_accum_steps"],
cfg.accelerator.gradient_accumulation.steps,
),
}
if snapshot is not None:
recorded.update(
{
"dp_replicate": (snapshot.get("dp_replicate"), cfg.parallelism.dp_replicate),
"dp_shard": (snapshot.get("dp_shard"), cfg.parallelism.dp_shard),
"ring_degree": (
snapshot.get("ring_degree"),
cfg.parallelism.context_parallel.ring_degree,
),
"ulysses_degree": (
snapshot.get("ulysses_degree"),
cfg.parallelism.context_parallel.ulysses_degree,
),
}
)
changed = [f"{key}: {was} -> {now}" for key, (was, now) in recorded.items() if was not in (None, now)]
if changed and is_main_process():
logging.warning(
"Resuming with settings that differ from the checkpoint: " + "; ".join(changed) + ". "
"Topology changes reshard safely via DCP; a changed grad_accum_steps shifts the "
"optimizer-update cadence (the step counter keeps counting micro-batches)."
)
def resume_after_prepare(
cfg: TrainPipelineConfig,
accelerator: "Accelerator",
policy: PreTrainedPolicy,
optimizer: Optimizer | dict[str, Optimizer],
scheduler: LRScheduler | None,
) -> None:
"""Phase 2 — after `accelerator.prepare()`: model (DCP) -> optimizer -> scheduler.
Collective under sharding: call on every rank. The model-weight source follows the
checkpoint's own recorded `checkpoint_format` (on resume, `cfg` was parsed from the
checkpoint's train_config.json): DCP-bearing formats load shards here into the prepared
model (whose construction skipped the safetensors load); the safetensors format was already
loaded by `from_pretrained` before sharding — no model step here.
Args:
cfg (TrainPipelineConfig): The resumed training config; `cfg.checkpoint_path` locates
the checkpoint and `cfg.checkpoint_format` selects the model-weight source.
accelerator (Accelerator): The accelerator the policy was prepared with; it unwraps the
model and owns the DCP load channels.
policy (PreTrainedPolicy): The prepared (possibly sharded) policy to load weights into.
optimizer (Optimizer | dict[str, Optimizer]): The prepared optimizer(s) to restore.
scheduler (LRScheduler | None): The scheduler to restore, or None if the run has none.
Raises:
FileNotFoundError: If the checkpoint format declares DCP model shards but the shard
directory is missing (e.g. it was pruned before upload).
"""
checkpoint_dir = cfg.checkpoint_path
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
training_state_dir = checkpoint_dir / TRAINING_STATE_DIR
unwrapped = accelerator.unwrap_model(policy)
sharded = is_sharded_module(unwrapped)
if cfg.checkpoint_format.wants_dcp:
from accelerate.utils.constants import FSDP_MODEL_NAME
dcp_dir = pretrained_dir / f"{FSDP_MODEL_NAME}_0"
if not dcp_dir.is_dir():
raise FileNotFoundError(
f"checkpoint_format={cfg.checkpoint_format.value} declares DCP model shards, "
f"but {dcp_dir} is missing. If the shards were pruned, convert what remains "
"with `lerobot-convert-dcp` or resume from a safetensors checkpoint."
)
load_sharded_model(accelerator, unwrapped, pretrained_dir)
if sharded:
# Requires the prepared optimizer: FSDP2's prepare rebinds param groups to DTensors but
# never migrates optimizer.state — DCP reshards it here (works across topology changes).
load_sharded_optimizer(accelerator, optimizer, unwrapped, training_state_dir)
else:
load_optimizer_state(optimizer, training_state_dir)
if scheduler is not None: if scheduler is not None:
scheduler = load_scheduler_state(scheduler, training_state_dir) load_scheduler_state(scheduler, training_state_dir)
return step, optimizer, scheduler
def gather_fsdp_state_dicts(model, optimizer) -> tuple[dict, dict]: # ---------------------------------------------------------------------------------------------
"""Gather the full (unsharded) model and optimizer state dicts under FSDP. # Hub: checkpoint push (resume artifact) and publishing (distribution artifact)
# ---------------------------------------------------------------------------------------------
`model.state_dict()` and `FSDP.optim_state_dict(...)` are cross-rank collectives, so this must be
called on *every* rank with the prepared (FSDP-wrapped) `model` and `optimizer`. With
`rank0_only=True` and `offload_to_cpu=True`, every rank runs the all-gather but only rank 0
materializes the full dicts (the others get empty dicts) and they are kept on CPU to bound GPU
memory. The returned optimizer state dict is keyed by parameter FQNs and is world-size
independent; `load_fsdp_optimizer_state` reshards it on resume.
Returns:
(model_state_dict, optim_state_dict): full dicts on rank 0, empty dicts on other ranks.
"""
from torch.distributed.fsdp import (
FullOptimStateDictConfig,
FullStateDictConfig,
FullyShardedDataParallel as FSDP, # noqa F401
StateDictType,
)
state_cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
optim_cfg = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg):
model_state_dict = model.state_dict()
optim_state_dict = FSDP.optim_state_dict(model, optimizer)
return model_state_dict, optim_state_dict
def load_fsdp_optimizer_state(model, optimizer, checkpoint_dir: Path) -> None:
"""Load the FSDP optimizer state (saved as safetensors) and reshard it into the optimizer.
This is a cross-rank collective and must be called on every rank *after* `accelerator.prepare()`
with the prepared (FSDP-wrapped) `model` and `optimizer`. The saved state is the full,
world-size-independent optimizer state (keyed by parameter FQNs); `FSDP.optim_state_dict_to_load`
reshards it to the current FSDP topology, so resume on a different number of GPUs works.
"""
from torch.distributed.fsdp import (
FullOptimStateDictConfig,
FullStateDictConfig,
FullyShardedDataParallel as FSDP, # noqa F401
StateDictType,
)
# Every rank reads the same full state from the (shared) checkpoint dir, so rank0_only=False.
full_osd = load_optimizer_state_dict(checkpoint_dir / TRAINING_STATE_DIR)
state_cfg = FullStateDictConfig(rank0_only=False)
optim_cfg = FullOptimStateDictConfig(rank0_only=False)
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, state_cfg, optim_cfg):
sharded_osd = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=full_osd)
optimizer.load_state_dict(sharded_osd)
def push_checkpoint_to_hub( def push_checkpoint_to_hub(
@@ -311,6 +510,16 @@ def push_checkpoint_to_hub(
The model repo is created idempotently, and the commit is tagged with the The model repo is created idempotently, and the commit is tagged with the
checkpoint step so a checkpoint can be recovered with checkpoint step so a checkpoint can be recovered with
--policy.pretrained_revision=<step> instead of a commit sha. --policy.pretrained_revision=<step> instead of a commit sha.
The directory is uploaded verbatim — including DCP shards under the DCP formats: this tree
exists for *resume*, not distribution, and `resolve_resume_checkpoint` downloads it back
symmetrically.
Args:
checkpoint_dir (Path): The local checkpoint step directory to upload.
repo_id (str): The Hub model repo to push to (created idempotently if missing).
private (bool | None): Whether a newly created repo should be private. Defaults to
None (public unless the organization's default is private).
""" """
api = HfApi() api = HfApi()
api.create_repo(repo_id=repo_id, repo_type="model", private=private, exist_ok=True) api.create_repo(repo_id=repo_id, repo_type="model", private=private, exist_ok=True)
@@ -338,6 +547,16 @@ def resolve_resume_checkpoint(repo_id: str, output_dir: Path) -> Path:
into `output_dir/checkpoints/<step>/`, recreate the local `last` symlink, and return that local into `output_dir/checkpoints/<step>/`, recreate the local `last` symlink, and return that local
checkpoint dir. Used to resume training from the Hub on a machine (or HF Jobs pod) that does not checkpoint dir. Used to resume training from the Hub on a machine (or HF Jobs pod) that does not
have the original local run dir. have the original local run dir.
Args:
repo_id (str): The Hub model repo holding `checkpoints/<step>/` subtrees.
output_dir (Path): The local run directory to download the checkpoint into.
Returns:
Path: The local checkpoint step directory, `output_dir/checkpoints/<step>`.
Raises:
FileNotFoundError: If the repo contains no checkpoints under `checkpoints/`.
""" """
latest = find_latest_hub_checkpoint(repo_id) latest = find_latest_hub_checkpoint(repo_id)
if latest is None: if latest is None:
@@ -354,3 +573,214 @@ def resolve_resume_checkpoint(repo_id: str, output_dir: Path) -> Path:
checkpoint_dir = output_dir / latest checkpoint_dir = output_dir / latest
update_last_checkpoint(checkpoint_dir) update_last_checkpoint(checkpoint_dir)
return checkpoint_dir return checkpoint_dir
def publish_trained_model(
cfg: TrainPipelineConfig,
model: "PreTrainedPolicy | PreTrainedRewardModel",
preprocessor: PolicyProcessorPipeline | None,
postprocessor: PolicyProcessorPipeline | None,
dataset_meta: "LeRobotDatasetMetadata | None",
*,
peft_model: Any | None = None,
) -> None:
"""Publish the complete training bundle as a distributable model repo.
Collective-safe: call on ALL ranks — the model commit gathers sharded weights through
`save_pretrained`; uploads happen on the main process only (gated inside
`HubMixin.push_to_hub` and here). Commits, in order: (1) the model (skipped for PEFT —
adapters replace full weights), (2) the preprocessor, (3) the postprocessor, (4) the bundle
sidecar: README.md model card + train_config.json (+ adapter weights and the wrapped
policy's config in the PEFT case). Every commit uploads a freshly assembled directory, so
a published repo carries only the distributable artifacts.
Args:
cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and used
to render the model card.
model (PreTrainedPolicy | PreTrainedRewardModel): The trained model to publish; its
config supplies the target repo id, visibility, license, and tags.
preprocessor (PolicyProcessorPipeline | None): The preprocessor pipeline to publish
alongside the model, if any.
postprocessor (PolicyProcessorPipeline | None): The postprocessor pipeline to publish
alongside the model, if any.
dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata for the model card, if
available.
peft_model (Any | None): The PEFT wrapper when training adapters; its adapter weights
replace the full model weights in the published repo. Defaults to None.
Raises:
ValueError: If the model config carries no repo id (`--policy.repo_id`).
"""
model_cfg = model.config
repo_id = model_cfg.repo_id
if not repo_id:
raise ValueError("Publishing requires a repo id (--policy.repo_id).")
ignore = ["*.tmp", "*.log"]
if peft_model is None:
# Calls are made on the exact objects that own each method (never through PEFT's
# attribute forwarding), so the peft branch below never touches this path.
model.push_to_hub(repo_id, private=model_cfg.private, ignore_patterns=ignore)
if preprocessor is not None:
preprocessor.push_to_hub(repo_id, private=model_cfg.private)
if postprocessor is not None:
postprocessor.push_to_hub(repo_id, private=model_cfg.private)
if is_main_process():
api = HfApi()
repo_id = api.create_repo(repo_id=repo_id, private=model_cfg.private, exist_ok=True).repo_id
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
saved_path = Path(tmp) / repo_id
saved_path.mkdir(parents=True, exist_ok=True)
if peft_model is not None:
peft_model.save_pretrained(saved_path) # adapter weights + adapter config
model.config.save_pretrained(saved_path) # PEFT cannot write the policy config
card = generate_model_card(model_cfg, cfg=cfg, dataset_meta=dataset_meta)
card.save(str(saved_path / "README.md"))
cfg.save_pretrained(saved_path) # train_config.json
commit_info = api.upload_folder(
repo_id=repo_id,
repo_type="model",
folder_path=saved_path,
commit_message="Upload model card and train config",
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
ignore_patterns=ignore,
)
# Contract: lerobot.jobs.hf.submit_to_hf watches for this exact "Model pushed to <url>"
# line to end a remote run early. Keep the wording and URL format in sync.
logging.info(f"Model pushed to {commit_info.repo_url.url}")
if dist.is_initialized():
dist.barrier()
# ---------------------------------------------------------------------------------------------
# Model card
# ---------------------------------------------------------------------------------------------
_BASE_MODEL_MAPPING = {
"smolvla": "lerobot/smolvla_base",
"pi0": "lerobot/pi0_base",
"pi05": "lerobot/pi05_base",
"pi0_fast": "lerobot/pi0fast-base",
"xvla": "lerobot/xvla-base",
}
def build_card_context(
cfg: TrainPipelineConfig | None,
dataset_meta: "LeRobotDatasetMetadata | None",
input_features: dict | None,
output_features: dict | None,
) -> dict:
"""Collect optional data for the model-card template.
Returns plain values only (no Markdown) — the template in
``lerobot/templates/lerobot_modelcard_template.md`` decides how and whether to show
each one. Everything is best-effort: anything unavailable is left empty/None and the
template simply skips that section, so this never breaks a Hub push.
Args:
cfg (TrainPipelineConfig | None): The training config supplying the training section,
if available.
dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata supplying the dataset,
robot-type, and camera sections, if available.
input_features (dict | None): The policy's input feature declarations, if any.
output_features (dict | None): The policy's output feature declarations, if any.
Returns:
dict: Template context with `training`, `input_features`, `output_features`,
`dataset`, `robot_type`, and `cameras` entries; unavailable pieces stay
empty/None.
"""
context = {
"training": None,
"input_features": input_features or {},
"output_features": output_features or {},
"dataset": None,
"robot_type": None,
"cameras": [],
}
if cfg is not None:
optimizer = getattr(cfg, "optimizer", None)
context["training"] = {
"steps": cfg.steps,
"batch_size": cfg.batch_size,
"seed": cfg.seed,
"optimizer": getattr(optimizer, "type", None) if optimizer else None,
"lr": getattr(optimizer, "lr", None) if optimizer else None,
"lerobot_version": __version__,
}
if dataset_meta is not None:
context["dataset"] = {
"repo_id": dataset_meta.repo_id,
"episodes": dataset_meta.total_episodes,
"frames": dataset_meta.total_frames,
"fps": dataset_meta.fps,
"tasks": [str(task) for task in dataset_meta.tasks.index],
}
context["robot_type"] = dataset_meta.robot_type
context["cameras"] = [key.split(".")[-1] for key in dataset_meta.camera_keys]
return context
def generate_model_card(
model_cfg: PreTrainedConfig | RewardModelConfig,
cfg: TrainPipelineConfig | None = None,
dataset_meta: "LeRobotDatasetMetadata | None" = None,
) -> ModelCard:
"""Render the LeRobot model card for a trained policy or reward model.
A free function on purpose: every template variable comes from arguments — the model
config, the training config, and the dataset metadata — none from a live model, so a card
can also be rendered from a checkpoint's `config.json` alone (see `lerobot-convert-dcp`).
The config type selects the template: reward models get the reward-model card, policies the
policy card with the training/dataset sections.
Args:
model_cfg (PreTrainedConfig | RewardModelConfig): The model config providing type,
license, tags, repo id, and — for policies — the feature declarations.
cfg (TrainPipelineConfig | None, optional): The training config for the training and
dataset card sections. Defaults to None.
dataset_meta (LeRobotDatasetMetadata | None, optional): Dataset metadata for the
dataset card sections. Defaults to None.
Returns:
ModelCard: The rendered and validated LeRobot model card.
"""
model_type = model_cfg.type
base_model = _BASE_MODEL_MAPPING.get(model_type)
if isinstance(model_cfg, RewardModelConfig):
tags = {"robotics", "lerobot", "reward-model", model_type}
template_card = (
files("lerobot.templates")
.joinpath("lerobot_rewardmodel_modelcard_template.md")
.read_text("utf-8")
)
context: dict[str, Any] = {} # the reward template renders from card_data alone
else:
tags = {"robotics", "lerobot", model_type}
template_card = (
files("lerobot.templates").joinpath("lerobot_modelcard_template.md").read_text("utf-8")
)
context = build_card_context(cfg, dataset_meta, model_cfg.input_features, model_cfg.output_features)
# Used by the template to pre-fill commands and the "Fine-tuned from" line.
context["policy_repo_id"] = model_cfg.repo_id
context["base_model"] = base_model
card_data = ModelCardData(
license=model_cfg.license or "apache-2.0",
library_name="lerobot",
pipeline_tag="robotics",
tags=list(tags.union(model_cfg.tags or [])),
model_name=model_type,
datasets=cfg.dataset.repo_id if cfg is not None else None,
base_model=base_model,
)
card = ModelCard.from_template(card_data, template_str=template_card, **context)
card.validate()
return card
+2 -1
View File
@@ -22,7 +22,7 @@ Import them directly: ``from lerobot.configs.train import TrainPipelineConfig``
""" """
from .dataset import DatasetRecordConfig 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 .policies import PreTrainedConfig
from .recipe import MessageTurn, TrainingRecipe, load_recipe from .recipe import MessageTurn, TrainingRecipe, load_recipe
from .types import ( from .types import (
@@ -57,6 +57,7 @@ __all__ = [
# Config classes # Config classes
"DatasetRecordConfig", "DatasetRecordConfig",
"DatasetConfig", "DatasetConfig",
"EMAConfig",
"EvalConfig", "EvalConfig",
"JobConfig", "JobConfig",
"MessageTurn", "MessageTurn",
+273
View File
@@ -0,0 +1,273 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Execution-runtime configuration: everything handed to (or applied by) the `Accelerator`.
Each sub-config mirrors the plain-typed subset of the corresponding accelerate object and
builds it at runtime (the way ``OptimizerConfig.build()`` constructs a ``torch.optim.Optimizer``),
so the whole tree round-trips through the CLI and ``train_config.json`` and parsing a config
never imports accelerate.
"""
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING
from lerobot.configs.parallelism import ParallelismConfig
if TYPE_CHECKING:
from accelerate import Accelerator
from accelerate.utils import (
DistributedDataParallelKwargs,
FullyShardedDataParallelPlugin,
GradientAccumulationPlugin,
)
@dataclass
class FSDPConfig:
"""Mirror of the `FullyShardedDataParallelPlugin` subset LeRobot supports (FSDP2 only).
Exactly one wrap policy applies: `wrap_modules` (module *class names* forming the FSDP
units — and, later, the activation-checkpointing units) or `min_num_params` (size-based).
When both are None, the policy's own `_fsdp_wrap_modules` declaration is used; a run where
no wrap source exists at all fails loudly rather than silently wrapping only the root.
"""
reshard_after_forward: bool = True
wrap_modules: list[str] | None = None
min_num_params: int | None = None
cpu_offload: bool = False
# Regex matched against module FQNs to exclude their parameters from sharding.
ignored_modules: str | None = None
def __post_init__(self) -> None:
"""Validate the wrap-policy fields.
Raises:
ValueError: If both ``wrap_modules`` and ``min_num_params`` are set (they are
mutually exclusive wrap policies), or if ``min_num_params`` is < 1.
"""
if self.wrap_modules is not None and self.min_num_params is not None:
raise ValueError(
"fsdp.wrap_modules and fsdp.min_num_params are mutually exclusive wrap policies."
)
if self.min_num_params is not None and self.min_num_params < 1:
raise ValueError(f"fsdp.min_num_params must be >= 1, got {self.min_num_params}.")
def build_plugin(self) -> "FullyShardedDataParallelPlugin":
"""Build the FSDP2 plugin for `Accelerator(fsdp_plugin=...)`.
Returns:
FullyShardedDataParallelPlugin: FSDP2 (`fsdp_version=2`) plugin carrying the
mirrored wrap policy, resharding, CPU-offload, and ignored-modules settings.
"""
from accelerate.utils import FullyShardedDataParallelPlugin
use_size_policy = self.min_num_params is not None
return FullyShardedDataParallelPlugin(
fsdp_version=2,
reshard_after_forward=self.reshard_after_forward,
auto_wrap_policy="size_based_wrap" if use_size_policy else "transformer_based_wrap",
# May legitimately still be None here: the policy-declared default is applied right
# before `accelerator.prepare()` (see lerobot.distributed.factory.set_fsdp_wrap_modules).
transformer_cls_names_to_wrap=list(self.wrap_modules) if self.wrap_modules else None,
min_num_params=self.min_num_params,
cpu_offload=self.cpu_offload,
ignored_modules=self.ignored_modules,
# state_dict_type stays at the FSDP2 default (SHARDED_STATE_DICT) and is never
# switched: full gathers go through torch's state-dict API, which does not consult
# the plugin. activation_checkpointing stays False: AC is LeRobot-owned.
)
@dataclass
class DDPConfig:
"""Mirror of the `DistributedDataParallelKwargs` subset LeRobot exposes."""
# Today's in-script default, kept for models with conditional computation.
find_unused_parameters: bool = True
gradient_as_bucket_view: bool = False
static_graph: bool = False
def build_kwargs_handler(self) -> "DistributedDataParallelKwargs":
"""Build the DDP kwargs handler for `Accelerator(kwargs_handlers=[...])`.
Returns:
DistributedDataParallelKwargs: Handler carrying the mirrored DDP fields, applied
by accelerate when it wraps the model in `DistributedDataParallel`.
"""
from accelerate.utils import DistributedDataParallelKwargs
return DistributedDataParallelKwargs(
find_unused_parameters=self.find_unused_parameters,
gradient_as_bucket_view=self.gradient_as_bucket_view,
static_graph=self.static_graph,
)
@dataclass
class GradientAccumulationConfig:
"""Mirror of the `GradientAccumulationPlugin` subset LeRobot supports.
Only the step count is a knob. ``sync_with_dataloader`` is pinned to False by
:meth:`build_plugin`: the training loop cycles a finite dataloader, so accelerate's default
of syncing at every dataloader end would force an optimizer step at every dataset epoch
boundary instead of every ``steps`` micro-batches.
"""
steps: int = 1
def __post_init__(self) -> None:
"""Validate the accumulation step count.
Raises:
ValueError: If ``steps`` is < 1.
"""
if self.steps < 1:
raise ValueError(f"gradient_accumulation.steps must be >= 1, got {self.steps}.")
def build_plugin(self) -> "GradientAccumulationPlugin":
"""Build the plugin for `Accelerator(gradient_accumulation_plugin=...)`.
A named plugin argument, not a `kwargs_handlers` entry: accelerate consumes this object
through its dedicated constructor parameter — the `KwargsHandler` base class only lends
it `to_kwargs()`, so the consumption site, not the inheritance, decides its role.
Returns:
GradientAccumulationPlugin: Carrying the mirrored step count, with
``sync_with_dataloader=False`` pinned (see the class docstring).
"""
from accelerate.utils import GradientAccumulationPlugin
return GradientAccumulationPlugin(num_steps=self.steps, sync_with_dataloader=False)
@dataclass
class CompileConfig:
"""torch.compile knobs — a configured placeholder: wiring lands in a later round.
The setup-order contract it will follow is already fixed: compile applies
after CP dispatch install and activation checkpointing, before `fully_shard`, regionally
(per wrap unit) — the only combination proven with FSDP2.
"""
enabled: bool = False
backend: str = "inductor"
mode: str | None = None
regional: bool = True
class ActivationCheckpointingMode(str, Enum):
NONE = "none"
FULL = "full"
@dataclass
class ActivationCheckpointingConfig:
"""Activation-checkpointing knobs — a configured placeholder: wiring lands in a later round.
AC units will coincide with the FSDP wrap units (one declaration drives both), applied
before torch.compile and `fully_shard` (the same ordering contract as CompileConfig).
"""
mode: ActivationCheckpointingMode = ActivationCheckpointingMode.NONE
@dataclass
class AcceleratorConfig:
"""Builds the `Accelerator` — the runtime counterpart of the `parallelism` topology.
`mixed_precision` selects accelerate-native AMP for DDP/single-GPU runs and the FSDP2
`MixedPrecisionPolicy` for sharded runs (accelerate derives it). Sharded runs support
"no" and "bf16" only; fp16's GradScaler-over-DTensor path is unverified and fails fast
at config validation.
"""
mixed_precision: str = "no"
gradient_accumulation: GradientAccumulationConfig = field(default_factory=GradientAccumulationConfig)
fsdp: FSDPConfig = field(default_factory=FSDPConfig)
ddp: DDPConfig = field(default_factory=DDPConfig)
compile: CompileConfig = field(default_factory=CompileConfig)
activation_checkpointing: ActivationCheckpointingConfig = field(
default_factory=ActivationCheckpointingConfig
)
def __post_init__(self) -> None:
"""Validate the accelerate-facing scalar fields.
Raises:
ValueError: If ``mixed_precision`` is not one of ``"no"``, ``"fp16"``, ``"bf16"``.
"""
if self.mixed_precision not in ("no", "fp16", "bf16"):
raise ValueError(
f"mixed_precision must be one of 'no', 'fp16', 'bf16', got {self.mixed_precision!r}."
)
def build(self, parallelism: ParallelismConfig, *, cpu: bool = False) -> "Accelerator":
"""Translate the mirrored fields into a ready `Accelerator` (call once per process).
`parallelism` must already be resolved against the world size. The degradation matrix
is encoded here and nowhere else: sharded -> FSDP2 (+HSDP via the accelerate
`ParallelismConfig` mesh), replicated-only -> DDP kwargs, single process -> plain.
Args:
parallelism (ParallelismConfig): The resolved process topology; selects which
accelerate path (FSDP2 mesh, DDP kwargs handler, or plain) is configured.
cpu (bool): Force CPU execution even when CUDA is available. Defaults to False.
Returns:
Accelerator: The configured accelerate entry point for this process.
"""
from accelerate import Accelerator
kwargs: dict = {
# LeRobot steps its scheduler manually once per training step; accelerate must not
# rescale scheduler stepping by num_processes.
"step_scheduler_with_optimizer": False,
"gradient_accumulation_plugin": self.gradient_accumulation.build_plugin(),
"mixed_precision": self.mixed_precision,
"cpu": cpu,
}
if parallelism.is_sharded:
kwargs["fsdp_plugin"] = self.fsdp.build_plugin()
kwargs["parallelism_config"] = _accelerate_parallelism_config(parallelism)
elif parallelism.is_replicated_only:
kwargs["kwargs_handlers"] = [self.ddp.build_kwargs_handler()]
return Accelerator(**kwargs)
def _accelerate_parallelism_config(parallelism: ParallelismConfig) -> object:
"""LeRobot topology -> accelerate `ParallelismConfig`.
CP is declared honestly (`cp_size = ring x ulysses`) so accelerate builds the canonical
mesh, folds CP into the FSDP shard group (`dp_shard_cp`), and duplicates batches within CP
groups. The ring/ulysses sub-structure stays private to `lerobot.distributed.ParallelDims`.
Args:
parallelism (ParallelismConfig): The resolved LeRobot topology to translate.
Returns:
object: The accelerate `ParallelismConfig` mirroring `dp_replicate`, `dp_shard`, and
the collapsed `cp_size` (annotated as `object` so importing this module never
imports accelerate).
"""
from accelerate.parallelism_config import ParallelismConfig as AccelerateParallelismConfig
return AccelerateParallelismConfig(
dp_replicate_size=parallelism.dp_replicate,
dp_shard_size=parallelism.dp_shard,
cp_size=parallelism.cp_size,
)
+53
View File
@@ -139,6 +139,59 @@ class EvalConfig:
return min(by_cpu, self.n_episodes, 64) 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 @dataclass
class PeftConfig: class PeftConfig:
# PEFT offers many fine-tuning methods, layer adapters being the most common and currently also the most # PEFT offers many fine-tuning methods, layer adapters being the most common and currently also the most
+190
View File
@@ -0,0 +1,190 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Declarative process topology for distributed training and inference.
The mesh convention (canonical row-major rank layout, outermost first)::
(dp_replicate, dp_shard, ring, ulysses)
- ``dp_replicate x dp_shard`` is the data-parallel world: HSDP replicates over
``dp_replicate`` and shards parameters over ``dp_shard``. FSDP2's actual shard
group folds context parallelism in (``dp_shard x ring x ulysses``), matching
accelerate's ``dp_shard_cp`` flattening and torchtitan's ``fsdp`` axis.
- ``ring`` is the outer and ``ulysses`` the inner context-parallel dim
(diffusers convention: ulysses all-to-all exchanges run over adjacent, typically
NVLink-connected ranks).
- ``cfg_parallel`` (classifier-free-guidance parallelism) is a branch-parallel,
inference-only dim that sits between dp and the sequence dims. It never
affects weight sharding or checkpoints.
This module is pure configuration: plain-typed dataclasses that draccus can
round-trip through the CLI and ``train_config.json``. Runtime objects (device
meshes, process groups) live in :mod:`lerobot.distributed`.
"""
import os
from dataclasses import dataclass, field
@dataclass
class ContextParallelConfig:
"""Ring x Ulysses context parallelism (sequence parallelism for attention).
Both degrees are configured placeholders in this release: the CP engine is not implemented
yet, and enabling either degree > 1 fails fast at config validation. The fields exist now so
that the CLI surface, checkpoint metadata, and mesh math are stable when the engine lands.
"""
ring_degree: int = 1
ulysses_degree: int = 1
def __post_init__(self) -> None:
"""Validate the declared context-parallel degrees.
Raises:
ValueError: If ``ring_degree`` or ``ulysses_degree`` is < 1.
"""
if self.ring_degree < 1 or self.ulysses_degree < 1:
raise ValueError(
f"Context-parallel degrees must be >= 1, got ring_degree={self.ring_degree}, "
f"ulysses_degree={self.ulysses_degree}."
)
@property
def size(self) -> int:
"""Total number of ranks a full sequence is sharded across."""
return self.ring_degree * self.ulysses_degree
@dataclass
class ParallelismConfig:
"""Degrees of every parallelism dim. Invariant: their product equals the world size.
Degradations are expressed purely through the degrees (no mode flags):
- single process: all degrees 1;
- DDP: ``dp_replicate == world_size`` (auto-filled when every sharding field is left at its
default — plain ``torchrun`` keeps today's out-of-the-box behavior);
- FSDP: ``dp_shard > 1`` (or ``-1`` to fill the remaining world into the shard dim);
- HSDP: ``dp_replicate > 1`` and ``dp_shard > 1``.
``resolve()`` turns the declared degrees into concrete ones once the world size is known and
is the single place the world-size equation is enforced. It is called by
:func:`lerobot.distributed.factory.make_accelerator`; the config is inert until then.
"""
dp_replicate: int = 1
# -1 is an explicit opt-in sentinel: shard over world_size // (dp_replicate * cp).
dp_shard: int = 1
context_parallel: ContextParallelConfig = field(default_factory=ContextParallelConfig)
# Classifier-free-guidance parallelism — inference-only (cosmos/vllm-omni precedent:
# cond/uncond branches on different ranks). Reserved for the serving round; training
# validates it to 1. Meaningful values are 1 or 2 (Cosmos3 has two CFG branches).
cfg_parallel: int = 1
def __post_init__(self) -> None:
"""Validate the declared degrees (world-size-independent checks only).
Raises:
ValueError: If ``dp_replicate`` is < 1, ``dp_shard`` is neither >= 1 nor the
``-1`` infer sentinel, or ``cfg_parallel`` is not 1 or 2.
"""
if self.dp_replicate < 1:
raise ValueError(f"dp_replicate must be >= 1, got {self.dp_replicate}.")
if self.dp_shard < 1 and self.dp_shard != -1:
raise ValueError(f"dp_shard must be >= 1, or -1 to infer, got {self.dp_shard}.")
if self.cfg_parallel not in (1, 2):
raise ValueError(f"cfg_parallel must be 1 or 2, got {self.cfg_parallel}.")
@property
def cp_size(self) -> int:
"""Total context-parallel size (``ring_degree * ulysses_degree``)."""
return self.context_parallel.size
@property
def is_sharded(self) -> bool:
"""True when the run uses FSDP2 (parameters sharded); selects the sharded engine path."""
return self.dp_shard != 1 or self.cp_size > 1
@property
def is_replicated_only(self) -> bool:
"""True for plain DDP (weights replicated, no sharding)."""
return not self.is_sharded and self.dp_replicate > 1
@property
def dp_world_size(self) -> int:
"""Number of distinct data-parallel workers (batches are sharded this many ways).
Returns:
int: ``dp_replicate * dp_shard``.
Raises:
RuntimeError: If accessed while ``dp_shard`` is still the ``-1`` sentinel, i.e.
before :meth:`resolve` has bound the degrees to a world size.
"""
if self.dp_shard == -1:
raise RuntimeError("dp_world_size is undefined before resolve() fills dp_shard=-1.")
return self.dp_replicate * self.dp_shard
def resolve(self, world_size: int) -> None:
"""Bind the declared degrees to a concrete world size (idempotent).
Fills the ``dp_shard=-1`` sentinel, auto-fills ``dp_replicate`` for the DDP degradation,
and enforces ``dp_replicate * dp_shard * cp == world_size`` with every degree echoed on
failure.
Args:
world_size (int): Total number of launched processes (torchrun's ``WORLD_SIZE``).
Raises:
ValueError: If a context-parallel degree is > 1 (the CP engine is not implemented
yet), if ``dp_shard=-1`` cannot be inferred because ``world_size`` is not
divisible by ``dp_replicate * cp``, or if the resolved degrees do not multiply
to ``world_size``.
"""
if self.cp_size > 1:
raise ValueError(
"Context parallelism is not implemented yet: ring_degree and ulysses_degree "
"must be 1. The fields are reserved for the CP engine round."
)
if self.is_sharded:
if self.dp_shard == -1:
self.dp_shard, remainder = divmod(world_size, self.dp_replicate * self.cp_size)
if remainder or self.dp_shard < 1:
raise ValueError(
f"Cannot infer dp_shard: world_size={world_size} is not divisible by "
f"dp_replicate={self.dp_replicate} * cp={self.cp_size}."
)
elif self.dp_replicate == 1:
# Untouched config on a multi-process launch: fill the DDP degradation.
self.dp_replicate = world_size
total = self.dp_replicate * self.dp_shard * self.cp_size
if total != world_size:
raise ValueError(
f"Parallelism degrees do not multiply to the world size: dp_replicate="
f"{self.dp_replicate} * dp_shard={self.dp_shard} * ring="
f"{self.context_parallel.ring_degree} * ulysses="
f"{self.context_parallel.ulysses_degree} = {total} != WORLD_SIZE={world_size}."
)
def world_size_from_env() -> int:
"""World size as set by torchrun (or 1 outside distributed launches).
Returns:
int: The ``WORLD_SIZE`` environment variable, or 1 when unset.
"""
return int(os.environ.get("WORLD_SIZE", "1"))
+95 -1
View File
@@ -18,6 +18,7 @@ import multiprocessing
import os import os
import tempfile import tempfile
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@@ -26,19 +27,49 @@ from huggingface_hub import hf_hub_download
from huggingface_hub.errors import HfHubHTTPError from huggingface_hub.errors import HfHubHTTPError
from lerobot import envs from lerobot import envs
from lerobot.configs.accelerator import AcceleratorConfig, ActivationCheckpointingMode
from lerobot.configs.parallelism import ParallelismConfig
from lerobot.optim import LRSchedulerConfig, OptimizerConfig from lerobot.optim import LRSchedulerConfig, OptimizerConfig
from lerobot.utils.constants import PRETRAINED_MODEL_DIR from lerobot.utils.constants import PRETRAINED_MODEL_DIR
from lerobot.utils.hub import HubMixin, find_latest_hub_checkpoint from lerobot.utils.hub import HubMixin, find_latest_hub_checkpoint
from lerobot.utils.sample_weighting import SampleWeightingConfig from lerobot.utils.sample_weighting import SampleWeightingConfig
from . import parser 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 .policies import PreTrainedConfig
from .rewards import RewardModelConfig from .rewards import RewardModelConfig
TRAIN_CONFIG_NAME = "train_config.json" TRAIN_CONFIG_NAME = "train_config.json"
class CheckpointFormat(str, Enum):
"""Model-artifact format inside training checkpoints.
Selects only the *model* artifact; the training_state layout is format-independent (the
optimizer channel is always DCP under sharded runs, safetensors+json otherwise).
- SAFETENSORS (default): a full `model.safetensors` — maximum compatibility, one gather per
save under sharding.
- DCP: sharded `pytorch_model_fsdp_0/*.distcp` only — fastest save/resume; convert with
`lerobot-convert-dcp` before distributing.
- SAFETENSORS_AND_DCP: both artifacts, written independently.
"""
SAFETENSORS = "safetensors"
DCP = "dcp"
SAFETENSORS_AND_DCP = "safetensors_dcp"
@property
def wants_safetensors(self) -> bool:
"""True when a full `model.safetensors` artifact should be written."""
return self in (CheckpointFormat.SAFETENSORS, CheckpointFormat.SAFETENSORS_AND_DCP)
@property
def wants_dcp(self) -> bool:
"""True when sharded DCP model shards (`pytorch_model_fsdp_0/`) should be written."""
return self in (CheckpointFormat.DCP, CheckpointFormat.SAFETENSORS_AND_DCP)
def _migrate_legacy_rabc_fields(config: dict[str, Any]) -> dict[str, Any] | None: def _migrate_legacy_rabc_fields(config: dict[str, Any]) -> dict[str, Any] | None:
"""Return migrated payload for legacy RA-BC fields, or None when no migration is needed.""" """Return migrated payload for legacy RA-BC fields, or None when no migration is needed."""
legacy_fields = ( legacy_fields = (
@@ -121,10 +152,19 @@ class TrainPipelineConfig(HubMixin):
# Checkpoint is saved every `save_freq` training iterations and after the last training step. # Checkpoint is saved every `save_freq` training iterations and after the last training step.
# A non-positive value disables periodic saving, keeping only the final checkpoint. # A non-positive value disables periodic saving, keeping only the final checkpoint.
save_freq: int = 20_000 save_freq: int = 20_000
# Model-artifact format inside checkpoints; non-default values require a sharded run.
checkpoint_format: CheckpointFormat = CheckpointFormat.SAFETENSORS
use_policy_training_preset: bool = True use_policy_training_preset: bool = True
optimizer: OptimizerConfig | None = None optimizer: OptimizerConfig | None = None
scheduler: LRSchedulerConfig | None = None scheduler: LRSchedulerConfig | None = None
# Process topology: dp_replicate / dp_shard (HSDP) and context-parallel degree placeholders.
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
# Execution runtime handed to the Accelerator: mixed precision, gradient accumulation,
# FSDP/DDP tuning knobs, compile & activation-checkpointing placeholders.
accelerator: AcceleratorConfig = field(default_factory=AcceleratorConfig)
eval: EvalConfig = field(default_factory=EvalConfig) 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) wandb: WandBConfig = field(default_factory=WandBConfig)
peft: PeftConfig | None = None peft: PeftConfig | None = None
@@ -291,6 +331,60 @@ class TrainPipelineConfig(HubMixin):
if self.save_checkpoint_to_hub and not (self.policy is not None and self.policy.repo_id): if self.save_checkpoint_to_hub and not (self.policy is not None and self.policy.repo_id):
raise ValueError("save_checkpoint_to_hub requires --policy.repo_id.") raise ValueError("save_checkpoint_to_hub requires --policy.repo_id.")
self._validate_distributed()
def _validate_distributed(self) -> None:
"""Fail-fasts for the distributed-training scope.
Raises:
ValueError: If the config requests anything outside the verified scope: context
parallelism or CFG parallelism (reserved placeholders), the compile or
activation-checkpointing placeholders, a DCP checkpoint format on a
non-sharded run, or — under sharded training — fp16 mixed precision, PEFT,
reward-model training, in-training environment evaluation, or multi-optimizer
configs.
"""
if self.parallelism.cp_size > 1:
raise ValueError(
"Context parallelism is not implemented yet: --parallelism.context_parallel.* "
"degrees must be 1 (reserved for the CP engine round)."
)
if self.parallelism.cfg_parallel != 1:
raise ValueError(
"CFG parallelism is inference-only and must be 1 for training "
"(cfg_parallel is reserved for the serving round)."
)
if self.accelerator.compile.enabled:
raise ValueError("--accelerator.compile is a placeholder and not wired yet.")
if self.accelerator.activation_checkpointing.mode is not ActivationCheckpointingMode.NONE:
raise ValueError("--accelerator.activation_checkpointing is a placeholder and not wired yet.")
if self.checkpoint_format is not CheckpointFormat.SAFETENSORS and not self.parallelism.is_sharded:
raise ValueError(
f"checkpoint_format={self.checkpoint_format.value} requires a sharded run "
"(--parallelism.dp_shard != 1); non-sharded checkpoints are always safetensors."
)
if self.parallelism.is_sharded:
if self.accelerator.mixed_precision == "fp16":
raise ValueError(
"fp16 is not supported under sharded training (GradScaler over DTensor "
"gradients is unverified); use bf16 or full precision."
)
if self.peft is not None:
raise ValueError("PEFT is not supported under sharded training yet.")
if self.is_reward_model_training:
raise ValueError(
"Reward-model training is not supported under sharded training yet "
"(reward models declare no FSDP wrap units and have no sharded save path)."
)
if self.env is not None and self.env_eval_freq > 0:
raise ValueError(
"In-training environment evaluation is not supported under sharded training "
"(a rank-0-only rollout of a sharded model deadlocks on collectives); set "
"--env_eval_freq=0 and evaluate with lerobot-eval on saved checkpoints."
)
if self.optimizer is not None and self.optimizer.builds_multiple_optimizers:
raise ValueError("Multi-optimizer configs are not supported under sharded training.")
@classmethod @classmethod
def __get_path_fields__(cls) -> list[str]: def __get_path_fields__(cls) -> list[str]:
"""Keys for draccus pretrained-path loading.""" """Keys for draccus pretrained-path loading."""
+9 -2
View File
@@ -613,8 +613,15 @@ def aggregate_feature_stats(stats_ft_list: list[dict[str, dict]]) -> dict[str, d
for q_key in quantile_keys: for q_key in quantile_keys:
if all(q_key in s for s in stats_ft_list): 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]) quantile_values = np.stack([s[q_key] for s in stats_ft_list])
weighted_quantiles = quantile_values * counts # Exact global quantiles cannot be recovered from quantile summaries.
aggregated[q_key] = weighted_quantiles.sum(axis=0) / total_count # 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 return aggregated
+43
View File
@@ -0,0 +1,43 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Distributed-training runtime for LeRobot.
This package owns everything that turns the declarative topology in
:class:`lerobot.configs.parallelism.ParallelismConfig` into a running engine:
mesh math (:class:`~lerobot.distributed.parallel_dims.ParallelDims`), the
`Accelerator` factory (:func:`~lerobot.distributed.factory.make_accelerator`),
sharding-aware checkpoint helpers, and small rank utilities.
Setup-order contract (normative):
CP dispatch install -> activation checkpointing -> torch.compile ->
``fully_shard``/DDP (via ``accelerator.prepare``) -> optimizer rebind.
Only the last two steps are active today; CP/AC/compile are configured
placeholders wired in later rounds.
"""
from .factory import guard_against_env_interference, make_accelerator, set_fsdp_wrap_modules
from .parallel_dims import ParallelDims
from .utils import finalize_sharded_policy, is_main_process, strip_accelerate_cp_hooks
__all__ = [
"ParallelDims",
"finalize_sharded_policy",
"guard_against_env_interference",
"is_main_process",
"make_accelerator",
"set_fsdp_wrap_modules",
"strip_accelerate_cp_hooks",
]
+195
View File
@@ -0,0 +1,195 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Sharding-aware checkpoint primitives.
Two artifact channels with distinct owners:
- the **distributable** ``model.safetensors``: produced by ``PreTrainedPolicy.save_pretrained``
through :func:`full_model_state_dict` — a collective full gather when the model is sharded;
- the **resume** channel (sharded runs): torch DCP directories written/read through accelerate's
``save/load_fsdp_model`` and ``save/load_fsdp_optimizer`` (``pytorch_model_fsdp_0/`` and
``optimizer_0/``, names imported from accelerate constants), which reshard on load across
topology changes.
Every function that touches sharded state is a collective and must run on ALL ranks.
"""
from pathlib import Path
from typing import TYPE_CHECKING
import torch
from torch import nn
if TYPE_CHECKING:
from accelerate import Accelerator
def is_sharded_module(module: nn.Module) -> bool:
"""True when `fully_shard` owns this module's parameters (FSDP2's in-place class swap).
Args:
module (nn.Module): The module to inspect (a torch.compile wrapper is looked through
via `_orig_mod`).
Returns:
bool: True when the module (or its compiled `_orig_mod`) is an `FSDPModule`.
"""
from torch.distributed.fsdp import FSDPModule
if isinstance(module, FSDPModule):
return True
# torch.compile wraps the sharded module; mirror accelerate's `_orig_mod` check.
orig_mod = getattr(module, "_orig_mod", None)
return orig_mod is not None and isinstance(orig_mod, FSDPModule)
def full_model_state_dict(module: nn.Module) -> dict[str, torch.Tensor]:
"""The module's full (unsharded) state dict, however its parameters are laid out.
Sharded modules gather through torch's DCP state-dict API: a COLLECTIVE that must run on
every rank; with ``cpu_offload=True`` the full dict materializes on the main rank only and
every other rank receives a literal ``{}`` (runtime-verified — a
rank-0-gated call deadlocks). Plain modules return ``module.state_dict()`` on every rank.
Args:
module (nn.Module): The (possibly sharded) module to read the state dict from.
Returns:
dict[str, torch.Tensor]: The full state dict — on the main rank only (``{}``
elsewhere) when the module is sharded, on every rank otherwise.
"""
if not is_sharded_module(module):
return module.state_dict()
from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict
return get_model_state_dict(module, options=StateDictOptions(full_state_dict=True, cpu_offload=True))
def _fsdp_plugin(accelerator: "Accelerator") -> object:
"""The accelerator's FSDP plugin, required by every DCP save/load helper below.
Args:
accelerator (Accelerator): The accelerator that prepared the sharded model.
Returns:
object: The FSDP plugin held by `accelerator.state`.
Raises:
RuntimeError: If the accelerator was not configured with an FSDP plugin.
"""
plugin = getattr(accelerator.state, "fsdp_plugin", None)
if plugin is None:
raise RuntimeError("Sharded checkpointing requires an FSDP-prepared Accelerator.")
return plugin
def save_sharded_model(accelerator: "Accelerator", model: nn.Module, output_dir: Path) -> None:
"""Write the DCP model shards (`pytorch_model_fsdp_0/`). Collective: call on all ranks.
Args:
accelerator (Accelerator): The accelerator that prepared the sharded model.
model (nn.Module): The prepared (sharded) model to save.
output_dir (Path): The directory the shard subdirectory is created in.
"""
from accelerate.utils import save_fsdp_model
# accelerate 1.14's DCP helpers do string containment checks on the path:
# always hand them str, never Path.
save_fsdp_model(_fsdp_plugin(accelerator), accelerator, model, str(output_dir))
def load_sharded_model(accelerator: "Accelerator", model: nn.Module, input_dir: Path) -> None:
"""Load DCP model shards into the prepared (sharded) model. Collective: call on all ranks.
Args:
accelerator (Accelerator): The accelerator that prepared the sharded model.
model (nn.Module): The prepared (sharded) model to load into.
input_dir (Path): The directory containing the `pytorch_model_fsdp_0/` shard
subdirectory.
"""
from accelerate.utils import load_fsdp_model
from accelerate.utils.constants import FSDP_MODEL_NAME
# Pass the exact shard directory: accelerate's load resolves it with a substring check
# ("pytorch_model_fsdp" in the path -> use as-is), which misfires on run paths that happen
# to contain the marker; the exact dir makes the check deterministic.
load_fsdp_model(_fsdp_plugin(accelerator), accelerator, model, str(input_dir / f"{FSDP_MODEL_NAME}_0"))
def save_sharded_optimizer(
accelerator: "Accelerator", optimizer: torch.optim.Optimizer, model: nn.Module, output_dir: Path
) -> None:
"""Write the DCP optimizer shards (`optimizer_0/`). Collective: call on all ranks.
Args:
accelerator (Accelerator): The accelerator that prepared the model and optimizer.
optimizer (torch.optim.Optimizer): The prepared optimizer to save the state from.
model (nn.Module): The prepared (sharded) model the optimizer state is keyed by.
output_dir (Path): The directory the shard subdirectory is created in.
"""
from accelerate.utils import save_fsdp_optimizer
save_fsdp_optimizer(_fsdp_plugin(accelerator), accelerator, optimizer, model, str(output_dir))
def load_sharded_optimizer(
accelerator: "Accelerator", optimizer: torch.optim.Optimizer, model: nn.Module, input_dir: Path
) -> None:
"""Load DCP optimizer shards into the prepared optimizer. Collective: call on all ranks.
Must run AFTER ``accelerator.prepare()``: FSDP2's prepare rebinds the optimizer's param
groups to sharded DTensors but never migrates ``optimizer.state`` — the resharding load is
the only correct way to restore it.
Args:
accelerator (Accelerator): The accelerator that prepared the model and optimizer.
optimizer (torch.optim.Optimizer): The prepared optimizer to restore the state into.
model (nn.Module): The prepared (sharded) model the optimizer state is keyed by.
input_dir (Path): The directory containing the `optimizer_0/` shard subdirectory.
"""
from accelerate.utils import load_fsdp_optimizer
from accelerate.utils.constants import OPTIMIZER_NAME
# Exact shard directory for the same reason as load_sharded_model: accelerate's substring
# check ("optimizer" in the path) would misread e.g. --job_name=optimizer_sweep run paths.
load_fsdp_optimizer(
_fsdp_plugin(accelerator), accelerator, optimizer, model, str(input_dir / f"{OPTIMIZER_NAME}_0")
)
def dcp_to_safetensors(dcp_dir: Path, output_dir: Path, *, delete_dcp: bool = False) -> Path:
"""Merge a DCP shard directory into a single `model.safetensors` (offline, single process).
Thin wrapper over `accelerate.utils.merge_fsdp_weights`, which loads the shards without a
process group, writes safetensors directly, and — when asked — removes the merged shard
directory itself, only on the main process and only once the merge has succeeded.
Args:
dcp_dir (Path): The DCP shard directory to merge (e.g. `.../pytorch_model_fsdp_0`).
output_dir (Path): The directory the merged `model.safetensors` is written into.
delete_dcp (bool): Whether to remove the shard directory once it has been merged.
Defaults to False.
Returns:
Path: The written `model.safetensors` file's path.
"""
from accelerate.utils import merge_fsdp_weights
merge_fsdp_weights(
str(dcp_dir), str(output_dir), safe_serialization=True, remove_checkpoint_dir=delete_dcp
)
return output_dir / "model.safetensors"
+139
View File
@@ -0,0 +1,139 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""The `Accelerator` factory — the only place accelerate gets configured.
`torchrun` is the launcher; every accelerate parameter comes from `TrainPipelineConfig`
(`cfg.parallelism` + `cfg.accelerator`) so a run is reproducible from its `train_config.json`
alone. `accelerate launch` without a `--config_file` remains equivalent (it only sets rendezvous
env vars in that mode); the yaml flow is superseded.
"""
import os
from typing import TYPE_CHECKING
from lerobot.configs.parallelism import world_size_from_env
from lerobot.configs.train import TrainPipelineConfig
if TYPE_CHECKING:
from accelerate import Accelerator
from lerobot.policies.pretrained import PreTrainedPolicy
# Env vars through which `accelerate launch --config_file` (or a stray shell) would configure
# accelerate behind the config system's back, making train_config.json lie about what ran.
_ACCELERATE_ENV_VARS = (
"ACCELERATE_USE_FSDP",
"ACCELERATE_USE_PARALLELISM_CONFIG",
"ACCELERATE_GRADIENT_ACCUMULATION_STEPS",
)
_ENV_OVERRIDE = "LEROBOT_ALLOW_ACCELERATE_ENV"
def guard_against_env_interference() -> None:
"""Hard-error when accelerate-configuring env vars are set.
A silently env-overridden "reproducible" config is worse than a stop: users migrating from
the old `accelerate launch --config_file fsdp.yaml` flow get a precise error instead of a
config that lies. Set LEROBOT_ALLOW_ACCELERATE_ENV=1 to acknowledge and proceed.
Raises:
RuntimeError: If any accelerate-configuring environment variable is set and the
LEROBOT_ALLOW_ACCELERATE_ENV override is not.
"""
if os.environ.get(_ENV_OVERRIDE):
return
offending = sorted(name for name in _ACCELERATE_ENV_VARS if name in os.environ)
if offending:
raise RuntimeError(
f"Accelerate-configuring environment variables are set: {', '.join(offending)}. "
"LeRobot manages accelerate exclusively through TrainPipelineConfig "
"(--parallelism.* / --accelerator.*); launch with plain torchrun and remove these "
"variables (the `accelerate launch --config_file` flow is superseded), or set "
f"{_ENV_OVERRIDE}=1 to acknowledge that they may override your config."
)
def make_accelerator(cfg: TrainPipelineConfig) -> "Accelerator":
"""Resolve the topology against the launched world and build the `Accelerator`.
Must run once per process, before any other component needs the device or the process
group (`Accelerator.__init__` initializes both and builds the device mesh).
Args:
cfg (TrainPipelineConfig): The full training config; `cfg.parallelism` is resolved in
place against the launched world size and `cfg.accelerator` builds the result.
Returns:
Accelerator: The configured accelerator, with device and process group initialized.
Raises:
ValueError: If `cfg.checkpoint_format` requires DCP but the topology resolved to a
non-sharded run.
"""
guard_against_env_interference()
cfg.parallelism.resolve(world_size_from_env())
# The parse-time format check ran against the declared degrees, where the dp_shard=-1
# sentinel counts as sharded; it may resolve to an unsharded run (e.g. -1 at world size 1).
# Re-check against the concrete degrees so the recorded format never lies about the
# artifacts a checkpoint will actually contain.
if cfg.checkpoint_format.wants_dcp and not cfg.parallelism.is_sharded:
raise ValueError(
f"checkpoint_format={cfg.checkpoint_format.value} requires a sharded run, but the "
f"topology resolved to a non-sharded one (dp_replicate={cfg.parallelism.dp_replicate}, "
f"dp_shard={cfg.parallelism.dp_shard}); non-sharded checkpoints are always safetensors."
)
return cfg.accelerator.build(
cfg.parallelism,
cpu=cfg.trainable_config.device == "cpu",
)
def set_fsdp_wrap_modules(accelerator: "Accelerator", policy: "PreTrainedPolicy") -> None:
"""Resolve the FSDP wrap-unit class names onto the plugin before `accelerator.prepare()`.
Resolution order: user override (`--accelerator.fsdp.wrap_modules`, already on the plugin)
-> the policy's `_fsdp_wrap_modules` declaration -> hard error. Root-only wrapping — the
silent default when no wrap source exists — is never accepted: it quietly forfeits all
sharding memory savings.
No-op for the size-based policy (`--accelerator.fsdp.min_num_params`), which needs no class
names, and for non-sharded runs (no fsdp plugin).
Args:
accelerator (Accelerator): The accelerator whose FSDP plugin receives the wrap-unit
class names.
policy (PreTrainedPolicy): The trainable whose class may declare `_fsdp_wrap_modules`.
Raises:
ValueError: If sharded class-based wrapping is configured but neither a user override
nor a policy declaration supplies wrap-unit class names.
"""
plugin = getattr(accelerator.state, "fsdp_plugin", None)
if plugin is None or plugin.min_num_params:
return
if plugin.transformer_cls_names_to_wrap: # user override, set at build time
return
# getattr, not attribute access: non-policy trainables (no `_fsdp_wrap_modules` attribute)
# must reach the actionable error below, not an AttributeError.
declared = getattr(type(policy), "_fsdp_wrap_modules", None)
if not declared:
raise ValueError(
f"Policy '{type(policy).__name__}' declares no FSDP wrap units. Sharded training "
"requires wrap-unit class names: set --accelerator.fsdp.wrap_modules='[\"MyBlock\"]' "
"(or --accelerator.fsdp.min_num_params for a size-based policy), or declare "
"`_fsdp_wrap_modules` on the policy class."
)
plugin.transformer_cls_names_to_wrap = list(declared)
+112
View File
@@ -0,0 +1,112 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Runtime mesh math derived from the declarative :class:`ParallelismConfig`.
`ParallelDims` is the training script's single source of truth for topology-derived numbers
(data-parallel world size and rank, sample accounting inputs) and — once the CP engine lands —
the owner of LeRobot's private ``(dp_replicate, dp_shard, ring, ulysses)`` mesh. It is a runtime
object and is never serialized (the config it derives from is what lands in
``train_config.json``).
"""
from dataclasses import dataclass
import torch.distributed as dist
from lerobot.configs.parallelism import ParallelismConfig
@dataclass(frozen=True)
class ParallelDims:
"""Concrete parallelism degrees bound to a world size (canonical row-major rank layout)."""
dp_replicate: int
dp_shard: int
ring: int
ulysses: int
world_size: int
device_type: str
@classmethod
def from_config(cls, cfg: ParallelismConfig, world_size: int, device_type: str) -> "ParallelDims":
"""Bind a *resolved* config to the actual runtime world size (cross-checked here).
Args:
cfg (ParallelismConfig): The declarative topology, already resolved via
`ParallelismConfig.resolve(world_size)`.
world_size (int): The launched world size the declared degrees must multiply to.
device_type (str): The accelerator device type backing the mesh (e.g. "cuda").
Returns:
ParallelDims: The concrete parallelism degrees bound to this world.
Raises:
ValueError: If the config is unresolved (`dp_shard == -1`) or its degrees do not
multiply to `world_size`.
"""
total = cfg.dp_replicate * cfg.dp_shard * cfg.cp_size
if cfg.dp_shard == -1 or total != world_size:
raise ValueError(
f"ParallelismConfig is not resolved against this world: dp_replicate="
f"{cfg.dp_replicate} * dp_shard={cfg.dp_shard} * cp={cfg.cp_size} != "
f"world_size={world_size}. Call ParallelismConfig.resolve(world_size) first "
"(make_accelerator does this)."
)
return cls(
dp_replicate=cfg.dp_replicate,
dp_shard=cfg.dp_shard,
ring=cfg.context_parallel.ring_degree,
ulysses=cfg.context_parallel.ulysses_degree,
world_size=world_size,
device_type=device_type,
)
@property
def cp_size(self) -> int:
"""Total context-parallel degree (`ring * ulysses`)."""
return self.ring * self.ulysses
@property
def is_sharded(self) -> bool:
"""Whether parameters are sharded (`dp_shard > 1` or any context parallelism)."""
return self.dp_shard > 1 or self.cp_size > 1
@property
def dp_world_size(self) -> int:
"""Number of distinct data-parallel workers — the divisor for all sample accounting."""
return self.dp_replicate * self.dp_shard
@property
def dp_rank(self) -> int:
"""This process's data-parallel coordinate (CP peers share one dp_rank).
With the canonical row-major layout and (ring, ulysses) innermost, CP peers are
contiguous global ranks, so the dp coordinate is the integer quotient by cp_size —
the same arithmetic accelerate's mesh-aware dataloader applies.
"""
global_rank = dist.get_rank() if dist.is_initialized() else 0
return global_rank // self.cp_size
def cp_mesh(self) -> None:
"""Private (ring, ulysses) mesh for the CP engine — reserved for the CP round.
Raises:
NotImplementedError: Always — context parallelism is not implemented yet.
"""
raise NotImplementedError(
"Context parallelism is not implemented yet; ParallelDims.cp_mesh is reserved for "
"the CP engine round (a private mesh aligned with accelerate's cp block)."
)
+94
View File
@@ -0,0 +1,94 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Rank utilities and post-`prepare()` sharding finalization."""
import logging
from typing import TYPE_CHECKING
import torch.distributed as dist
from torch import nn
if TYPE_CHECKING:
from lerobot.distributed.parallel_dims import ParallelDims
def is_main_process() -> bool:
"""True on the process that owns rank-0-only side effects (file writes, uploads, logging).
Torch-native on purpose: persistence code must not depend on an `Accelerator` handle —
`_save_pretrained` and the hub publishers run in contexts that have none. Outside
distributed runs every process is the main process.
Returns:
bool: True when this process is rank 0 or no process group is initialized.
"""
return not dist.is_initialized() or dist.get_rank() == 0
def strip_accelerate_cp_hooks(model: nn.Module) -> int:
"""Remove accelerate's context-parallel forward-pre-hooks from every module.
When `cp_size > 1` is declared, `accelerator.prepare()` unconditionally attaches hooks that
silently replace any `attention_mask` kwarg of `*self_attn` modules with `is_causal=True`
(`accelerate.big_modeling._attach_context_parallel_hooks`) — mask corruption for policies
with non-causal attention. LeRobot implements CP itself and never enters accelerate's CP
context, so these hooks are pure hazard. Deterministically identified by their defining
module; a version canary pins that identity.
Args:
model (nn.Module): The prepared model to strip the hooks from (all submodules are
visited).
Returns:
int: The number of hooks removed.
"""
removed = 0
for module in model.modules():
for hook_id, hook in list(module._forward_pre_hooks.items()):
if getattr(hook, "__module__", None) == "accelerate.big_modeling":
del module._forward_pre_hooks[hook_id]
module._forward_pre_hooks_with_kwargs.pop(hook_id, None)
removed += 1
return removed
def finalize_sharded_policy(policy: nn.Module, parallel_dims: "ParallelDims") -> None:
"""Sharding correctness protocol, applied once, immediately after `accelerator.prepare()`.
1. Strip accelerate's CP mask hooks (only attached when cp > 1 was declared).
2. Register the policy's non-`forward` entry points (`_fsdp_forward_methods`) so FSDP2
unshards parameters around `select_action` & co. — without this, any inference-style
call on a sharded policy crashes on mixed Tensor/DTensor.
No-op for DDP/single-process runs.
Args:
policy (nn.Module): The policy as returned by `accelerator.prepare()`.
parallel_dims (ParallelDims): The run's resolved topology; decides whether the protocol
applies.
"""
if not parallel_dims.is_sharded:
return
if parallel_dims.cp_size > 1:
removed = strip_accelerate_cp_hooks(policy)
logging.info("Stripped %d accelerate context-parallel attention-mask hooks.", removed)
from torch.distributed.fsdp import FSDPModule, register_fsdp_forward_method
if isinstance(policy, FSDPModule):
for method_name in getattr(type(policy), "_fsdp_forward_methods", ()):
if callable(getattr(policy, method_name, None)):
register_fsdp_forward_method(policy, method_name)
+1 -1
View File
@@ -432,7 +432,7 @@ def submit_to_hf(cfg: TrainPipelineConfig) -> None:
# Finish as soon as the model is pushed, rather than waiting out the platform's # Finish as soon as the model is pushed, rather than waiting out the platform's
# post-run finalization before the job stage flips to COMPLETED. This matches the # post-run finalization before the job stage flips to COMPLETED. This matches the
# exact log line emitted by PreTrainedPolicy.push_model_to_hub — the two must stay # exact log line emitted by lerobot.common.train_utils.publish_trained_model — the two must stay
# in sync. If it ever stops matching we just fall back to stage-based completion # in sync. If it ever stops matching we just fall back to stage-based completion
# (~30s slower), so the contract is an optimization, not a correctness requirement. # (~30s slower), so the contract is an optimization, not a correctness requirement.
success_marker = f"Model pushed to https://huggingface.co/{repo_id}" success_marker = f"Model pushed to https://huggingface.co/{repo_id}"
-2
View File
@@ -20,7 +20,6 @@ from .optimizers import (
SGDConfig as SGDConfig, SGDConfig as SGDConfig,
XVLAAdamWConfig as XVLAAdamWConfig, XVLAAdamWConfig as XVLAAdamWConfig,
load_optimizer_state, load_optimizer_state,
load_optimizer_state_dict,
save_optimizer_state, save_optimizer_state,
) )
from .schedulers import ( from .schedulers import (
@@ -51,7 +50,6 @@ __all__ = [
"VQBeTSchedulerConfig", "VQBeTSchedulerConfig",
# State management # State management
"load_optimizer_state", "load_optimizer_state",
"load_optimizer_state_dict",
"load_scheduler_state", "load_scheduler_state",
"save_optimizer_state", "save_optimizer_state",
"save_scheduler_state", "save_scheduler_state",
+14 -29
View File
@@ -27,7 +27,7 @@ from lerobot.utils.constants import (
OPTIMIZER_PARAM_GROUPS, OPTIMIZER_PARAM_GROUPS,
OPTIMIZER_STATE, OPTIMIZER_STATE,
) )
from lerobot.utils.io_utils import deserialize_json_into_object, load_json, write_json from lerobot.utils.io_utils import deserialize_json_into_object, write_json
from lerobot.utils.utils import flatten_dict, unflatten_dict from lerobot.utils.utils import flatten_dict, unflatten_dict
# Type alias for parameters accepted by optimizer build() methods. # Type alias for parameters accepted by optimizer build() methods.
@@ -52,6 +52,11 @@ class OptimizerConfig(draccus.ChoiceRegistry, abc.ABC):
def type(self) -> str: def type(self) -> str:
return self.get_choice_name(self.__class__) return self.get_choice_name(self.__class__)
@property
def builds_multiple_optimizers(self) -> bool:
"""True when build() returns a dict of optimizers (unsupported under sharded training)."""
return False
@classmethod @classmethod
def default_choice_name(cls) -> str | None: def default_choice_name(cls) -> str | None:
return "adam" return "adam"
@@ -245,6 +250,10 @@ class MultiAdamConfig(OptimizerConfig):
grad_clip_norm: float = 10.0 grad_clip_norm: float = 10.0
optimizer_groups: dict[str, dict[str, Any]] = field(default_factory=dict) optimizer_groups: dict[str, dict[str, Any]] = field(default_factory=dict)
@property
def builds_multiple_optimizers(self) -> bool:
return True
def build(self, params: OptimizerParams) -> dict[str, torch.optim.Optimizer]: def build(self, params: OptimizerParams) -> dict[str, torch.optim.Optimizer]:
"""Build multiple Adam optimizers. """Build multiple Adam optimizers.
@@ -283,35 +292,27 @@ class MultiAdamConfig(OptimizerConfig):
def save_optimizer_state( def save_optimizer_state(
optimizer: torch.optim.Optimizer | dict[str, torch.optim.Optimizer], optimizer: torch.optim.Optimizer | dict[str, torch.optim.Optimizer],
save_dir: Path, save_dir: Path,
optim_state_dict: dict | None = None,
) -> None: ) -> None:
"""Save optimizer state to disk. """Save optimizer state to disk (non-sharded runs; sharded runs use the DCP channel).
Args: Args:
optimizer: Either a single optimizer or a dictionary of optimizers. optimizer: Either a single optimizer or a dictionary of optimizers.
save_dir: Directory to save the optimizer state. save_dir: Directory to save the optimizer state.
optim_state_dict: Pre-gathered optimizer state dict (for FSDP, where the sharded state must
be gathered across ranks first). If provided, it is saved directly instead of calling
``optimizer.state_dict()``. Only supported for a single optimizer. Defaults to None.
""" """
if isinstance(optimizer, dict): if isinstance(optimizer, dict):
# Handle dictionary of optimizers # Handle dictionary of optimizers
if optim_state_dict is not None:
raise ValueError("optim_state_dict is not supported for a dict of optimizers")
for name, opt in optimizer.items(): for name, opt in optimizer.items():
optimizer_dir = save_dir / name optimizer_dir = save_dir / name
optimizer_dir.mkdir(exist_ok=True, parents=True) optimizer_dir.mkdir(exist_ok=True, parents=True)
_save_single_optimizer_state(opt, optimizer_dir) _save_single_optimizer_state(opt, optimizer_dir)
else: else:
# Handle single optimizer # Handle single optimizer
_save_single_optimizer_state(optimizer, save_dir, optim_state_dict=optim_state_dict) _save_single_optimizer_state(optimizer, save_dir)
def _save_single_optimizer_state( def _save_single_optimizer_state(optimizer: torch.optim.Optimizer, save_dir: Path) -> None:
optimizer: torch.optim.Optimizer, save_dir: Path, optim_state_dict: dict | None = None
) -> None:
"""Save a single optimizer's state to disk.""" """Save a single optimizer's state to disk."""
state = dict(optim_state_dict) if optim_state_dict is not None else optimizer.state_dict() state = optimizer.state_dict()
param_groups = state.pop("param_groups") param_groups = state.pop("param_groups")
flat_state = flatten_dict(state) flat_state = flatten_dict(state)
save_file(flat_state, save_dir / OPTIMIZER_STATE) save_file(flat_state, save_dir / OPTIMIZER_STATE)
@@ -365,19 +366,3 @@ def _load_single_optimizer_state(optimizer: torch.optim.Optimizer, save_dir: Pat
optimizer.load_state_dict(loaded_state_dict) optimizer.load_state_dict(loaded_state_dict)
return optimizer return optimizer
def load_optimizer_state_dict(save_dir: Path) -> dict:
"""Read a saved optimizer state dict (safetensors + json) back into a plain dict.
Unlike `load_optimizer_state`, this does not load into an optimizer and preserves the original
``state`` keys verbatim (e.g. FSDP parameter FQNs, which are not integer-castable). It is used by
the FSDP resume path, where the full state must be resharded via `FSDP.optim_state_dict_to_load`
before being loaded into the (sharded) optimizer.
"""
flat_state = load_file(save_dir / OPTIMIZER_STATE)
state = unflatten_dict(flat_state)
return {
"state": state.get("state", {}),
"param_groups": load_json(save_dir / OPTIMIZER_PARAM_GROUPS),
}
+2
View File
@@ -47,6 +47,8 @@ class ACTPolicy(PreTrainedPolicy):
config_class = ACTConfig config_class = ACTConfig
name = "act" name = "act"
# FSDP2 wrap units: one unit per transformer layer of both stacks.
_fsdp_wrap_modules = ["ACTEncoderLayer", "ACTDecoderLayer"]
def __init__( def __init__(
self, self,
+29 -16
View File
@@ -242,6 +242,7 @@ def make_policy(
ds_meta: LeRobotDatasetMetadata | None = None, ds_meta: LeRobotDatasetMetadata | None = None,
env_cfg: EnvConfig | None = None, env_cfg: EnvConfig | None = None,
rename_map: dict[str, str] | None = None, rename_map: dict[str, str] | None = None,
defer_weight_load: bool = False,
) -> PreTrainedPolicy: ) -> PreTrainedPolicy:
""" """
Instantiate a policy model. Instantiate a policy model.
@@ -252,22 +253,27 @@ def make_policy(
can either initialize a new policy from scratch or load a pretrained one. can either initialize a new policy from scratch or load a pretrained one.
Args: Args:
cfg: The configuration for the policy to be created. If `cfg.pretrained_path` is cfg (PreTrainedConfig): The configuration for the policy to be created. If
set, the policy will be loaded with weights from that path. `cfg.pretrained_path` is set, the policy will be loaded with weights from that path.
ds_meta: Dataset metadata used to infer feature shapes and types. Also provides ds_meta (LeRobotDatasetMetadata | None): Dataset metadata used to infer feature shapes and
statistics for normalization layers. types. Also provides statistics for normalization layers.
env_cfg: Environment configuration used to infer feature shapes and types. env_cfg (EnvConfig | None): Environment configuration used to infer feature shapes and
One of `ds_meta` or `env_cfg` must be provided. types. One of `ds_meta` or `env_cfg` must be provided.
rename_map: Optional mapping of dataset or environment feature keys to match rename_map (dict[str, str] | None): Optional mapping of dataset or environment feature
expected policy feature names (e.g., `"left"` → `"camera1"`). keys to match expected policy feature names (e.g., `"left"` → `"camera1"`).
defer_weight_load (bool): Build the exact policy `from_pretrained` would build — same
config resolution, same stats-derived buffers, same device placement and eval mode —
but skip the safetensors weight load. Used when resuming from a DCP checkpoint, whose
sharded weights stream in after `accelerator.prepare()` (the distributed checkpoint
engine overwrites the random init).
Returns: Returns:
An instantiated and device-placed policy model. PreTrainedPolicy: An instantiated and device-placed policy model.
Raises: Raises:
ValueError: If both or neither of `ds_meta` and `env_cfg` are provided. ValueError: If both or neither of `ds_meta` and `env_cfg` are provided.
NotImplementedError: If attempting to use an unsupported policy-backend NotImplementedError: If attempting to use an unsupported policy-backend combination
combination (e.g., VQBeT with 'mps'). (e.g., VQBeT with 'mps').
""" """
if bool(ds_meta) == bool(env_cfg): if bool(ds_meta) == bool(env_cfg):
raise ValueError("Either one of a dataset metadata or a sim env must be provided.") raise ValueError("Either one of a dataset metadata or a sim env must be provided.")
@@ -332,11 +338,18 @@ def make_policy(
) )
if cfg.pretrained_path and not cfg.use_peft: if cfg.pretrained_path and not cfg.use_peft:
# Load a pretrained policy and override the config if needed (for example, if there are inference-time if defer_weight_load:
# hyperparameters that we want to vary). # Same construction path as from_pretrained (config already resolved from the
kwargs["pretrained_name_or_path"] = cfg.pretrained_path # checkpoint by the caller; dataset_stats/dataset_meta kwargs identical), minus the
kwargs["revision"] = cfg.pretrained_revision # weight load — parity by construction.
policy = policy_cls.from_pretrained(**kwargs) policy = policy_cls(**kwargs)
policy.eval()
else:
# Load a pretrained policy and override the config if needed (for example, if there
# are inference-time hyperparameters that we want to vary).
kwargs["pretrained_name_or_path"] = cfg.pretrained_path
kwargs["revision"] = cfg.pretrained_revision
policy = policy_cls.from_pretrained(**kwargs)
elif cfg.pretrained_path and cfg.use_peft: elif cfg.pretrained_path and cfg.use_peft:
# Load a pretrained PEFT model on top of the policy. The pretrained path points to the folder/repo # Load a pretrained PEFT model on top of the policy. The pretrained path points to the folder/repo
# of the adapter and the adapter's config contains the path to the base policy. So we need the # of the adapter and the adapter's config contains the path to the base policy. So we need the
@@ -54,6 +54,9 @@ class FastWAMPolicy(PreTrainedPolicy):
config_class = FastWAMConfig config_class = FastWAMConfig
name = "fastwam" name = "fastwam"
# FSDP2 wrap units: MoTLayer is the single FSDP owner of each layer's expert blocks
# (the blocks are re-parented onto it precisely so sharding has one boundary to hook).
_fsdp_wrap_modules = ["MoTLayer"]
def __init__( def __init__(
self, self,
+75 -167
View File
@@ -18,20 +18,17 @@ import builtins
import dataclasses import dataclasses
import logging import logging
import os import os
from importlib.resources import files import warnings
from pathlib import Path from pathlib import Path
from tempfile import TemporaryDirectory from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, TypeVar, Unpack
from typing import TYPE_CHECKING, TypedDict, TypeVar, Unpack
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download, save_torch_state_dict from huggingface_hub import hf_hub_download, save_torch_state_dict
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
from huggingface_hub.errors import HfHubHTTPError from huggingface_hub.errors import HfHubHTTPError
from safetensors.torch import load_model as load_model_as_safetensor, save_model as save_model_as_safetensor from safetensors.torch import load_model as load_model_as_safetensor
from torch import Tensor, nn from torch import Tensor, nn
from lerobot.__version__ import __version__
from lerobot.configs import PreTrainedConfig from lerobot.configs import PreTrainedConfig
from lerobot.configs.train import TrainPipelineConfig
from lerobot.utils.device_utils import resolve_safetensors_device from lerobot.utils.device_utils import resolve_safetensors_device
from lerobot.utils.hub import HubMixin from lerobot.utils.hub import HubMixin
from lerobot.utils.import_utils import _peft_available, require_package from lerobot.utils.import_utils import _peft_available, require_package
@@ -46,56 +43,14 @@ else:
get_peft_model = None get_peft_model = None
if TYPE_CHECKING: if TYPE_CHECKING:
from lerobot.configs.train import TrainPipelineConfig
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
T = TypeVar("T", bound="PreTrainedPolicy") T = TypeVar("T", bound="PreTrainedPolicy")
# Pinned far above any policy's total size so save_torch_state_dict always emits exactly one
def _build_card_context( # `model.safetensors` (no shards, no index) — a constant, not a computed byte count.
cfg: TrainPipelineConfig | None, _SINGLE_FILE_SHARD_SIZE = "1TB"
dataset_meta: LeRobotDatasetMetadata | None,
input_features: dict | None,
output_features: dict | None,
) -> dict:
"""Collect optional data for the model-card template.
Returns plain values only (no Markdown) — the template in
``lerobot/templates/lerobot_modelcard_template.md`` decides how and whether to show
each one. Everything is best-effort: anything unavailable is left empty/None and the
template simply skips that section, so this never breaks a Hub push.
"""
context = {
"training": None,
"input_features": input_features or {},
"output_features": output_features or {},
"dataset": None,
"robot_type": None,
"cameras": [],
}
if cfg is not None:
optimizer = getattr(cfg, "optimizer", None)
context["training"] = {
"steps": cfg.steps,
"batch_size": cfg.batch_size,
"seed": cfg.seed,
"optimizer": getattr(optimizer, "type", None) if optimizer else None,
"lr": getattr(optimizer, "lr", None) if optimizer else None,
"lerobot_version": __version__,
}
if dataset_meta is not None:
context["dataset"] = {
"repo_id": dataset_meta.repo_id,
"episodes": dataset_meta.total_episodes,
"frames": dataset_meta.total_frames,
"fps": dataset_meta.fps,
"tasks": [str(task) for task in dataset_meta.tasks.index],
}
context["robot_type"] = dataset_meta.robot_type
context["cameras"] = [key.split(".")[-1] for key in dataset_meta.camera_keys]
return context
class ActionSelectKwargs(TypedDict, total=False): class ActionSelectKwargs(TypedDict, total=False):
@@ -110,6 +65,22 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
config_class: None config_class: None
name: None name: None
# --- declarative parallelism/acceleration surface ----------------------------------------
# Module CLASS names forming the FSDP2 wrap units (and, once wired, the activation-
# checkpointing units). Resolved onto the accelerate plugin right before
# `accelerator.prepare()` by `lerobot.distributed.set_fsdp_wrap_modules`; sharded training
# with no wrap source anywhere fails loudly instead of silently wrapping only the root.
_fsdp_wrap_modules: ClassVar[list[str] | None] = None
# Non-`forward` entry points that must trigger FSDP2 unshard/reshard hooks when called on a
# sharded policy (registered post-prepare via `torch.distributed.fsdp
# .register_fsdp_forward_method`); calling them unregistered crashes on mixed Tensor/DTensor.
_fsdp_forward_methods: ClassVar[tuple[str, ...]] = ("select_action", "predict_action_chunk")
# Capability gate for the (future) activation-checkpointing wiring.
supports_gradient_checkpointing: ClassVar[bool] = False
# Declarative context-parallel plan (diffusers `ContextParallelModelPlan` semantics:
# module FQN -> sequence split/gather spec). Reserved for the CP engine round.
_cp_plan: ClassVar[dict[str, Any] | None] = None
def __init__(self, config: PreTrainedConfig, *inputs, **kwargs): def __init__(self, config: PreTrainedConfig, *inputs, **kwargs):
super().__init__() super().__init__()
if not isinstance(config, PreTrainedConfig): if not isinstance(config, PreTrainedConfig):
@@ -127,43 +98,33 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
if not getattr(cls, "name", None): if not getattr(cls, "name", None):
raise TypeError(f"Class {cls.__name__} must define 'name'") raise TypeError(f"Class {cls.__name__} must define 'name'")
def save_pretrained( def _save_pretrained(self, save_directory: Path) -> None:
self, """Serialize this policy's parameters (and config) into `save_directory`.
save_directory: str | Path,
*,
state_dict: dict[str, Tensor] | None = None,
repo_id: str | None = None,
push_to_hub: bool = False,
card_kwargs: dict | None = None,
**push_to_hub_kwargs,
) -> str | None:
"""Save the policy to a directory (and optionally push to the Hub).
Overrides `HubMixin.save_pretrained` to add a `state_dict` argument (mirroring Sharding is handled internally: under FSDP2 the full state dict is gathered through a
`transformers.PreTrainedModel.save_pretrained`). Under FSDP, `self.state_dict()` would COLLECTIVE, so when the policy is sharded this method (via `save_pretrained`) must be
return sharded tensors, so the caller gathers the full state dict via a cross-rank called on EVERY rank — a rank-0-gated call deadlocks. File writes happen on the main
collective and passes it here for `_save_pretrained` to write directly. process only, in all layouts (single, DDP, sharded).
Args:
save_directory (Path): Target directory for the policy config (`config.json`) and the
safetensors weight file(s).
""" """
save_directory = Path(save_directory) # Lazy imports: the persistence layer pulls in lerobot.distributed only when saving.
save_directory.mkdir(parents=True, exist_ok=True) from lerobot.distributed.checkpoint import full_model_state_dict, is_sharded_module
self._save_pretrained(save_directory, state_dict=state_dict) from lerobot.distributed.utils import is_main_process
if push_to_hub:
if repo_id is None:
repo_id = save_directory.name
return self.push_to_hub(repo_id=repo_id, card_kwargs=card_kwargs, **push_to_hub_kwargs)
return None
def _save_pretrained(self, save_directory: Path, state_dict: dict[str, Tensor] | None = None) -> None:
self.config._save_pretrained(save_directory)
model_to_save = self.module if hasattr(self, "module") else self model_to_save = self.module if hasattr(self, "module") else self
if state_dict is None: if is_sharded_module(model_to_save):
save_model_as_safetensor(model_to_save, str(save_directory / SAFETENSORS_SINGLE_FILE)) logging.info("Gathering the full state dict from all ranks (sharded policy).")
state_dict = full_model_state_dict(model_to_save) # collective when sharded; {} off-main
if not state_dict or not is_main_process():
# Sharded: the gather materializes on the main rank only (emptiness check).
# Non-sharded multi-rank (DDP): every rank holds a full dict — the explicit rank
# gate prevents N ranks racing on the same files. Single process: never taken.
return return
# A pre-gathered (e.g. FSDP full) state dict was supplied: write it directly. self.config._save_pretrained(save_directory)
# `save_torch_state_dict` discards shared-tensor duplicates just like `save_model` does; save_torch_state_dict(state_dict, str(save_directory), max_shard_size=_SINGLE_FILE_SHARD_SIZE)
# pin `max_shard_size` above the total size so the output stays a single `model.safetensors`
total_bytes = sum(t.numel() * t.element_size() for t in state_dict.values())
save_torch_state_dict(state_dict, str(save_directory), max_shard_size=max(total_bytes, 1))
@classmethod @classmethod
def from_pretrained( def from_pretrained(
@@ -291,92 +252,39 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
peft_model=None, peft_model=None,
state_dict: dict[str, Tensor] | None = None, state_dict: dict[str, Tensor] | None = None,
dataset_meta: LeRobotDatasetMetadata | None = None, dataset_meta: LeRobotDatasetMetadata | None = None,
): ) -> None:
api = HfApi() """Publish this policy to the Hub.
repo_id = api.create_repo(
repo_id=self.config.repo_id, private=self.config.private, exist_ok=True
).repo_id
# Push the files to the repo in a single commit Deprecated: use :func:`lerobot.common.train_utils.publish_trained_model` instead, which
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp: also publishes the pre/post-processors alongside the model.
saved_path = Path(tmp) / repo_id
if peft_model is not None: Args:
# Since PEFT just forwards calls to `push_model_to_hub`, `self` is not the PeftModel wrapper cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and
# but the actual policy which is why we need the PEFT model passed to us to save the adapter. used to render the model card.
# That also means that we need to store the policy config ourselves since PEFT can't. peft_model: The PEFT wrapper when training adapters, whose weights replace the full
peft_model.save_pretrained(saved_path) model weights in the published repo. Defaults to None.
self.config.save_pretrained(saved_path) state_dict (dict[str, Tensor] | None): Ignored; weights are now gathered internally
else: when the policy is sharded. Defaults to None.
# Calls _save_pretrained and stores model tensors dataset_meta (LeRobotDatasetMetadata | None): Dataset metadata for the model card,
self.save_pretrained(saved_path, state_dict=state_dict) if available. Defaults to None.
"""
from lerobot.common.train_utils import publish_trained_model
card = self.generate_model_card( warnings.warn(
cfg.dataset.repo_id, "PreTrainedPolicy.push_model_to_hub is deprecated and will be removed in a future "
self.config.type, "version. Use lerobot.common.train_utils.publish_trained_model(cfg, model, "
self.config.license, "preprocessor, postprocessor, dataset_meta) instead.",
self.config.tags, FutureWarning,
cfg=cfg, stacklevel=2,
dataset_meta=dataset_meta, )
if state_dict is not None:
warnings.warn(
"The `state_dict` argument is ignored: sharded weights are gathered internally "
"when the policy is saved.",
FutureWarning,
stacklevel=2,
) )
card.save(str(saved_path / "README.md")) publish_trained_model(cfg, self, None, None, dataset_meta, peft_model=peft_model)
cfg.save_pretrained(saved_path) # Calls _save_pretrained and stores train config
commit_info = api.upload_folder(
repo_id=repo_id,
repo_type="model",
folder_path=saved_path,
commit_message="Upload policy weights, train config and readme",
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
ignore_patterns=["*.tmp", "*.log"],
)
# Contract: lerobot.jobs.hf.submit_to_hf watches for this exact
# "Model pushed to <url>" line to end a remote run early. Keep the wording
# and URL format in sync (it falls back to status polling if they drift).
logging.info(f"Model pushed to {commit_info.repo_url.url}")
def generate_model_card(
self,
dataset_repo_id: str,
model_type: str,
license: str | None,
tags: list[str] | None,
cfg: TrainPipelineConfig | None = None,
dataset_meta: LeRobotDatasetMetadata | None = None,
) -> ModelCard:
base_model_mapping = {
"smolvla": "lerobot/smolvla_base",
"pi0": "lerobot/pi0_base",
"pi05": "lerobot/pi05_base",
"pi0_fast": "lerobot/pi0fast-base",
"xvla": "lerobot/xvla-base",
}
card_data = ModelCardData(
license=license or "apache-2.0",
library_name="lerobot",
pipeline_tag="robotics",
tags=list(set(tags or []).union({"robotics", "lerobot", model_type})),
model_name=model_type,
datasets=dataset_repo_id,
base_model=base_model_mapping.get(model_type),
)
context = _build_card_context(
cfg, dataset_meta, self.config.input_features, self.config.output_features
)
# Used by the template to pre-fill commands and the "Fine-tuned from" line.
context["policy_repo_id"] = getattr(self.config, "repo_id", None)
context["base_model"] = base_model_mapping.get(model_type)
template_card = (
files("lerobot.templates").joinpath("lerobot_modelcard_template.md").read_text(encoding="utf-8")
)
card = ModelCard.from_template(card_data, template_str=template_card, **context)
card.validate()
return card
def wrap_with_peft( def wrap_with_peft(
self, self,
@@ -647,10 +647,15 @@ def main():
tags = set(tags).union({"robotics", "lerobot", policy_type}) tags = set(tags).union({"robotics", "lerobot", policy_type})
tags = list(tags) tags = list(tags)
# Generate model card # Generate model card through the free helper (PreTrainedPolicy.generate_model_card was
card = policy.generate_model_card( # removed with the publisher redesign), then apply the metadata recovered above — the
dataset_repo_id=dataset_repo_id, model_type=policy_type, license=license, tags=tags # migrated policy config does not carry the original repo's card fields.
) from lerobot.common.train_utils import generate_model_card
card = generate_model_card(policy.config)
card.data.datasets = dataset_repo_id
card.data.license = license
card.data.tags = sorted(tags)
# Save model card locally # Save model card locally
card.save(str(output_dir / "README.md")) card.save(str(output_dir / "README.md"))
+33 -49
View File
@@ -16,12 +16,11 @@ import abc
import builtins import builtins
import logging import logging
import os import os
from importlib.resources import files import warnings
from pathlib import Path from pathlib import Path
from tempfile import TemporaryDirectory
from typing import TYPE_CHECKING, Any, TypeVar from typing import TYPE_CHECKING, Any, TypeVar
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download from huggingface_hub import hf_hub_download
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
from huggingface_hub.errors import HfHubHTTPError from huggingface_hub.errors import HfHubHTTPError
from safetensors.torch import load_model as load_model_as_safetensor, save_model as save_model_as_safetensor from safetensors.torch import load_model as load_model_as_safetensor, save_model as save_model_as_safetensor
@@ -61,6 +60,22 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
raise TypeError(f"Class {cls.__name__} must define 'name'") raise TypeError(f"Class {cls.__name__} must define 'name'")
def _save_pretrained(self, save_directory: Path) -> None: def _save_pretrained(self, save_directory: Path) -> None:
"""Serialize this reward model's parameters (and config) into `save_directory`.
Safe to call on every rank: replicas carry identical weights, so only the main process
writes (sharded reward models are rejected at config validation — no collective gather).
Args:
save_directory (Path): Target directory for the reward model config (`config.json`)
and `model.safetensors`.
"""
from lerobot.distributed.utils import is_main_process
# save_checkpoint calls this on every rank; replicas carry identical
# weights, so the main process is the only writer. Sharded reward models are rejected
# at config validation, so no collective gather is needed here.
if not is_main_process():
return
self.config._save_pretrained(save_directory) self.config._save_pretrained(save_directory)
model_to_save = self.module if hasattr(self, "module") else self model_to_save = self.module if hasattr(self, "module") else self
save_model_as_safetensor(model_to_save, str(save_directory / SAFETENSORS_SINGLE_FILE)) save_model_as_safetensor(model_to_save, str(save_directory / SAFETENSORS_SINGLE_FILE))
@@ -175,53 +190,22 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
""" """
return type(self).forward is not PreTrainedRewardModel.forward return type(self).forward is not PreTrainedRewardModel.forward
def push_model_to_hub(self, cfg: "TrainPipelineConfig"): def push_model_to_hub(self, cfg: "TrainPipelineConfig") -> None:
api = HfApi() """Publish this reward model to the Hub.
repo_id = api.create_repo(
repo_id=self.config.repo_id, private=self.config.private, exist_ok=True
).repo_id
# Push the files to the repo in a single commit Deprecated: use :func:`lerobot.common.train_utils.publish_trained_model` instead.
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
saved_path = Path(tmp) / repo_id
self.save_pretrained(saved_path) # Calls _save_pretrained and stores model tensors Args:
cfg (TrainPipelineConfig): The training config; saved as `train_config.json` and
used to render the model card.
"""
from lerobot.common.train_utils import publish_trained_model
card = self.generate_model_card( warnings.warn(
cfg.dataset.repo_id, self.config.type, self.config.license, self.config.tags "PreTrainedRewardModel.push_model_to_hub is deprecated and will be removed in a "
) "future version. Use lerobot.common.train_utils.publish_trained_model(cfg, model, "
card.save(str(saved_path / "README.md")) "preprocessor, postprocessor, dataset_meta) instead.",
FutureWarning,
cfg.save_pretrained(saved_path) # Calls _save_pretrained and stores train config stacklevel=2,
commit_info = api.upload_folder(
repo_id=repo_id,
repo_type="model",
folder_path=saved_path,
commit_message="Upload reward model weights, train config and readme",
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
ignore_patterns=["*.tmp", "*.log"],
)
logging.info(f"Model pushed to {commit_info.repo_url.url}")
def generate_model_card(
self, dataset_repo_id: str, model_type: str, license: str | None, tags: list[str] | None
) -> ModelCard:
card_data = ModelCardData(
license=license or "apache-2.0",
library_name="lerobot",
pipeline_tag="robotics",
tags=list(set(tags or []).union({"robotics", "lerobot", "reward-model", model_type})),
model_name=model_type,
datasets=dataset_repo_id,
) )
publish_trained_model(cfg, self, None, None, None)
template_card = (
files("lerobot.templates")
.joinpath("lerobot_rewardmodel_modelcard_template.md")
.read_text(encoding="utf-8")
)
card = ModelCard.from_template(card_data, template_str=template_card)
card.validate()
return card
@@ -58,12 +58,11 @@ import builtins
import logging import logging
import os import os
from pathlib import Path from pathlib import Path
from tempfile import TemporaryDirectory
from typing import TYPE_CHECKING, Any, TypeVar from typing import TYPE_CHECKING, Any, TypeVar
import numpy as np import numpy as np
import torch import torch
from huggingface_hub import HfApi, hf_hub_download from huggingface_hub import hf_hub_download
from huggingface_hub.constants import CONFIG_NAME from huggingface_hub.constants import CONFIG_NAME
from huggingface_hub.errors import HfHubHTTPError from huggingface_hub.errors import HfHubHTTPError
from torch import Tensor from torch import Tensor
@@ -75,9 +74,6 @@ from lerobot.rewards.topreward.configuration_topreward import TOPRewardConfig
from lerobot.rewards.topreward.processor_topreward import TOPREWARD_FEATURE_PREFIX, TOPREWARD_INPUT_KEYS from lerobot.rewards.topreward.processor_topreward import TOPREWARD_FEATURE_PREFIX, TOPREWARD_INPUT_KEYS
from lerobot.utils.import_utils import _transformers_available, require_package from lerobot.utils.import_utils import _transformers_available, require_package
if TYPE_CHECKING:
from lerobot.configs.train import TrainPipelineConfig
if TYPE_CHECKING or _transformers_available: if TYPE_CHECKING or _transformers_available:
from transformers import Qwen3VLForConditionalGeneration from transformers import Qwen3VLForConditionalGeneration
else: else:
@@ -205,34 +201,3 @@ class TOPRewardModel(PreTrainedRewardModel):
instance.to(config.device) instance.to(config.device)
instance.eval() instance.eval()
return instance return instance
def push_model_to_hub(self, cfg: TrainPipelineConfig):
"""Push the TOPReward ``config.json`` + model card to the Hub."""
api = HfApi()
repo_id = api.create_repo(
repo_id=self.config.repo_id, private=self.config.private, exist_ok=True
).repo_id
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
saved_path = Path(tmp) / repo_id
saved_path.mkdir(parents=True, exist_ok=True)
self.config._save_pretrained(saved_path)
card = self.generate_model_card(
cfg.dataset.repo_id, self.config.type, self.config.license, self.config.tags
)
card.save(str(saved_path / "README.md"))
cfg.save_pretrained(saved_path)
commit_info = api.upload_folder(
repo_id=repo_id,
repo_type="model",
folder_path=saved_path,
commit_message="Upload TOPReward config and readme",
allow_patterns=["*.json", "*.yaml", "*.md"],
ignore_patterns=["*.tmp", "*.log", "*.safetensors"],
)
logger.info(f"Model pushed to {commit_info.repo_url.url}")
+17 -10
View File
@@ -74,13 +74,14 @@ from torch.optim.optimizer import Optimizer
from lerobot.cameras import opencv # noqa: F401 from lerobot.cameras import opencv # noqa: F401
from lerobot.common.train_utils import ( from lerobot.common.train_utils import (
get_step_checkpoint_dir, get_step_checkpoint_dir,
load_training_state as utils_load_training_state, load_training_metadata,
save_checkpoint, save_checkpoint,
update_last_checkpoint, update_last_checkpoint,
) )
from lerobot.common.wandb_utils import WandBLogger from lerobot.common.wandb_utils import WandBLogger
from lerobot.configs import parser from lerobot.configs import parser
from lerobot.datasets import LeRobotDataset, make_dataset from lerobot.datasets import LeRobotDataset, make_dataset
from lerobot.optim import load_optimizer_state
from lerobot.policies import make_policy, make_pre_post_processors from lerobot.policies import make_policy, make_pre_post_processors
from lerobot.robots import so_follower # noqa: F401 from lerobot.robots import so_follower # noqa: F401
from lerobot.teleoperators import gamepad, so_leader # noqa: F401 from lerobot.teleoperators import gamepad, so_leader # noqa: F401
@@ -103,7 +104,7 @@ from lerobot.utils.constants import (
from lerobot.utils.device_utils import get_safe_torch_device from lerobot.utils.device_utils import get_safe_torch_device
from lerobot.utils.io_utils import load_json, write_json from lerobot.utils.io_utils import load_json, write_json
from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method
from lerobot.utils.random_utils import set_seed from lerobot.utils.random_utils import load_rng_state, set_seed
from lerobot.utils.utils import ( from lerobot.utils.utils import (
format_big_number, format_big_number,
init_logging, init_logging,
@@ -716,15 +717,18 @@ def load_training_state(
algorithm-owned tensors) from the most recent checkpoint. algorithm-owned tensors) from the most recent checkpoint.
Args: Args:
cfg: Training configuration. cfg (TrainRLServerPipelineConfig): Training configuration; `cfg.resume` gates the load and
optimizers: Optimizers to load state into. `cfg.output_dir` locates the last checkpoint.
algorithm: Algorithm whose state dict should be restored. optimizers (Optimizer | dict[str, Optimizer]): Optimizers to load state into.
Required for full main-equivalent resume; algorithm (RLAlgorithm | None, optional): Algorithm whose state dict should be restored.
the policy itself is restored separately via ``make_policy``. Required for full main-equivalent resume; the policy itself is restored separately via
device: 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: Returns:
tuple: (optimization_step, interaction_step) or (None, None) if not resuming tuple[int | None, int | None]: `(optimization_step, interaction_step)`, or `(None, None)`
when not resuming or when loading the training state fails.
""" """
if not cfg.resume: if not cfg.resume:
return None, None return None, None
@@ -736,7 +740,10 @@ def load_training_state(
try: try:
# Restore optimizers + RNG + step from the standard `training_state/` folder # Restore optimizers + RNG + step from the standard `training_state/` folder
step, optimizers, _ = utils_load_training_state(checkpoint_dir, optimizers, None) training_state_dir = checkpoint_dir / TRAINING_STATE_DIR
load_rng_state(training_state_dir)
step = load_training_metadata(training_state_dir)["step"]
optimizers = load_optimizer_state(optimizers, training_state_dir)
# Restore algorithm-owned tensors # Restore algorithm-owned tensors
if algorithm is not None: if algorithm is not None:
@@ -25,6 +25,11 @@ quantile statistics (q01, q10, q50, q90, q99) in their metadata. This script:
3. If missing, computes quantile statistics for all features 3. If missing, computes quantile statistics for all features
4. Updates the dataset metadata with the new quantile statistics 4. Updates the dataset metadata with the new quantile statistics
Statistics are accumulated into a single running histogram per feature across
all episodes rather than aggregating per-episode quantile summaries. The
resulting quantiles are histogram approximations, subject to discretization and
range-rebinning error; image/video frames are sampled by default.
Usage: Usage:
```bash ```bash
@@ -34,9 +39,7 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
""" """
import argparse import argparse
import concurrent.futures
import logging import logging
import os
from pathlib import Path from pathlib import Path
import numpy as np import numpy as np
@@ -49,11 +52,10 @@ from lerobot.datasets import (
CODEBASE_VERSION, CODEBASE_VERSION,
DEFAULT_QUANTILES, DEFAULT_QUANTILES,
LeRobotDataset, LeRobotDataset,
aggregate_stats,
get_feature_stats, get_feature_stats,
write_stats, write_stats,
) )
from lerobot.datasets.compute_stats import sample_indices from lerobot.datasets.compute_stats import RunningQuantileStats, sample_indices
from lerobot.utils.utils import init_logging from lerobot.utils.utils import init_logging
@@ -79,20 +81,25 @@ def has_quantile_stats(stats: dict[str, dict] | None, quantile_list_keys: list[s
return False return False
def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampling: bool = True) -> dict: def collect_episode_arrays(
"""Process a single episode and return its statistics. dataset: LeRobotDataset,
episode_idx: int,
use_sampling: bool = True,
skip_images: bool = False,
) -> dict[str, tuple[np.ndarray, int]]:
"""Collect one episode's frames per feature, flattened to (num_samples, dim).
Args: Args:
dataset: The LeRobot dataset dataset: The LeRobot dataset
episode_idx: Index of the episode to process episode_idx: Index of the episode to read
use_sampling: If True, sub-sample image/video frames per episode to bound use_sampling: If True, sub-sample image/video frames to bound memory.
memory. If False, use every frame (exact, higher memory). If False, use every frame (higher memory).
skip_images: If True, skip image/video features entirely.
Returns: Returns:
Dictionary containing episode statistics Mapping of feature name to that episode's values and the number of frames
they came from (which differs from the row count for image features).
""" """
logging.info(f"Computing stats for episode {episode_idx}")
start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"] start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"]
end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"] end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"]
@@ -102,7 +109,9 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
# numeric columns are cheap, so read them in full (exact). # numeric columns are cheap, so read them in full (exact).
image_keys = [k for k in dataset.features if dataset.features[k]["dtype"] in ("image", "video")] image_keys = [k for k in dataset.features if dataset.features[k]["dtype"] in ("image", "video")]
numeric_keys = [ numeric_keys = [
k for k in dataset.features if dataset.features[k]["dtype"] not in ("image", "video", "string") k
for k in dataset.features
if dataset.features[k]["dtype"] not in ("image", "video", "string", "language")
] ]
collected_data: dict[str, list] = {} collected_data: dict[str, list] = {}
@@ -114,7 +123,7 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
collected_data[key] = [torch.as_tensor(v) for v in numeric_cols[key]] collected_data[key] = [torch.as_tensor(v) for v in numeric_cols[key]]
# Image/video features: decode only a sampled subset of frames. # Image/video features: decode only a sampled subset of frames.
if image_keys: if image_keys and not skip_images:
sampled_offsets = sample_indices(episode_len) if use_sampling else list(range(episode_len)) sampled_offsets = sample_indices(episode_len) if use_sampling else list(range(episode_len))
for offset in sampled_offsets: for offset in sampled_offsets:
item = dataset[start_idx + offset] item = dataset[start_idx + offset]
@@ -122,87 +131,82 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
if key in item: if key in item:
collected_data.setdefault(key, []).append(item[key]) collected_data.setdefault(key, []).append(item[key])
ep_stats = {} episode_arrays: dict[str, tuple[np.ndarray, int]] = {}
for key, data_list in collected_data.items(): for key, data_list in collected_data.items():
if dataset.features[key]["dtype"] == "string":
continue
data = torch.stack(data_list).cpu().numpy() data = torch.stack(data_list).cpu().numpy()
if dataset.features[key]["dtype"] in ["image", "video"]: if dataset.features[key]["dtype"] in ["image", "video"]:
if data.dtype == np.uint8: if data.dtype == np.uint8:
data = data.astype(np.float32) / 255.0 data = data.astype(np.float32) / 255.0
# (N, C, H, W) -> (N * H * W, C) so quantiles are computed per channel.
axes_to_reduce = (0, 2, 3) channels = data.shape[1]
keepdims = True values = data.transpose(0, 2, 3, 1).reshape(-1, channels)
else: else:
axes_to_reduce = 0 values = data.reshape(-1, data.shape[-1]) if data.ndim > 1 else data.reshape(-1, 1)
keepdims = data.ndim == 1 episode_arrays[key] = (values, len(data_list))
ep_stats[key] = get_feature_stats( return episode_arrays
data, axis=axes_to_reduce, keepdims=keepdims, quantile_list=DEFAULT_QUANTILES
)
if dataset.features[key]["dtype"] in ["image", "video"]:
ep_stats[key] = {
k: v if k == "count" else np.squeeze(v, axis=0) for k, v in ep_stats[key].items()
}
return ep_stats
def compute_quantile_stats_for_dataset(dataset: LeRobotDataset, use_sampling: bool = True) -> dict[str, dict]: def compute_quantile_stats_for_dataset(
"""Compute quantile statistics for all episodes in the dataset. dataset: LeRobotDataset,
use_sampling: bool = True,
skip_images: bool = False,
) -> dict[str, dict]:
"""Compute whole-dataset statistics with one running histogram per feature.
Args: Args:
dataset: The LeRobot dataset to compute statistics for dataset: The LeRobot dataset to compute statistics for
use_sampling: If True, sub-sample image/video frames per episode to bound use_sampling: If True, sub-sample image/video frames per episode to bound
memory. If False, use every frame (exact, higher memory). memory. If False, use every frame (higher memory).
skip_images: If True, skip image/video features and leave their stats untouched.
Returns: Returns:
Dictionary containing aggregated statistics with quantiles Dictionary containing statistics with histogram-based global quantile estimates
Note: Note:
Video decoding operations are not thread-safe, so we process episodes sequentially Episodes are accumulated sequentially because the running accumulators are
when video keys are present. For datasets without videos, we use parallel processing shared across all of them.
with ThreadPoolExecutor for better performance.
""" """
logging.info(f"Computing quantile statistics for dataset with {dataset.num_episodes} episodes") logging.info(f"Computing quantile statistics for dataset with {dataset.num_episodes} episodes")
episode_stats_list = [] running_stats: dict[str, RunningQuantileStats] = {}
has_videos = len(dataset.meta.video_keys) > 0 frame_counts: dict[str, int] = {}
row_counts: dict[str, int] = {}
# Kept only while a feature has a single row, so it can still be finalized.
single_row_arrays: dict[str, np.ndarray] = {}
if has_videos: for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"):
logging.info("Dataset contains video keys - using sequential processing for thread safety") episode_arrays = collect_episode_arrays(
for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"): dataset, episode_idx, use_sampling=use_sampling, skip_images=skip_images
ep_stats = process_single_episode(dataset, episode_idx, use_sampling) )
episode_stats_list.append(ep_stats) for key, (array, num_frames) in episode_arrays.items():
else: running_stats.setdefault(key, RunningQuantileStats()).update(array)
logging.info("Dataset has no video keys - using parallel processing for better performance") frame_counts[key] = frame_counts.get(key, 0) + num_frames
max_workers = min(dataset.num_episodes, int(os.environ.get("LEROBOT_STATS_MAX_WORKERS", 16))) row_counts[key] = row_counts.get(key, 0) + len(array)
if row_counts[key] < 2:
single_row_arrays[key] = array
else:
single_row_arrays.pop(key, None)
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor: if not running_stats:
future_to_episode = {
executor.submit(process_single_episode, dataset, episode_idx, use_sampling): episode_idx
for episode_idx in range(dataset.num_episodes)
}
episode_results = {}
with tqdm(total=dataset.num_episodes, desc="Processing episodes") as pbar:
for future in concurrent.futures.as_completed(future_to_episode):
episode_idx = future_to_episode[future]
ep_stats = future.result()
episode_results[episode_idx] = ep_stats
pbar.update(1)
for episode_idx in range(dataset.num_episodes):
if episode_idx in episode_results:
episode_stats_list.append(episode_results[episode_idx])
if not episode_stats_list:
raise ValueError("No episode data found for computing statistics") raise ValueError("No episode data found for computing statistics")
logging.info(f"Aggregating statistics from {len(episode_stats_list)} episodes") aggregated_stats: dict[str, dict] = {}
return aggregate_stats(episode_stats_list) for key, accumulator in running_stats.items():
if row_counts[key] < 2:
# Histograms need at least two samples; mirror get_feature_stats' basic-stats path.
stats = get_feature_stats(single_row_arrays[key], axis=0, keepdims=False)
else:
stats = accumulator.get_statistics()
if dataset.features[key]["dtype"] in ["image", "video"]:
# Image stats are stored as (C, 1, 1) to broadcast over height and width.
stats = {k: v if k == "count" else v[:, np.newaxis, np.newaxis] for k, v in stats.items()}
# `get_feature_stats` counts frames, not the per-channel rows the accumulator sees.
stats["count"] = np.array([frame_counts[key]])
aggregated_stats[key] = stats
logging.info(f"Computed global histogram statistics for {len(aggregated_stats)} features")
return aggregated_stats
def augment_dataset_with_quantile_stats( def augment_dataset_with_quantile_stats(
@@ -210,6 +214,7 @@ def augment_dataset_with_quantile_stats(
root: str | Path | None = None, root: str | Path | None = None,
overwrite: bool = False, overwrite: bool = False,
use_sampling: bool = True, use_sampling: bool = True,
skip_images: bool = False,
) -> None: ) -> None:
"""Augment a dataset with quantile statistics if they are missing. """Augment a dataset with quantile statistics if they are missing.
@@ -218,7 +223,8 @@ def augment_dataset_with_quantile_stats(
root: Local root directory for the dataset root: Local root directory for the dataset
overwrite: Overwrite existing quantile statistics if they already exist overwrite: Overwrite existing quantile statistics if they already exist
use_sampling: If True, sub-sample image/video frames per episode to bound use_sampling: If True, sub-sample image/video frames per episode to bound
memory. If False, use every frame (exact, higher memory). memory. If False, use every frame (higher memory).
skip_images: If True, skip image/video features and keep their existing stats
""" """
logging.info(f"Loading dataset: {repo_id}") logging.info(f"Loading dataset: {repo_id}")
dataset = LeRobotDataset( dataset = LeRobotDataset(
@@ -232,7 +238,13 @@ def augment_dataset_with_quantile_stats(
logging.info("Dataset does not contain quantile statistics. Computing them now...") logging.info("Dataset does not contain quantile statistics. Computing them now...")
new_stats = compute_quantile_stats_for_dataset(dataset, use_sampling=use_sampling) new_stats = compute_quantile_stats_for_dataset(
dataset, use_sampling=use_sampling, skip_images=skip_images
)
if skip_images and dataset.meta.stats:
for key, feature_stats in dataset.meta.stats.items():
new_stats.setdefault(key, feature_stats)
logging.info("Updating dataset metadata with new quantile statistics") logging.info("Updating dataset metadata with new quantile statistics")
dataset.meta.stats = new_stats dataset.meta.stats = new_stats
@@ -276,10 +288,15 @@ def main():
"--no-sampling", "--no-sampling",
action="store_true", action="store_true",
help=( help=(
"Compute stats over every frame (exact, higher memory). By default, " "Compute stats over every frame (higher memory). By default, "
"image/video frames are sub-sampled per episode to bound memory." "image/video frames are sub-sampled per episode to bound memory."
), ),
) )
parser.add_argument(
"--skip-images",
action="store_true",
help="Skip image/video features and preserve their existing stats",
)
args = parser.parse_args() args = parser.parse_args()
root = Path(args.root) if args.root else None root = Path(args.root) if args.root else None
@@ -291,6 +308,7 @@ def main():
root=root, root=root,
overwrite=args.overwrite, overwrite=args.overwrite,
use_sampling=not args.no_sampling, use_sampling=not args.no_sampling,
skip_images=args.skip_images,
) )
+174
View File
@@ -0,0 +1,174 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Convert a DCP-format checkpoint into a distributable safetensors model, offline.
Runs single-process (no GPUs, no process group). Example:
```bash
lerobot-convert-dcp --checkpoint_dir=outputs/train/run/checkpoints/005000
lerobot-convert-dcp --checkpoint_dir=... --delete_dcp=true --push_to_hub=user/my-policy
```
`--push_to_hub` publishes the converted directory as a model repo, degrading gracefully: the
core artifacts (model.safetensors, config.json, processor files) always upload; the README
model card is enriched with training/dataset metadata only when `train_config.json` (and the
dataset it names) are reachable, with a WARNING naming exactly what was skipped otherwise.
DCP shard artifacts are never uploaded — published repos carry safetensors only.
"""
import logging
from dataclasses import dataclass
from pathlib import Path
from huggingface_hub import HfApi
from lerobot.configs import parser
from lerobot.distributed.checkpoint import dcp_to_safetensors
from lerobot.utils.constants import PRETRAINED_MODEL_DIR
from lerobot.utils.utils import init_logging
@dataclass
class ConvertDcpConfig:
"""CLI config for the offline DCP-to-safetensors checkpoint conversion."""
# A checkpoint step directory (containing pretrained_model/) or a pretrained_model
# directory itself.
checkpoint_dir: Path
# Remove the DCP shard directory after a successful conversion.
delete_dcp: bool = False
# Publish the converted directory to this Hub repo id (e.g. "user/my-policy").
push_to_hub: str | None = None
private: bool | None = None
def _locate_pretrained_dir(checkpoint_dir: Path) -> Path:
"""Resolve the pretrained_model/ directory from a user-supplied checkpoint path.
Args:
checkpoint_dir (Path): A checkpoint step directory (containing `pretrained_model/`) or a
`pretrained_model` directory itself.
Returns:
Path: The nested `pretrained_model/` directory when present, otherwise `checkpoint_dir`
unchanged.
"""
nested = checkpoint_dir / PRETRAINED_MODEL_DIR
return nested if nested.is_dir() else checkpoint_dir
def _publish_converted(pretrained_dir: Path, repo_id: str, private: bool | None) -> None:
"""Best-effort publish of a converted checkpoint dir, degrading gracefully.
The core artifacts (model.safetensors, config.json, processor files) always upload; the README
model card gains training/dataset metadata only when `train_config.json` (and the dataset it
names) are reachable, with a WARNING naming what was skipped otherwise. DCP shard artifacts are
excluded from the upload.
Args:
pretrained_dir (Path): The converted `pretrained_model/` directory to upload.
repo_id (str): Target Hub model repo id (e.g. "user/my-policy"); created if missing.
private (bool | None): Repo visibility passed to `create_repo`; None keeps the Hub (or
existing repo's) default.
"""
from lerobot.common.train_utils import generate_model_card
from lerobot.configs.policies import PreTrainedConfig
from lerobot.configs.train import TRAIN_CONFIG_NAME, TrainPipelineConfig
train_cfg = None
dataset_meta = None
if (pretrained_dir / TRAIN_CONFIG_NAME).is_file():
try:
train_cfg = TrainPipelineConfig.from_pretrained(pretrained_dir)
except Exception as e: # noqa: BLE001 — degrade, never block the upload
logging.warning(f"Could not parse {TRAIN_CONFIG_NAME} ({e}); README will lack training metadata.")
else:
logging.warning(f"{TRAIN_CONFIG_NAME} missing; README will lack training metadata.")
if train_cfg is not None:
try:
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
dataset_meta = LeRobotDatasetMetadata(
repo_id=train_cfg.dataset.repo_id,
root=train_cfg.dataset.root,
revision=train_cfg.dataset.revision,
)
except Exception as e: # noqa: BLE001
logging.warning(
f"Dataset '{train_cfg.dataset.repo_id}' unreachable ({e}); README will lack dataset metadata."
)
try:
model_cfg = PreTrainedConfig.from_pretrained(pretrained_dir)
card = generate_model_card(model_cfg, cfg=train_cfg, dataset_meta=dataset_meta)
card.save(str(pretrained_dir / "README.md"))
except Exception as e: # noqa: BLE001
logging.warning(f"Could not build the model card ({e}); publishing without README.")
api = HfApi()
repo_id = api.create_repo(repo_id=repo_id, private=private, exist_ok=True).repo_id
commit_info = api.upload_folder(
repo_id=repo_id,
repo_type="model",
folder_path=str(pretrained_dir),
commit_message="Upload converted policy (DCP -> safetensors)",
allow_patterns=["*.safetensors", "*.json", "*.yaml", "*.md"],
# The checkpoint keeps its DCP shard directory unless --delete_dcp was passed; the
# allow list above admits neither `.distcp` shards nor their `.metadata` sidecar.
ignore_patterns=["*.tmp", "*.log"],
)
logging.info(f"Model pushed to {commit_info.repo_url.url}")
@parser.wrap()
def convert_checkpoint(cfg: ConvertDcpConfig) -> Path:
"""Merge a checkpoint's DCP shards into `model.safetensors`, then optionally publish it.
Args:
cfg (ConvertDcpConfig): Conversion options — the checkpoint directory to convert, whether
to delete the DCP shards after a successful merge, and the optional Hub repo id (and
visibility) to publish the converted directory to.
Returns:
Path: The path to the merged `model.safetensors` file.
Raises:
FileNotFoundError: If the checkpoint has no DCP shard directory, i.e. it was not saved
with `checkpoint_format=dcp` (or `safetensors_dcp`).
"""
from accelerate.utils.constants import FSDP_MODEL_NAME
pretrained_dir = _locate_pretrained_dir(cfg.checkpoint_dir)
dcp_dir = pretrained_dir / f"{FSDP_MODEL_NAME}_0"
if not dcp_dir.is_dir():
raise FileNotFoundError(
f"No DCP shard directory at {dcp_dir}. Point --checkpoint_dir at a checkpoint "
"saved with checkpoint_format=dcp (or safetensors_dcp)."
)
logging.info(f"Merging {dcp_dir} -> {pretrained_dir / 'model.safetensors'}")
safetensors_path = dcp_to_safetensors(dcp_dir, pretrained_dir, delete_dcp=cfg.delete_dcp)
if cfg.push_to_hub:
_publish_converted(pretrained_dir, cfg.push_to_hub, cfg.private)
return safetensors_path
def main() -> None:
"""`lerobot-convert-dcp` console entry point: set up logging and run the conversion."""
init_logging()
convert_checkpoint()
if __name__ == "__main__":
main()
File diff suppressed because it is too large Load Diff
+205
View File
@@ -0,0 +1,205 @@
# 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.
"""Doctest plumbing so the examples in our docstrings actually run.
Adapted from `transformers.testing_utils`. Two stdlib limitations make this necessary:
1. Ruff is configured with `docstring-code-format = true`, which reformats code inside docstrings and
removes the blank line before the closing fence. stdlib's `_EXAMPLE_RE` then swallows the ` ``` ` into
the expected-output group, so every example that has output fails. [`LeRobotDocTestParser`] patches the
regex to stop at a fence.
2. `doctest.DocTestFinder` reports the wrong line number for `@property` and `functools.wraps` objects
(https://bugs.python.org/issue17446). Our hardware API is property-heavy — `observation_features`,
`action_features`, `is_connected`, `is_calibrated` are all abstract properties — so
[`LeRobotDoctestModule`] unwraps them before locating the example.
Two environment variables skip whole example blocks by content:
- `SKIP_CUDA_DOCTEST=1` skips examples that need a GPU.
- `SKIP_HARDWARE_DOCTEST=1` skips examples that need a physical robot or a Hub download.
Both are heuristics over the example source. They are deliberately blunt: an example that is skipped
needlessly costs nothing, whereas one that runs on a machine without the hardware hangs or fails.
"""
import doctest
import functools
import inspect
import os
import re
import sys
from collections.abc import Iterable
from _pytest.doctest import (
DoctestItem,
DoctestModule,
_get_checker,
_get_continue_on_failure,
_get_runner,
get_optionflags,
)
from _pytest.nodes import Collector
from _pytest.outcomes import skip
# Calls whose progress bars would otherwise be compared against the expected output. The lookahead leaves
# lines that already carry a directive alone.
_NOISY_CALL_PATTERN = re.compile(r"(>>> (?!.*# doctest:).*(?:load_dataset|LeRobotDataset)\(.*)")
_CUDA_PATTERN = re.compile(r"cuda|to\(0\)|device=0")
# Serial ports, video devices, and the connect/scan calls that talk to real hardware.
_HARDWARE_PATTERN = re.compile(r"/dev/tty|/dev/video|COM\d|\.connect\(|find_cameras\(|find_port\(")
# Anything that reaches the Hub over the network.
_HUB_PATTERN = re.compile(r"from_pretrained\(|push_to_hub\(|snapshot_download\(|load_dataset\(")
def preprocess_string(string: str, skip_cuda_tests: bool, skip_hardware_tests: bool) -> str:
"""Prepare a docstring or `.mdx` file to be run by doctest.
Args:
string (`str`):
A whole file's contents for `.mdx`, or a single docstring for a Python file. Either may hold
several fenced examples.
skip_cuda_tests (`bool`):
Whether to drop examples that look like they need a GPU.
skip_hardware_tests (`bool`):
Whether to drop examples that look like they need a robot or a Hub download.
Returns:
`str`: The input with `# doctest: +IGNORE_RESULT` injected on noisy calls, or an empty string if
the examples were skipped — in which case no doctest is collected for it at all.
"""
# Match against the example lines only, not the surrounding prose, so that a docstring merely
# *describing* CUDA or a serial port is not mistaken for one that uses them.
example_lines = "\n".join(
line for line in string.splitlines() if line.lstrip().startswith((">>>", "..."))
)
if not example_lines:
return string
if skip_cuda_tests and _CUDA_PATTERN.search(example_lines):
return ""
if skip_hardware_tests and (
_HARDWARE_PATTERN.search(example_lines) or _HUB_PATTERN.search(example_lines)
):
return ""
return _NOISY_CALL_PATTERN.sub(r"\1 # doctest: +IGNORE_RESULT", string)
class LeRobotDocTestParser(doctest.DocTestParser):
"""A `DocTestParser` that understands fenced, auto-formatted code blocks.
Ruff's `docstring-code-format` removes the blank line before a closing fence, after which stdlib's
`_EXAMPLE_RE` reads the fence itself as part of the expected output and every example with output
fails. The regex below is the stdlib one plus a clause that stops matching at a fence.
"""
# fmt: off
_EXAMPLE_RE = re.compile(r'''
# Source consists of a PS1 line followed by zero or more PS2 lines.
(?P<source>
(?:^(?P<indent> [ ]*) >>> .*) # PS1 line
(?:\n [ ]* \.\.\. .*)*) # PS2 lines
\n?
# Want consists of any non-blank lines that do not start with PS1.
(?P<want> (?:(?![ ]*$) # Not a blank line
(?![ ]*>>>) # Not a line starting with PS1
(?:(?!```).)* # Stop at a closing fence: formatting drops the blank line before it
(?:\n|$) # Match a new line or end of string
)*)
''', re.MULTILINE | re.VERBOSE
)
# fmt: on
skip_cuda_tests: bool = os.environ.get("SKIP_CUDA_DOCTEST", "0") == "1"
skip_hardware_tests: bool = os.environ.get("SKIP_HARDWARE_DOCTEST", "0") == "1"
def parse(self, string, name="<string>"):
"""Preprocess `string`, then parse it as stdlib would.
Args:
string (`str`):
The docstring or file contents to parse.
name (`str`, *optional*, defaults to `"<string>"`):
Name used in failure messages.
Returns:
`list`: The examples and interleaved text, as returned by `doctest.DocTestParser.parse`.
"""
string = preprocess_string(string, self.skip_cuda_tests, self.skip_hardware_tests)
return super().parse(string, name)
class LeRobotDoctestModule(DoctestModule):
"""A pytest `DoctestModule` that collects with [`LeRobotDocTestParser`].
`doctest.DocTestFinder` binds its default parser at class-definition time, so patching
`doctest.DocTestParser` in `conftest.py` does not reach the finder pytest builds. The parser has to be
passed in explicitly, which means reimplementing `collect`. It mirrors pytest's own implementation.
"""
def collect(self) -> Iterable[DoctestItem]:
"""Collect the doctests in this module.
Returns:
`Iterable[DoctestItem]`: One item per example-bearing docstring. Docstrings whose examples were
dropped by `preprocess_string` yield nothing.
"""
class MockAwareDocTestFinder(doctest.DocTestFinder):
"""A doctest finder that reports correct line numbers for properties and wrapped callables."""
# Fixed upstream in CPython 3.11.9 / 3.12.3; kept for older interpreters. Our hardware API is
# property-heavy (`observation_features`, `is_connected`, ...), so a wrong line number here
# would point every failure at the decorator. https://github.com/python/cpython/issues/61648
def _find_lineno(self, obj, source_lines):
if isinstance(obj, property):
obj = getattr(obj, "fget", obj)
if hasattr(obj, "__wrapped__"):
obj = inspect.unwrap(obj)
return super()._find_lineno(obj, source_lines)
if sys.version_info < (3, 13):
# `cached_property` is otherwise never considered part of the current module and its
# examples are silently skipped. https://github.com/python/cpython/issues/107995
def _from_module(self, module, object):
if isinstance(object, functools.cached_property):
object = object.func
return super()._from_module(module, object)
try:
module = self.obj
except Collector.CollectError:
if self.config.getvalue("doctest_ignore_import_errors"):
skip(f"unable to import module {self.path!r}")
else:
raise
# Doctests support fixtures via `getfixture` and autouse.
self.session._fixturemanager.parsefactories(self)
finder = MockAwareDocTestFinder(parser=LeRobotDocTestParser())
optionflags = get_optionflags(self.config)
runner = _get_runner(
verbose=False,
optionflags=optionflags,
checker=_get_checker(),
continue_on_failure=_get_continue_on_failure(self.config),
)
for test in finder.find(module, module.__name__):
if test.examples: # Skip docstrings with no examples, and blocks dropped by the parser.
yield DoctestItem.from_parent(self, name=test.name, runner=runner, dtest=test)
+31 -5
View File
@@ -25,6 +25,9 @@ from .constants import CHECKPOINTS_DIR
T = TypeVar("T", bound="HubMixin") T = TypeVar("T", bound="HubMixin")
# Sharded-training resume artifacts (torch DCP shard dirs + shard files). Published model repos
# carry safetensors only, so publishing uploads exclude these — checkpoint pushes (which exist
# for resume, not distribution) deliberately do not.
def find_latest_hub_checkpoint( def find_latest_hub_checkpoint(
repo_id: str, repo_id: str,
*, *,
@@ -36,6 +39,16 @@ def find_latest_hub_checkpoint(
Training runs push checkpoints to ``checkpoints/<step>/`` (see Training runs push checkpoints to ``checkpoints/<step>/`` (see
``push_checkpoint_to_hub``). This lists those step dirs and returns ``push_checkpoint_to_hub``). This lists those step dirs and returns
``checkpoints/<highest-step>``, or ``None`` if the repo has no checkpoints. ``checkpoints/<highest-step>``, or ``None`` if the repo has no checkpoints.
Args:
repo_id (str): The Hub model repo to inspect.
token (str | bool | None): Hub authentication token. Defaults to None (the token
cached by `huggingface-cli login`).
revision (str | None): Repo revision to list. Defaults to None (the default branch).
Returns:
str | None: The repo-relative path `checkpoints/<highest-step>`, or None if the repo
has no checkpoints.
""" """
files = HfApi().list_repo_files(repo_id=repo_id, repo_type="model", revision=revision, token=token) files = HfApi().list_repo_files(repo_id=repo_id, repo_type="model", revision=revision, token=token)
prefix = f"{CHECKPOINTS_DIR}/" prefix = f"{CHECKPOINTS_DIR}/"
@@ -164,7 +177,7 @@ class HubMixin:
ignore_patterns: list[str] | str | None = None, ignore_patterns: list[str] | str | None = None,
delete_patterns: list[str] | str | None = None, delete_patterns: list[str] | str | None = None,
card_kwargs: dict[str, Any] | None = None, card_kwargs: dict[str, Any] | None = None,
) -> str: ) -> str | None:
""" """
Upload model checkpoint to the Hub. Upload model checkpoint to the Hub.
@@ -172,6 +185,10 @@ class HubMixin:
`delete_patterns` to delete existing remote files in the same commit. See [`upload_folder`] reference for more `delete_patterns` to delete existing remote files in the same commit. See [`upload_folder`] reference for more
details. details.
Distributed contract: call on EVERY rank. `save_pretrained` runs on all ranks — for
sharded objects it can contain a collective gather (rank-gating it would deadlock) —
while repo creation and the upload happen on the main process only.
Args: Args:
repo_id (`str`): repo_id (`str`):
ID of the repository to push to (example: `"username/my-model"`). ID of the repository to push to (example: `"username/my-model"`).
@@ -197,11 +214,17 @@ class HubMixin:
Additional arguments passed to the card template to customize the card. Additional arguments passed to the card template to customize the card.
Returns: Returns:
The url of the commit of your object in the given repository. `str` or `None`: The url of the commit of your object in the given repository, or
`None` on non-main ranks of a distributed run (only the main process uploads).
""" """
api = HfApi(token=token) # Lazy import: hub code must not import the distributed package at module load
repo_id = api.create_repo(repo_id=repo_id, private=private, exist_ok=True).repo_id # (configs -> hub is on the import path of lerobot.distributed itself).
from lerobot.distributed.utils import is_main_process
# Distributed contract: `save_pretrained` runs on EVERY rank — for sharded policies it
# contains a collective gather (rank-gating it would deadlock) and it writes into this
# rank's private tmpdir only on the main process. Repo creation and upload are then
# main-process-only.
if commit_message is None: if commit_message is None:
if "Policy" in self.__class__.__name__: if "Policy" in self.__class__.__name__:
commit_message = "Upload policy" commit_message = "Upload policy"
@@ -210,10 +233,13 @@ class HubMixin:
else: else:
commit_message = f"Upload {self.__class__.__name__}" commit_message = f"Upload {self.__class__.__name__}"
# Push the files to the repo in a single commit
with TemporaryDirectory(ignore_cleanup_errors=True) as tmp: with TemporaryDirectory(ignore_cleanup_errors=True) as tmp:
saved_path = Path(tmp) / repo_id saved_path = Path(tmp) / repo_id
self.save_pretrained(saved_path, card_kwargs=card_kwargs) self.save_pretrained(saved_path, card_kwargs=card_kwargs)
if not is_main_process():
return None
api = HfApi(token=token)
repo_id = api.create_repo(repo_id=repo_id, private=private, exist_ok=True).repo_id
return api.upload_folder( return api.upload_folder(
repo_id=repo_id, repo_id=repo_id,
repo_type="model", repo_type="model",
+51 -16
View File
@@ -14,10 +14,10 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from collections import defaultdict from collections import defaultdict
from collections.abc import Callable
from typing import Any from typing import Any
import torch import torch
import torch.distributed as dist
from .utils import format_big_number from .utils import format_big_number
@@ -69,12 +69,31 @@ class MetricsTracker:
""" """
A helper class to track and log metrics over time. A helper class to track and log metrics over time.
Args:
batch_size (int): Per-process batch size (samples per micro-batch on each
data-parallel worker).
num_frames (int): Total number of frames in the training dataset.
num_episodes (int): Total number of episodes in the training dataset.
metrics (dict[str, AverageMeter]): The meters to track, keyed by metric name.
initial_step (int): Step counter to start from (non-zero when resuming a run).
Defaults to 0.
dp_world_size (int): Number of distinct data-parallel workers
(`dp_replicate * dp_shard`), used to scale sample accounting; context-parallel
peers consume the same batch and must not be double counted. Defaults to 1.
Usage pattern: Usage pattern:
```python ```python
# initialize, potentially with non-zero initial step (e.g. if resuming run) # initialize, potentially with non-zero initial step (e.g. if resuming run)
metrics = {"loss": AverageMeter("loss", ":.3f")} metrics = {"loss": AverageMeter("loss", ":.3f")}
train_metrics = MetricsTracker(cfg, dataset, metrics, initial_step=step) train_metrics = MetricsTracker(
batch_size,
dataset.num_frames,
dataset.num_episodes,
metrics,
initial_step=step,
dp_world_size=dp_world,
)
# update metrics derived from step (samples, episodes, epochs) at each training step # update metrics derived from step (samples, episodes, epochs) at each training step
train_metrics.step() train_metrics.step()
@@ -98,12 +117,12 @@ class MetricsTracker:
"_batch_size", "_batch_size",
"_num_frames", "_num_frames",
"_avg_samples_per_ep", "_avg_samples_per_ep",
"_dp_world_size",
"metrics", "metrics",
"steps", "steps",
"samples", "samples",
"episodes", "episodes",
"epochs", "epochs",
"accelerator",
"_caller_metrics", "_caller_metrics",
] ]
@@ -114,22 +133,25 @@ class MetricsTracker:
num_episodes: int, num_episodes: int,
metrics: dict[str, AverageMeter], metrics: dict[str, AverageMeter],
initial_step: int = 0, initial_step: int = 0,
accelerator: Callable | None = None, dp_world_size: int = 1,
): ):
self.__dict__.update(dict.fromkeys(self.__keys__)) self.__dict__.update(dict.fromkeys(self.__keys__))
self._batch_size = batch_size self._batch_size = batch_size
self._num_frames = num_frames self._num_frames = num_frames
self._avg_samples_per_ep = num_frames / num_episodes self._avg_samples_per_ep = num_frames / num_episodes
# Sample accounting scales by the number of DISTINCT data-parallel workers, which is
# dp_replicate * dp_shard — not the world size: context-parallel peers consume the same
# batch and must not be double counted. `step` counts micro-batches, so no
# grad-accumulation factor belongs here either.
self._dp_world_size = dp_world_size
self.metrics = metrics self.metrics = metrics
self.steps = initial_step self.steps = initial_step
world_size = accelerator.num_processes if accelerator else 1
# A sample is an (observation,action) pair, where observation and action # A sample is an (observation,action) pair, where observation and action
# can be on multiple timestamps. In a batch, we have `batch_size` number of samples. # can be on multiple timestamps. In a batch, we have `batch_size` number of samples.
self.samples = self.steps * self._batch_size * world_size self.samples = self.steps * self._batch_size * self._dp_world_size
self.episodes = self.samples / self._avg_samples_per_ep self.episodes = self.samples / self._avg_samples_per_ep
self.epochs = self.samples / self._num_frames self.epochs = self.samples / self._num_frames
self.accelerator = accelerator
# Meter names the caller registered up front. update_metrics() leaves these untouched, so a # Meter names the caller registered up front. update_metrics() leaves these untouched, so a
# policy that echoes e.g. "loss" in its output dict can't clobber the aggregated meter. # policy that echoes e.g. "loss" in its output dict can't clobber the aggregated meter.
self._caller_metrics: set[str] = set(self.metrics) self._caller_metrics: set[str] = set(self.metrics)
@@ -155,8 +177,7 @@ class MetricsTracker:
Updates metrics that depend on 'step' for one step. Updates metrics that depend on 'step' for one step.
""" """
self.steps += 1 self.steps += 1
world_size = self.accelerator.num_processes if self.accelerator else 1 self.samples += self._batch_size * self._dp_world_size
self.samples += self._batch_size * world_size
self.episodes = self.samples / self._avg_samples_per_ep self.episodes = self.samples / self._avg_samples_per_ep
self.epochs = self.samples / self._num_frames self.epochs = self.samples / self._num_frames
@@ -181,11 +202,16 @@ class MetricsTracker:
across all distributed processes (in-place). across all distributed processes (in-place).
This is a collective operation and MUST be invoked on every rank — typically just before This is a collective operation and MUST be invoked on every rank — typically just before
logging. With no accelerator or in single-process runs it is a no-op. Without it, metrics logging. Outside distributed runs it is a no-op. Without it, metrics reported by the
reported by the main process only reflect rank 0; for bottleneck-style timings main process only reflect rank 0; for bottleneck-style timings (``dataloading_s``,
(``dataloading_s``, ``update_s``, ...) that means the slowest worker's stall is invisible. ``update_s``, ...) that means the slowest worker's stall is invisible.
Torch-native on purpose: metrics code carries no Accelerator dependency.
Note the reduction spans the WORLD group — correct for count-free averages (loss values
are identical within a context-parallel group, so including CP peers is a weighted
no-op).
""" """
if self.accelerator is None or self.accelerator.num_processes <= 1: if not dist.is_initialized() or dist.get_world_size() <= 1:
return return
buckets: dict[str, list[str]] = defaultdict(list) buckets: dict[str, list[str]] = defaultdict(list)
@@ -195,11 +221,20 @@ class MetricsTracker:
if not buckets: if not buckets:
return return
device = self.accelerator.device device = (
torch.device("cuda", torch.cuda.current_device())
if torch.cuda.is_available()
else torch.device("cpu")
)
reduce_ops = {
"mean": dist.ReduceOp.AVG,
"sum": dist.ReduceOp.SUM,
"max": dist.ReduceOp.MAX,
}
for reduction, names in buckets.items(): for reduction, names in buckets.items():
tensor = torch.tensor([self.metrics[n].avg for n in names], dtype=torch.float32, device=device) tensor = torch.tensor([self.metrics[n].avg for n in names], dtype=torch.float32, device=device)
reduced = self.accelerator.reduce(tensor, reduction=reduction) dist.all_reduce(tensor, op=reduce_ops[reduction])
for name, value in zip(names, reduced.tolist(), strict=True): for name, value in zip(names, tensor.tolist(), strict=True):
meter = self.metrics[name] meter = self.metrics[name]
# Preserve avg == sum / count so a later .update() on this meter accumulates # Preserve avg == sum / count so a later .update() on this meter accumulates
# against the cluster view, not the stale per-rank history. # against the cluster view, not the stale per-rank history.
@@ -0,0 +1,221 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Legacy-checkpoint contracts.
Two contracts are pinned here so they are documented behavior, not accidents:
- **The v0.6.0 hard break.** The v0.6.0 #3810 FSDP checkpoint layout
(full gathered ``model.safetensors`` + full ``optimizer_state.safetensors``, no DCP dirs,
no ``checkpoint_format`` in ``train_config.json``) is a hard break with ZERO v0.6.0-aware
runtime code — not even layout detection. A sharded resume pointed at such a checkpoint
must fail through the ORDINARY missing-artifact path (torch DCP erroring on the absent
``training_state/optimizer_0/``), while the model weights remain loadable forever via
``from_pretrained`` and the old ``num_processes`` key keeps feeding the topology reader.
- **Converter equivalence.** ``dcp_to_safetensors`` (real ``merge_fsdp_weights``, no mocks)
on accelerate's ``save_fsdp_model`` DCP layout reproduces exactly the tensors that the
direct-gather ``save_pretrained`` artifact contains.
"""
import json
from pathlib import Path
from types import SimpleNamespace
import pytest
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
import torch
import torch.distributed.checkpoint as dist_cp
from accelerate.utils.constants import FSDP_MODEL_NAME, OPTIMIZER_NAME
from safetensors.torch import load_file
from torch.distributed.checkpoint.api import CheckpointException
from torch.distributed.fsdp import FSDPModule
from lerobot.common.train_utils import (
load_training_metadata,
resume_after_prepare,
resume_before_prepare,
)
from lerobot.configs.accelerator import FSDPConfig
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import TRAIN_CONFIG_NAME, CheckpointFormat, TrainPipelineConfig
from lerobot.distributed.checkpoint import dcp_to_safetensors, is_sharded_module
from lerobot.optim.optimizers import save_optimizer_state
from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR, TRAINING_STEP
from lerobot.utils.io_utils import write_json
from lerobot.utils.random_utils import save_rng_state
from tests.fixtures.dummy_checkpoint_policy import DummyCheckpointPolicy, make_dummy_policy
@pytest.fixture
def accelerate_state():
"""accelerate's process state, as the trainer's `Accelerator()` would have initialized it.
`load_fsdp_optimizer` and `merge_fsdp_weights` both consult `PartialState` internals
(logging and main-process gating). Single-process CPU state; reset on teardown so no
global accelerate state leaks into other tests.
"""
from accelerate.state import AcceleratorState, PartialState
PartialState()
yield
AcceleratorState._reset_state(reset_partial_state=True)
def make_v060_fsdp_checkpoint(checkpoint_dir: Path) -> dict[str, torch.Tensor]:
"""Reproduce the v0.6.0 #3810 FSDP checkpoint layout with real artifacts.
- ``pretrained_model/``: ``config.json`` + full gathered ``model.safetensors`` (real
``save_pretrained`` outputs) and a ``train_config.json`` predating the v0.7 fields
(``checkpoint_format``/``parallelism``/``accelerator`` stripped from the draccus dump);
- ``training_state/``: old-style ``training_step.json`` (``{"step", "num_processes"}``,
no ``dp_world_size``), ``rng_state.safetensors``, and the gathered full optimizer
channel (``optimizer_state.safetensors`` + ``optimizer_param_groups.json``) — and,
crucially, NO ``optimizer_0/`` DCP directory.
Returns the saved model weights for later comparison.
"""
policy = make_dummy_policy()
optimizer = torch.optim.Adam(policy.parameters())
policy.forward({"observation.state": torch.randn(2, 4)})[0].backward()
optimizer.step() # real optimizer state, applied before the weights are saved
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
policy.save_pretrained(pretrained_dir)
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"), batch_size=3)
cfg._save_pretrained(pretrained_dir)
config_path = pretrained_dir / TRAIN_CONFIG_NAME
raw = json.loads(config_path.read_text())
assert "checkpoint_format" in raw # draccus dumps defaults; a v0.6.0 config predates the key
for key in ("checkpoint_format", "parallelism", "accelerator"):
raw.pop(key, None)
config_path.write_text(json.dumps(raw, indent=4))
training_state_dir = checkpoint_dir / TRAINING_STATE_DIR
training_state_dir.mkdir()
write_json({"step": 5000, "num_processes": 4}, training_state_dir / TRAINING_STEP)
save_rng_state(training_state_dir)
save_optimizer_state(optimizer, training_state_dir)
return {key: tensor.clone() for key, tensor in policy.state_dict().items()}
def as_fsdp2_module(policy: DummyCheckpointPolicy) -> DummyCheckpointPolicy:
"""Give the policy FSDP2's runtime identity via the in-place class swap `fully_shard` performs.
torch's `fully_shard` swaps ``module.__class__`` to a ``(FSDPModule, type(module))``
subclass; mirroring that swap is what makes `is_sharded_module` (and thus the sharded
branch of `resume_after_prepare`) see a sharded model on a CPU-only single process. The
parameters stay plain tensors — sufficient here, because the resume must fail at the DCP
read before any sharded state is touched.
"""
policy.__class__ = type(f"FSDP{type(policy).__name__}", (FSDPModule, type(policy)), {})
assert is_sharded_module(policy)
return policy
def sharded_passthrough_accelerator() -> SimpleNamespace:
"""The accelerator surface the sharded resume touches, carrying the trainer's real plugin.
`FSDPConfig.build_plugin()` is the exact FSDP2 plugin construction `make_accelerator`
hands to accelerate (state_dict_type stays at the FSDP2 default, SHARDED_STATE_DICT).
"""
return SimpleNamespace(
unwrap_model=lambda m: m,
wait_for_everyone=lambda: None,
state=SimpleNamespace(fsdp_plugin=FSDPConfig().build_plugin()),
)
class TestV060HardBreak:
"""Pin the v0.6.0 hard break as a contract.
Zero v0.6.0-aware code ships — not even layout detection — so every assertion here must
hold through ORDINARY code paths only: the recorded config parses with plain defaults,
phase-1 resume and the weights stay loadable, and the sharded phase-2 resume fails with
torch DCP's own missing-artifact error, never a bespoke v0.6.0 message.
"""
def test_sharded_resume_fails_with_ordinary_missing_artifact_error(self, tmp_path, accelerate_state):
make_v060_fsdp_checkpoint(tmp_path)
# No checkpoint_format recorded -> plain draccus default, no layout detection anywhere.
cfg = TrainPipelineConfig.from_pretrained(tmp_path / PRETRAINED_MODEL_DIR / TRAIN_CONFIG_NAME)
assert cfg.checkpoint_format is CheckpointFormat.SAFETENSORS
cfg.checkpoint_path = tmp_path
# Phase 1 (RNG + step counter) is format-independent and still succeeds.
assert resume_before_prepare(cfg) == 5000
# Phase 2 under sharding: the recorded format skips the DCP model preflight (the
# weights were already loaded by from_pretrained), then the sharded optimizer load
# hits the absent optimizer_0/ and fails inside torch DCP — the ordinary error path.
assert not (tmp_path / TRAINING_STATE_DIR / f"{OPTIMIZER_NAME}_0").exists()
policy = as_fsdp2_module(make_dummy_policy())
optimizer = torch.optim.Adam(policy.parameters())
with pytest.raises(CheckpointException) as excinfo:
resume_after_prepare(cfg, sharded_passthrough_accelerator(), policy, optimizer, None)
message = str(excinfo.value)
assert "lerobot-convert-dcp" not in message # the converter hint belongs to recorded-format=DCP
assert "v0.6" not in message # no bespoke wording: the explanation lives in the migration docs
def test_weights_remain_loadable_via_from_pretrained(self, tmp_path):
saved_weights = make_v060_fsdp_checkpoint(tmp_path)
policy = DummyCheckpointPolicy.from_pretrained(tmp_path / PRETRAINED_MODEL_DIR)
for key, tensor in policy.state_dict().items():
assert torch.equal(tensor, saved_weights[key]), key
def test_topology_reader_falls_back_to_legacy_num_processes(self, tmp_path):
make_v060_fsdp_checkpoint(tmp_path)
assert load_training_metadata(tmp_path / TRAINING_STATE_DIR)["dp_world_size"] == 4
class TestConverterEquivalence:
def test_dcp_to_safetensors_output_equals_direct_gather(self, tmp_path, accelerate_state):
"""DCP -> safetensors conversion is exactly the direct-gather artifact.
The DCP checkpoint is written with torch's real `dist_cp.save` (single process, no
process group), replicating accelerate's `save_fsdp_model` SHARDED_STATE_DICT branch
byte for byte: the ``{"model": state_dict}`` nesting and the ``pytorch_model_fsdp_0``
directory name. The conversion runs the real `merge_fsdp_weights` — no mocks.
"""
policy = make_dummy_policy()
with torch.no_grad():
for param in policy.parameters():
param.add_(torch.randn_like(param)) # make every tensor distinct from init
reference = {key: tensor.clone() for key, tensor in policy.state_dict().items()}
# The direct-gather artifact (on a single process the gather is state_dict itself).
direct_dir = tmp_path / "direct"
policy.save_pretrained(direct_dir)
# The DCP artifact, laid out exactly as accelerate's save_fsdp_model writes it.
pretrained_dir = tmp_path / "checkpoint" / PRETRAINED_MODEL_DIR
dcp_dir = pretrained_dir / f"{FSDP_MODEL_NAME}_0"
dcp_dir.mkdir(parents=True)
dist_cp.save(
state_dict={"model": policy.state_dict()},
storage_writer=dist_cp.FileSystemWriter(str(dcp_dir)),
)
merged_file = dcp_to_safetensors(dcp_dir, pretrained_dir)
assert merged_file == pretrained_dir / "model.safetensors"
merged = load_file(merged_file)
direct = load_file(direct_dir / "model.safetensors")
assert set(merged) == set(direct) == set(reference)
for key, tensor in reference.items():
assert torch.equal(merged[key], tensor), key
assert torch.equal(direct[key], tensor), key
assert merged[key].dtype == tensor.dtype, key
+186
View File
@@ -0,0 +1,186 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Checkpoint save/resume round-trips on the non-sharded paths."""
from types import SimpleNamespace
import pytest
import torch
from safetensors.torch import load_file
from lerobot.common.train_utils import (
load_training_metadata,
resume_after_prepare,
resume_before_prepare,
save_checkpoint,
)
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import CheckpointFormat, TrainPipelineConfig
from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR, TRAINING_STEP
from lerobot.utils.io_utils import load_json, write_json
from tests.fixtures.dummy_checkpoint_policy import make_dummy_policy
def make_cfg(**overrides) -> TrainPipelineConfig:
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"), batch_size=3)
cfg.parallelism.resolve(1)
for name, value in overrides.items():
setattr(cfg, name, value)
return cfg
def passthrough_accelerator() -> SimpleNamespace:
"""The accelerator surface save/resume touches on non-sharded runs."""
return SimpleNamespace(unwrap_model=lambda m: m, wait_for_everyone=lambda: None)
class TestSaveCheckpoint:
def test_non_sharded_layout(self, tmp_path):
policy = make_dummy_policy()
optimizer = torch.optim.Adam(policy.parameters())
save_checkpoint(
tmp_path,
step=7,
cfg=make_cfg(),
policy=policy,
optimizer=optimizer,
accelerator=passthrough_accelerator(),
)
pretrained = tmp_path / PRETRAINED_MODEL_DIR
state = tmp_path / TRAINING_STATE_DIR
assert (pretrained / "model.safetensors").is_file()
assert (pretrained / "config.json").is_file()
assert (pretrained / "train_config.json").is_file()
assert (state / TRAINING_STEP).is_file()
assert (state / "rng_state.safetensors").is_file()
assert (state / "optimizer_state.safetensors").is_file()
# single-file artifact, no index, weights intact
weights = load_file(pretrained / "model.safetensors")
assert torch.allclose(weights["net.weight"], torch.full_like(weights["net.weight"], 0.5))
assert not list(pretrained.glob("*.index.json"))
def test_training_step_records_topology(self, tmp_path):
cfg = make_cfg()
cfg.accelerator.gradient_accumulation.steps = 4
policy = make_dummy_policy()
save_checkpoint(
tmp_path,
step=11,
cfg=cfg,
policy=policy,
optimizer=torch.optim.Adam(policy.parameters()),
accelerator=passthrough_accelerator(),
)
metadata = load_training_metadata(tmp_path / TRAINING_STATE_DIR)
assert metadata["dp_world_size"] == 1
assert metadata["batch_size"] == 3
assert metadata["grad_accum_steps"] == 4
def test_dp_world_size_legacy_fallback(self, tmp_path):
"""Pre-v0.7 checkpoints recorded num_processes; the reader falls back to it."""
state_dir = tmp_path / TRAINING_STATE_DIR
state_dir.mkdir(parents=True)
write_json({"step": 5, "num_processes": 4}, state_dir / TRAINING_STEP)
metadata = load_training_metadata(tmp_path / TRAINING_STATE_DIR)
assert metadata["dp_world_size"] == 4
assert metadata["batch_size"] is None
class TestResume:
def _checkpointed_run(self, tmp_path):
policy = make_dummy_policy()
optimizer = torch.optim.Adam(policy.parameters(), lr=0.123)
# give the optimizer real state
policy.forward({"observation.state": torch.randn(2, 4)})[0].backward()
optimizer.step()
cfg = make_cfg()
save_checkpoint(
tmp_path,
step=42,
cfg=cfg,
policy=policy,
optimizer=optimizer,
accelerator=passthrough_accelerator(),
)
cfg.checkpoint_path = tmp_path
return cfg, policy, optimizer
def test_two_phase_resume_round_trip(self, tmp_path):
cfg, _, optimizer = self._checkpointed_run(tmp_path)
assert resume_before_prepare(cfg) == 42
fresh_policy = make_dummy_policy()
fresh_optimizer = torch.optim.Adam(fresh_policy.parameters(), lr=0.999)
resume_after_prepare(cfg, passthrough_accelerator(), fresh_policy, fresh_optimizer, None)
restored = fresh_optimizer.state_dict()
original = optimizer.state_dict()
assert restored["param_groups"][0]["lr"] == original["param_groups"][0]["lr"]
for key, tensor in original["state"][0].items():
assert torch.equal(restored["state"][0][key], tensor), key
def test_resume_warns_on_changed_cadence_and_topology(self, tmp_path, caplog):
"""The recorded grad-accum factor and parallelism snapshot must be compared on
resume, with one warning naming the diff."""
import logging
cfg, _, _ = self._checkpointed_run(tmp_path)
cfg.accelerator.gradient_accumulation.steps = 4
cfg.parallelism.dp_replicate = 2 # same dp_world_size story is irrelevant here
with caplog.at_level(logging.WARNING):
assert resume_before_prepare(cfg) == 42
warning = next(m for m in caplog.messages if "differ from the checkpoint" in m)
assert "grad_accum_steps: 1 -> 4" in warning
assert "dp_replicate: 1 -> 2" in warning
def test_resume_unchanged_settings_stay_silent(self, tmp_path, caplog):
import logging
cfg, _, _ = self._checkpointed_run(tmp_path)
with caplog.at_level(logging.WARNING):
resume_before_prepare(cfg)
assert not [m for m in caplog.messages if "differ from the checkpoint" in m]
def test_resume_rejects_non_sharded_checkpoint_on_sharded_run(self, tmp_path):
"""Resharding works across sizes, not across kinds: non-sharded -> sharded is rejected."""
cfg, _, _ = self._checkpointed_run(tmp_path)
cfg.parallelism.dp_shard = 2
with pytest.raises(ValueError, match="Cannot resume"):
resume_before_prepare(cfg)
def test_resume_rejects_sharded_checkpoint_on_non_sharded_run(self, tmp_path):
"""The symmetric direction: a checkpoint recorded sharded cannot resume non-sharded."""
cfg, _, _ = self._checkpointed_run(tmp_path)
state_file = tmp_path / TRAINING_STATE_DIR / TRAINING_STEP
state = load_json(state_file)
state["parallelism"]["dp_shard"] = 2
write_json(state, state_file)
with pytest.raises(ValueError, match="Cannot resume"):
resume_before_prepare(cfg)
def test_resume_before_prepare_requires_training_state(self, tmp_path):
cfg = make_cfg()
cfg.checkpoint_path = tmp_path
with pytest.raises(NotADirectoryError):
resume_before_prepare(cfg)
def test_dcp_format_integrity_preflight(self, tmp_path):
"""A checkpoint declaring DCP shards without the shard dir fails with the converter hint."""
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
cfg, policy, optimizer = self._checkpointed_run(tmp_path)
cfg.parallelism.dp_shard = 2 # pretend the recorded run was sharded
cfg.checkpoint_format = CheckpointFormat.DCP
with pytest.raises(FileNotFoundError, match="lerobot-convert-dcp"):
resume_after_prepare(cfg, passthrough_accelerator(), policy, optimizer, None)
+148
View File
@@ -0,0 +1,148 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""publish_trained_model: commit set, card, log-line contract, PEFT branch (hub fully mocked)."""
import logging
from pathlib import Path
from types import SimpleNamespace
import pytest
import lerobot.common.train_utils as train_utils
import lerobot.utils.hub as hub
from lerobot.common.train_utils import generate_model_card, publish_trained_model
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import TrainPipelineConfig
from tests.fixtures.dummy_checkpoint_policy import make_dummy_policy
class FakeHfApi:
"""Records every repo/upload interaction; shared across both HfApi import sites."""
calls: list[dict] = []
def __init__(self, *args, **kwargs):
pass
def create_repo(self, repo_id, private=None, exist_ok=False, **kwargs):
return SimpleNamespace(repo_id=repo_id)
def upload_folder(self, *, repo_id, folder_path, commit_message, **kwargs):
FakeHfApi.calls.append(
{
"repo_id": repo_id,
"commit_message": commit_message,
"files": sorted(p.name for p in Path(folder_path).iterdir()),
"ignore_patterns": kwargs.get("ignore_patterns"),
}
)
return SimpleNamespace(repo_url=SimpleNamespace(url=f"https://huggingface.co/{repo_id}"))
@pytest.fixture
def mocked_hub(monkeypatch):
FakeHfApi.calls = []
monkeypatch.setattr(train_utils, "HfApi", FakeHfApi)
monkeypatch.setattr(hub, "HfApi", FakeHfApi)
# card.validate() hits the Hub; publishing must work offline in tests
monkeypatch.setattr(train_utils.ModelCard, "validate", lambda self: None)
return FakeHfApi
def make_cfg() -> TrainPipelineConfig:
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="user/dataset"))
cfg.parallelism.resolve(1)
return cfg
class RecordingProcessor:
def __init__(self):
self.pushed_to = None
def push_to_hub(self, repo_id, **kwargs):
self.pushed_to = repo_id
class TestPublishTrainedModel:
def test_commit_set_and_log_contract(self, mocked_hub, caplog):
policy = make_dummy_policy(repo_id="user/policy")
pre, post = RecordingProcessor(), RecordingProcessor()
with caplog.at_level(logging.INFO):
publish_trained_model(make_cfg(), policy, pre, post, dataset_meta=None)
# commit 1: the model through HubMixin (config.json + model.safetensors in a tmpdir)
model_commit = mocked_hub.calls[0]
assert {"config.json", "model.safetensors"} <= set(model_commit["files"])
# commits 2-3: processors
assert pre.pushed_to == "user/policy" and post.pushed_to == "user/policy"
# commit 4: the bundle sidecar
bundle = mocked_hub.calls[-1]
assert {"README.md", "train_config.json"} <= set(bundle["files"])
# the exact line lerobot.jobs.hf watches to end remote runs early
assert any(
m.startswith("Model pushed to https://huggingface.co/user/policy") for m in caplog.messages
)
def test_peft_branch_skips_model_commit(self, mocked_hub):
policy = make_dummy_policy(repo_id="user/policy")
class FakePeftModel:
def save_pretrained(self, path):
(Path(path) / "adapter_model.safetensors").write_bytes(b"x")
publish_trained_model(make_cfg(), policy, None, None, dataset_meta=None, peft_model=FakePeftModel())
assert len(mocked_hub.calls) == 1 # only the bundle commit
bundle = mocked_hub.calls[0]
# adapter weights + the wrapped policy's config + card + train config, no full weights
assert {"README.md", "adapter_model.safetensors", "config.json", "train_config.json"} <= set(
bundle["files"]
)
assert "model.safetensors" not in bundle["files"]
def test_missing_repo_id_fails_loudly(self, mocked_hub):
policy = make_dummy_policy(repo_id=None)
with pytest.raises(ValueError, match="repo id"):
publish_trained_model(make_cfg(), policy, None, None, dataset_meta=None)
class TestGenerateModelCard:
def test_free_function_renders_from_arguments(self, monkeypatch):
monkeypatch.setattr(train_utils.ModelCard, "validate", lambda self: None)
policy = make_dummy_policy(repo_id="user/policy")
card = generate_model_card(policy.config, cfg=make_cfg(), dataset_meta=None)
assert card.data.library_name == "lerobot"
assert card.data.datasets == "user/dataset"
assert "lerobot" in card.data.tags
class TestDeprecatedPushModelToHub:
"""`push_model_to_hub` stays callable for external scripts, delegating to the publisher."""
def test_policy_shim_warns_and_publishes(self, mocked_hub):
policy = make_dummy_policy(repo_id="user/policy")
with pytest.warns(FutureWarning, match="push_model_to_hub is deprecated"):
policy.push_model_to_hub(make_cfg())
# Same artifacts the method produced before: weights + config, then card + train config.
model_commit = mocked_hub.calls[0]
assert {"config.json", "model.safetensors"} <= set(model_commit["files"])
bundle = mocked_hub.calls[-1]
assert {"README.md", "train_config.json"} <= set(bundle["files"])
def test_policy_shim_warns_that_state_dict_is_ignored(self, mocked_hub):
policy = make_dummy_policy(repo_id="user/policy")
with pytest.warns(FutureWarning, match="`state_dict` argument is ignored"):
policy.push_model_to_hub(make_cfg(), state_dict=policy.state_dict())
+135
View File
@@ -0,0 +1,135 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import draccus
import pytest
from lerobot.configs.accelerator import (
AcceleratorConfig,
ActivationCheckpointingConfig,
ActivationCheckpointingMode,
CompileConfig,
DDPConfig,
FSDPConfig,
GradientAccumulationConfig,
)
from lerobot.configs.parallelism import ParallelismConfig
class TestFieldValidation:
def test_wrap_policies_mutually_exclusive(self):
with pytest.raises(ValueError, match="mutually exclusive"):
FSDPConfig(wrap_modules=["Block"], min_num_params=1000)
def test_min_num_params_positive(self):
with pytest.raises(ValueError, match="min_num_params"):
FSDPConfig(min_num_params=0)
def test_mixed_precision_choices(self):
with pytest.raises(ValueError, match="mixed_precision"):
AcceleratorConfig(mixed_precision="tf32")
def test_gradient_accumulation_positive(self):
with pytest.raises(ValueError, match="gradient_accumulation.steps"):
GradientAccumulationConfig(steps=0)
class TestDraccusRoundTrip:
@pytest.mark.parametrize(
"cfg",
[
AcceleratorConfig(),
AcceleratorConfig(
mixed_precision="bf16",
gradient_accumulation=GradientAccumulationConfig(steps=4),
fsdp=FSDPConfig(
reshard_after_forward=False,
wrap_modules=["ACTEncoderLayer", "ACTDecoderLayer"],
cpu_offload=True,
ignored_modules=r".*pos_embed.*",
),
ddp=DDPConfig(find_unused_parameters=False, static_graph=True),
compile=CompileConfig(enabled=True, mode="max-autotune", regional=False),
activation_checkpointing=ActivationCheckpointingConfig(mode=ActivationCheckpointingMode.FULL),
),
AcceleratorConfig(fsdp=FSDPConfig(min_num_params=1_000_000)),
],
)
def test_encode_json_decode_identity(self, cfg):
payload = json.loads(json.dumps(draccus.encode(cfg)))
assert draccus.decode(AcceleratorConfig, payload) == cfg
def test_pre_existing_config_without_fields_gets_defaults(self):
assert draccus.decode(AcceleratorConfig, {}) == AcceleratorConfig()
class TestRuntimeBuilders:
"""The mirrors must translate into real accelerate objects (plugins built lazily)."""
@pytest.fixture(autouse=True)
def _requires_accelerate(self):
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
def test_fsdp_plugin_translation(self):
plugin = FSDPConfig(
reshard_after_forward=False, wrap_modules=["MyBlock"], cpu_offload=True
).build_plugin()
assert plugin.fsdp_version == 2
assert plugin.reshard_after_forward is False
assert plugin.transformer_cls_names_to_wrap == ["MyBlock"]
# bools are normalized into torch offload policies by the plugin itself
assert type(plugin.cpu_offload).__name__ == "CPUOffloadPolicy"
# LeRobot never switches state_dict_type: FSDP2's SHARDED default must hold
assert plugin.state_dict_type.name == "SHARDED_STATE_DICT"
assert not plugin.activation_checkpointing
def test_fsdp_plugin_size_based_policy(self):
plugin = FSDPConfig(min_num_params=1024).build_plugin()
assert plugin.min_num_params == 1024
assert plugin.transformer_cls_names_to_wrap is None
def test_ddp_kwargs_translation(self):
handler = DDPConfig(find_unused_parameters=False, gradient_as_bucket_view=True).build_kwargs_handler()
assert handler.find_unused_parameters is False
assert handler.gradient_as_bucket_view is True
def test_gradient_accumulation_plugin_translation(self):
plugin = GradientAccumulationConfig(steps=4).build_plugin()
assert plugin.num_steps == 4
assert plugin.sync_with_dataloader is False
def test_gradient_accumulation_never_syncs_with_dataloader(self, monkeypatch):
"""The loop cycles a finite dataloader, so accelerate's default
sync_with_dataloader=True would force an optimizer step at every dataset epoch
boundary instead of every num_steps micro-batches."""
captured = {}
class FakeAccelerator:
def __init__(self, **kwargs):
captured.update(kwargs)
monkeypatch.setattr("accelerate.Accelerator", FakeAccelerator)
parallelism = ParallelismConfig()
parallelism.resolve(1)
AcceleratorConfig(gradient_accumulation=GradientAccumulationConfig(steps=4)).build(
parallelism, cpu=True
)
ga_plugin = captured["gradient_accumulation_plugin"]
assert ga_plugin.num_steps == 4
assert ga_plugin.sync_with_dataloader is False
assert "gradient_accumulation_steps" not in captured
+118
View File
@@ -0,0 +1,118 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import draccus
import pytest
from lerobot.configs.parallelism import ContextParallelConfig, ParallelismConfig
class TestResolve:
def test_single_process_defaults(self):
cfg = ParallelismConfig()
cfg.resolve(1)
assert (cfg.dp_replicate, cfg.dp_shard) == (1, 1)
assert not cfg.is_sharded and not cfg.is_replicated_only
assert cfg.dp_world_size == 1
def test_untouched_config_fills_ddp(self):
"""Plain `torchrun --nproc-per-node=8` with a default config resolves to DDP."""
cfg = ParallelismConfig()
cfg.resolve(8)
assert cfg.dp_replicate == 8
assert cfg.is_replicated_only and not cfg.is_sharded
assert cfg.dp_world_size == 8
def test_full_shard_sentinel(self):
cfg = ParallelismConfig(dp_shard=-1)
assert cfg.is_sharded # sharded even before resolve: -1 is an explicit opt-in
cfg.resolve(8)
assert cfg.dp_shard == 8 and cfg.dp_replicate == 1
def test_hsdp_sentinel_infers_shard(self):
cfg = ParallelismConfig(dp_replicate=2, dp_shard=-1)
cfg.resolve(8)
assert (cfg.dp_replicate, cfg.dp_shard) == (2, 4)
assert cfg.dp_world_size == 8
def test_explicit_hsdp(self):
cfg = ParallelismConfig(dp_replicate=2, dp_shard=4)
cfg.resolve(8)
assert cfg.is_sharded and not cfg.is_replicated_only
def test_product_mismatch_lists_all_degrees(self):
cfg = ParallelismConfig(dp_replicate=2, dp_shard=2)
with pytest.raises(ValueError, match=r"dp_replicate=2 \* dp_shard=2.*WORLD_SIZE=8"):
cfg.resolve(8)
def test_explicit_replicate_must_match_world(self):
cfg = ParallelismConfig(dp_replicate=4)
with pytest.raises(ValueError, match="WORLD_SIZE=8"):
cfg.resolve(8)
def test_sentinel_indivisible_world(self):
cfg = ParallelismConfig(dp_replicate=3, dp_shard=-1)
with pytest.raises(ValueError, match="not divisible"):
cfg.resolve(8)
def test_cp_fails_fast(self):
cfg = ParallelismConfig(dp_shard=-1, context_parallel=ContextParallelConfig(ulysses_degree=2))
with pytest.raises(ValueError, match="not implemented"):
cfg.resolve(8)
class TestFieldValidation:
@pytest.mark.parametrize("kwargs", [{"dp_replicate": 0}, {"dp_shard": 0}, {"dp_shard": -2}])
def test_bad_dp_degrees(self, kwargs):
with pytest.raises(ValueError):
ParallelismConfig(**kwargs)
def test_cfg_parallel_capped_at_two(self):
ParallelismConfig(cfg_parallel=2) # reserved but representable
with pytest.raises(ValueError, match="cfg_parallel"):
ParallelismConfig(cfg_parallel=3)
@pytest.mark.parametrize("kwargs", [{"ring_degree": 0}, {"ulysses_degree": -1}])
def test_bad_cp_degrees(self, kwargs):
with pytest.raises(ValueError):
ContextParallelConfig(**kwargs)
def test_dp_world_size_undefined_before_resolve(self):
with pytest.raises(RuntimeError, match="resolve"):
_ = ParallelismConfig(dp_shard=-1).dp_world_size
class TestDraccusRoundTrip:
@pytest.mark.parametrize(
"cfg",
[
ParallelismConfig(),
ParallelismConfig(dp_replicate=2, dp_shard=4, cfg_parallel=2),
ParallelismConfig(
dp_shard=-1,
context_parallel=ContextParallelConfig(ring_degree=2, ulysses_degree=4),
),
],
)
def test_encode_json_decode_identity(self, cfg):
payload = json.loads(json.dumps(draccus.encode(cfg)))
assert draccus.decode(ParallelismConfig, payload) == cfg
def test_pre_existing_config_without_fields_gets_defaults(self):
"""Checkpoints written before this feature parse with default topology."""
assert draccus.decode(ParallelismConfig, {}) == ParallelismConfig()
@@ -0,0 +1,118 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""TrainPipelineConfig integration for the distributed fields: fail-fasts + config compat."""
import draccus
import pytest
from lerobot.configs.accelerator import ActivationCheckpointingMode
from lerobot.configs.default import DatasetConfig, PeftConfig
from lerobot.configs.parallelism import ContextParallelConfig, ParallelismConfig
from lerobot.configs.train import CheckpointFormat, TrainPipelineConfig
from lerobot.optim.optimizers import AdamConfig, MultiAdamConfig
def make_cfg(**overrides) -> TrainPipelineConfig:
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"))
for name, value in overrides.items():
setattr(cfg, name, value)
return cfg
def sharded() -> ParallelismConfig:
return ParallelismConfig(dp_shard=-1)
class TestDistributedFailFasts:
def test_defaults_pass(self):
make_cfg()._validate_distributed()
def test_cp_reserved(self):
cfg = make_cfg(parallelism=ParallelismConfig(context_parallel=ContextParallelConfig(ring_degree=2)))
with pytest.raises(ValueError, match="not implemented"):
cfg._validate_distributed()
def test_cfg_parallel_training_rejected(self):
cfg = make_cfg(parallelism=ParallelismConfig(cfg_parallel=2))
with pytest.raises(ValueError, match="inference-only"):
cfg._validate_distributed()
def test_compile_placeholder(self):
cfg = make_cfg()
cfg.accelerator.compile.enabled = True
with pytest.raises(ValueError, match="compile"):
cfg._validate_distributed()
def test_activation_checkpointing_placeholder(self):
cfg = make_cfg()
cfg.accelerator.activation_checkpointing.mode = ActivationCheckpointingMode.FULL
with pytest.raises(ValueError, match="activation_checkpointing"):
cfg._validate_distributed()
def test_dcp_format_requires_sharding(self):
cfg = make_cfg(checkpoint_format=CheckpointFormat.DCP)
with pytest.raises(ValueError, match="sharded"):
cfg._validate_distributed()
cfg.parallelism = sharded()
cfg._validate_distributed()
def test_fp16_rejected_when_sharded(self):
cfg = make_cfg(parallelism=sharded())
cfg.accelerator.mixed_precision = "fp16"
with pytest.raises(ValueError, match="fp16"):
cfg._validate_distributed()
cfg.accelerator.mixed_precision = "bf16"
cfg._validate_distributed()
def test_peft_rejected_when_sharded(self):
cfg = make_cfg(parallelism=sharded(), peft=PeftConfig())
with pytest.raises(ValueError, match="PEFT"):
cfg._validate_distributed()
def test_env_eval_rejected_when_sharded(self):
cfg = make_cfg(parallelism=sharded(), env_eval_freq=1000)
cfg.env = object() # any configured env triggers the check
with pytest.raises(ValueError, match="environment evaluation"):
cfg._validate_distributed()
def test_multi_optimizer_rejected_when_sharded(self):
cfg = make_cfg(parallelism=sharded(), optimizer=MultiAdamConfig())
with pytest.raises(ValueError, match="Multi-optimizer"):
cfg._validate_distributed()
cfg.optimizer = AdamConfig()
cfg._validate_distributed()
class TestConfigCompat:
def test_checkpoint_format_round_trip(self):
for fmt in CheckpointFormat:
assert draccus.decode(CheckpointFormat, draccus.encode(fmt)) is fmt
def test_wants_predicates(self):
assert CheckpointFormat.SAFETENSORS.wants_safetensors
assert not CheckpointFormat.SAFETENSORS.wants_dcp
assert CheckpointFormat.DCP.wants_dcp and not CheckpointFormat.DCP.wants_safetensors
both = CheckpointFormat.SAFETENSORS_AND_DCP
assert both.wants_safetensors and both.wants_dcp
def test_reward_model_rejected_when_sharded():
"""Sharded reward runs previously failed late (missing wrap
units, DTensor serialization at the first checkpoint) instead of at validation."""
cfg = make_cfg(parallelism=sharded())
cfg.reward_model = object() # any configured reward model triggers the check
with pytest.raises(ValueError, match="Reward-model"):
cfg._validate_distributed()
+115 -1
View File
@@ -12,8 +12,11 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from types import SimpleNamespace
import numpy as np import numpy as np
import pytest import pytest
import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])") pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
@@ -24,7 +27,9 @@ from lerobot.scripts.augment_dataset_quantile_stats import (
def _numeric_keys(dataset): def _numeric_keys(dataset):
return [k for k, v in dataset.features.items() if v["dtype"] not in ("image", "video", "string")] return [
k for k, v in dataset.features.items() if v["dtype"] not in ("image", "video", "string", "language")
]
def _image_keys(dataset): def _image_keys(dataset):
@@ -102,3 +107,112 @@ def test_quantile_stats_present_after_compute(tmp_path, lerobot_dataset_factory)
) )
stats = compute_quantile_stats_for_dataset(dataset, use_sampling=True) stats = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
assert has_quantile_stats(stats) assert has_quantile_stats(stats)
class FakeHFDataset:
"""Minimal stand-in exposing the column slicing used by the augment script."""
def __init__(self, columns: dict[str, list]):
self._columns = columns
def select_columns(self, keys):
return FakeHFDataset({key: self._columns[key] for key in keys})
def __getitem__(self, index):
return {key: values[index] for key, values in self._columns.items()}
def test_compute_quantile_stats_skips_language_features():
class FakeDataset:
num_episodes = 1
features = {
"action": {"dtype": "float32"},
"observation.language": {"dtype": "language"},
}
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
hf_dataset = FakeHFDataset(
{
"action": [[0.0], [1.0]],
"observation.language": [
[{"role": "user", "content": "pick"}],
[{"role": "assistant", "content": "done"}],
],
}
)
stats = compute_quantile_stats_for_dataset(FakeDataset())
assert set(stats) == {"action"}
def test_compute_quantile_stats_skip_images_avoids_decoding():
class FakeDataset:
num_episodes = 1
features = {
"action": {"dtype": "float32"},
"observation.images.cam": {"dtype": "video"},
}
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
hf_dataset = FakeHFDataset({"action": [[0.0], [1.0]]})
def __getitem__(self, index):
raise AssertionError(f"video frame {index} was decoded despite skip_images=True")
stats = compute_quantile_stats_for_dataset(FakeDataset(), skip_images=True)
assert set(stats) == {"action"}
def test_compute_quantile_stats_handles_single_frame():
class FakeDataset:
num_episodes = 1
features = {"action": {"dtype": "float32"}}
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 1}])
hf_dataset = FakeHFDataset({"action": [[5.0, 7.0]]})
stats = compute_quantile_stats_for_dataset(FakeDataset())
np.testing.assert_array_equal(stats["action"]["count"], np.array([1]))
for key in ("min", "max", "mean", "q01", "q10", "q50", "q90", "q99"):
np.testing.assert_allclose(stats["action"][key], np.array([5.0, 7.0]))
def test_compute_quantile_stats_image_count_uses_frames():
frames = [torch.zeros(3, 2, 2), torch.ones(3, 2, 2)]
class FakeDataset:
num_episodes = 1
features = {"observation.images.cam": {"dtype": "video"}}
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
hf_dataset = FakeHFDataset({})
def __getitem__(self, index):
return {"observation.images.cam": frames[index]}
stats = compute_quantile_stats_for_dataset(FakeDataset(), use_sampling=False)
image_stats = stats["observation.images.cam"]
np.testing.assert_array_equal(image_stats["count"], np.array([2]))
assert image_stats["mean"].shape == (3, 1, 1)
np.testing.assert_allclose(image_stats["mean"], np.full((3, 1, 1), 0.5))
def test_compute_quantile_stats_accumulates_across_episodes():
values = [[float(value)] for value in range(100)] + [[float(value)] for value in range(1000, 1010)]
class FakeDataset:
num_episodes = 2
features = {"action": {"dtype": "float32"}}
meta = SimpleNamespace(
episodes=[
{"dataset_from_index": 0, "dataset_to_index": 100},
{"dataset_from_index": 100, "dataset_to_index": 110},
]
)
hf_dataset = FakeHFDataset({"action": values})
stats = compute_quantile_stats_for_dataset(FakeDataset())
np.testing.assert_array_equal(stats["action"]["count"], np.array([110]))
expected_q90 = np.percentile(np.asarray(values), 90, axis=0)
np.testing.assert_allclose(stats["action"]["q90"], expected_q90, atol=0.1)
+70 -11
View File
@@ -688,7 +688,7 @@ def test_compute_episode_stats_string_features_skipped():
def test_aggregate_feature_stats_with_quantiles(): def test_aggregate_feature_stats_with_quantiles():
"""Test aggregating feature stats that include quantiles.""" """Test aggregating feature stats that include quantiles uses conservative bounds."""
stats_ft_list = [ stats_ft_list = [
{ {
"min": np.array([1.0]), "min": np.array([1.0]),
@@ -697,6 +697,9 @@ def test_aggregate_feature_stats_with_quantiles():
"std": np.array([2.0]), "std": np.array([2.0]),
"count": np.array([100]), "count": np.array([100]),
"q01": np.array([1.5]), "q01": np.array([1.5]),
"q10": np.array([2.0]),
"q50": np.array([5.0]),
"q90": np.array([9.0]),
"q99": np.array([9.5]), "q99": np.array([9.5]),
}, },
{ {
@@ -706,22 +709,21 @@ def test_aggregate_feature_stats_with_quantiles():
"std": np.array([2.5]), "std": np.array([2.5]),
"count": np.array([150]), "count": np.array([150]),
"q01": np.array([2.5]), "q01": np.array([2.5]),
"q10": np.array([3.0]),
"q50": np.array([6.0]),
"q90": np.array([11.0]),
"q99": np.array([11.5]), "q99": np.array([11.5]),
}, },
] ]
result = aggregate_feature_stats(stats_ft_list) result = aggregate_feature_stats(stats_ft_list)
# Should preserve quantiles # Lower quantiles use min; upper quantiles use max, regardless of counts.
assert "q01" in result np.testing.assert_allclose(result["q01"], np.array([1.5]), atol=1e-6)
assert "q99" in result np.testing.assert_allclose(result["q10"], np.array([2.0]), atol=1e-6)
np.testing.assert_allclose(result["q50"], np.array([5.0]), atol=1e-6)
# Verify quantile aggregation (weighted average) np.testing.assert_allclose(result["q90"], np.array([11.0]), atol=1e-6)
expected_q01 = (1.5 * 100 + 2.5 * 150) / 250 # ≈ 2.1 np.testing.assert_allclose(result["q99"], np.array([11.5]), atol=1e-6)
expected_q99 = (9.5 * 100 + 11.5 * 150) / 250 # ≈ 10.7
np.testing.assert_allclose(result["q01"], np.array([expected_q01]), atol=1e-6)
np.testing.assert_allclose(result["q99"], np.array([expected_q99]), atol=1e-6)
def test_aggregate_stats_mixed_quantiles(): def test_aggregate_stats_mixed_quantiles():
@@ -878,3 +880,60 @@ def test_fixed_quantiles_always_computed():
for q_key in expected_quantiles: for q_key in expected_quantiles:
assert q_key in episode_stats[key] assert q_key in episode_stats[key]
assert episode_stats[key][q_key].shape == (features[key]["shape"][0],) assert episode_stats[key][q_key].shape == (features[key]["shape"][0],)
def test_aggregate_stats_incremental_resume():
"""Verify conservative bounds remain associative across incremental additions."""
# Start with episode 1 stats (narrow distribution)
ep1_stats = {
"action": {
"min": np.array([-10.0, -5.0]),
"max": np.array([10.0, 5.0]),
"mean": np.array([0.0, 0.0]),
"std": np.array([3.0, 1.5]),
"count": np.array([500]),
"q01": np.array([-9.0, -4.5]),
"q99": np.array([9.0, 4.5]),
},
}
# Episode 2: wider distribution on dim 0
ep2_stats = {
"action": {
"min": np.array([-30.0, -5.0]),
"max": np.array([40.0, 6.0]),
"mean": np.array([5.0, 0.5]),
"std": np.array([15.0, 2.0]),
"count": np.array([100]),
"q01": np.array([-25.0, -4.0]),
"q99": np.array([35.0, 5.5]),
},
}
# First aggregation: ep1 + ep2 (simulates save_episode for ep2)
cumulative = aggregate_stats([ep1_stats, ep2_stats])
# q01 should take min (conservative lower bound)
np.testing.assert_allclose(cumulative["action"]["q01"], np.array([-25.0, -4.5]))
# q99 should take max (conservative upper bound)
np.testing.assert_allclose(cumulative["action"]["q99"], np.array([35.0, 5.5]))
# Episode 3: even wider on dim 1
ep3_stats = {
"action": {
"min": np.array([-8.0, -20.0]),
"max": np.array([8.0, 25.0]),
"mean": np.array([0.0, 3.0]),
"std": np.array([2.0, 8.0]),
"count": np.array([50]),
"q01": np.array([-7.0, -18.0]),
"q99": np.array([7.0, 22.0]),
},
}
# Second aggregation: cumulative + ep3 (simulates save_episode for ep3)
cumulative2 = aggregate_stats([cumulative, ep3_stats])
# Bounds should widen monotonically
np.testing.assert_allclose(cumulative2["action"]["q01"], np.array([-25.0, -18.0]))
np.testing.assert_allclose(cumulative2["action"]["q99"], np.array([35.0, 22.0]))
@@ -0,0 +1,139 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Version canaries for the accelerate/torch seams LeRobot's distributed engine relies on.
LeRobot deliberately builds on a few accelerate internals that are not covered by a public
stability promise. These tests exist to fail LOUDLY on a
dependency upgrade — on a CPU runner, before any distributed job can be corrupted — whenever one
of those seams moves. If a canary fails, re-audit the corresponding integration seam before bumping
the pin; do not simply update the assertion.
"""
import inspect
import pytest
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
def test_fsdp_checkpoint_name_constants():
"""Checkpoint dir names are imported from accelerate; the on-disk layout depends on them."""
from accelerate.utils.constants import FSDP_MODEL_NAME, OPTIMIZER_NAME
assert FSDP_MODEL_NAME == "pytorch_model_fsdp"
assert OPTIMIZER_NAME == "optimizer"
def test_parallelism_config_mesh_dim_contract():
"""FSDP2 shards over the flattened dp_shard_cp dim; the dataloader keys on exact root names."""
from accelerate.parallelism_config import ParallelismConfig
pc = ParallelismConfig(dp_replicate_size=2, dp_shard_size=2, cp_size=2)
assert pc.fsdp_dim_names == ["dp_replicate", "dp_shard_cp"]
assert pc.dp_shard_cp_dim_names == ["dp_shard", "cp"]
assert pc.dp_cp_dim_names == ["dp_replicate", "dp_shard", "cp"]
# Degenerate FSDP-only case still shards over the flattened name.
pc_fsdp = ParallelismConfig(dp_replicate_size=1, dp_shard_size=4)
assert pc_fsdp.fsdp_dim_names == ["dp_shard_cp"]
def test_accelerator_accepts_parallelism_config():
from accelerate import Accelerator
params = inspect.signature(Accelerator.__init__).parameters
assert "parallelism_config" in params
assert "fsdp_plugin" in params
assert "gradient_accumulation_plugin" in params
def test_dataloader_is_mesh_aware():
"""prepare_data_loader must accept the device mesh that makes CP peers share batches."""
from accelerate.data_loader import prepare_data_loader
assert "torch_device_mesh" in inspect.signature(prepare_data_loader).parameters
def test_cp_mask_stripping_hook_seam():
"""finalize_sharded_policy strips this exact hook.
If accelerate renames or moves it, the strip becomes a silent no-op and CP training would
inherit mask-corrupting hooks — hence a canary rather than a runtime hasattr.
"""
from accelerate.big_modeling import _attach_context_parallel_hooks
assert callable(_attach_context_parallel_hooks)
assert _attach_context_parallel_hooks.__module__ == "accelerate.big_modeling"
def test_fsdp_plugin_mirrored_fields_exist():
"""AcceleratorConfig mirrors a plain-typed subset of the plugin; the fields must survive."""
from accelerate.utils import FullyShardedDataParallelPlugin
fields = {f.name for f in FullyShardedDataParallelPlugin.__dataclass_fields__.values()}
assert {
"fsdp_version",
"reshard_after_forward",
"auto_wrap_policy",
"transformer_cls_names_to_wrap",
"min_num_params",
"cpu_offload",
"ignored_modules",
"activation_checkpointing",
"state_dict_type",
} <= fields
def test_merge_fsdp_weights_signature():
"""The DCP->safetensors converter is a thin wrapper over this accelerate utility."""
from accelerate.utils import merge_fsdp_weights
params = inspect.signature(merge_fsdp_weights).parameters
assert {"checkpoint_dir", "output_path", "safe_serialization"} <= set(params)
def test_fsdp_save_load_helpers_exist():
from accelerate.utils import (
load_fsdp_model,
load_fsdp_optimizer,
save_fsdp_model,
save_fsdp_optimizer,
)
for fn in (save_fsdp_model, load_fsdp_model, save_fsdp_optimizer, load_fsdp_optimizer):
assert callable(fn)
def test_torch_fsdp2_seams():
"""isinstance(FSDPModule) discrimination + non-forward entry registration + full gather."""
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict, # noqa: F401
)
from torch.distributed.fsdp import FSDPModule, register_fsdp_forward_method # noqa: F401
options = inspect.signature(StateDictOptions).parameters
assert {"full_state_dict", "cpu_offload"} <= set(options)
def test_accelerate_version_floor():
import accelerate
from packaging import version
if version.parse(accelerate.__version__) < version.parse("1.14.0"):
pytest.fail(
f"accelerate {accelerate.__version__} < 1.14.0: the FSDP2 auto-wrap fallback fix "
"(#3999) and the bf16->fp32 master-weight upcast this design relies on are absent."
)
@@ -0,0 +1,70 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""The DCP wrappers must hand accelerate exact shard directories.
accelerate 1.14 resolves the load directory with a substring check ("optimizer" /
"pytorch_model_fsdp" in the path -> use as-is) while the save side joins the shard name
unconditionally, so a run path like `--job_name=optimizer_sweep` would save to
`training_state/optimizer_0/` but load from `training_state/` itself. Passing the exact
shard dir makes the containment check deterministically a no-op.
"""
from pathlib import Path
from types import SimpleNamespace
import pytest
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
def fake_accelerator() -> SimpleNamespace:
return SimpleNamespace(state=SimpleNamespace(fsdp_plugin=object()))
# A parent path that trips both of accelerate's substring checks at once.
POISONED_PARENT = Path("/outputs/train/optimizer_sweep_pytorch_model_fsdp_repro/training_state")
def test_load_sharded_optimizer_passes_exact_shard_dir(monkeypatch):
import accelerate.utils
from lerobot.distributed.checkpoint import load_sharded_optimizer
seen = {}
monkeypatch.setattr(
accelerate.utils,
"load_fsdp_optimizer",
lambda plugin, accelerator, optimizer, model, input_dir: seen.update(path=input_dir),
)
load_sharded_optimizer(fake_accelerator(), optimizer=object(), model=object(), input_dir=POISONED_PARENT)
assert seen["path"] == str(POISONED_PARENT / "optimizer_0")
assert isinstance(seen["path"], str) # str, never Path (accelerate does string checks)
def test_load_sharded_model_passes_exact_shard_dir(monkeypatch):
import accelerate.utils
from lerobot.distributed.checkpoint import load_sharded_model
seen = {}
monkeypatch.setattr(
accelerate.utils,
"load_fsdp_model",
lambda plugin, accelerator, model, input_dir: seen.update(path=input_dir),
)
load_sharded_model(fake_accelerator(), model=object(), input_dir=POISONED_PARENT)
assert seen["path"] == str(POISONED_PARENT / "pytorch_model_fsdp_0")
assert isinstance(seen["path"], str)
+486
View File
@@ -0,0 +1,486 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""End-to-end multi-GPU tests for the distributed core.
Sized for a 4-GPU CI lane, these tests execute the sharded code paths nothing else in the tree can
reach — ``fully_shard`` via ``accelerator.prepare``, the DCP branches of ``save_checkpoint`` /
``save_training_state`` / ``resume_after_prepare``, the collective gather inside
``save_pretrained``, and HSDP/DDP gradient reduction — against the tiny
``DummyCheckpointPolicy`` fixture on synthetic data (no datasets, no network, no site paths).
Run on a node with at least 4 GPUs::
pytest -m multigpu tests/distributed/test_multigpu_training.py -v
Mechanics:
- Plain pytest, no ``torchrun``: each test launches its own ranks with
``torch.multiprocessing.spawn`` (spawn start method) and a per-test free TCP port; workers set
the torchrun-equivalent env (``RANK``/``LOCAL_RANK``/``WORLD_SIZE``/``MASTER_*``) that
accelerate's ``env://`` initialization consumes.
- Deadlock watchdog (:func:`_spawn`): the spawn context is polled with a deadline instead of a
blocking join, so a hung collective — the exact failure mode the all-ranks contracts guard
against — fails the test with ``TimeoutError`` (all workers SIGKILLed) rather than hanging CI.
A worker exception propagates through ``ProcessContext.join``, which tears down the survivors.
- Workers configure accelerate exclusively through the LeRobot config mirrors
(``AcceleratorConfig.build(ParallelismConfig)`` after ``resolve(world_size)``) — the same
construction path ``make_accelerator`` takes; see :func:`_build_accelerator` for why the
factory itself is not called.
- Without GPUs every test skips (``torch.cuda.device_count()`` gate), so the file is safe to
collect and run in the CPU lanes.
"""
import json
import os
import socket
import time
from pathlib import Path
import pytest
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from safetensors.torch import load_file
from lerobot.common.train_utils import resume_after_prepare, resume_before_prepare, save_checkpoint
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import CheckpointFormat, TrainPipelineConfig
from lerobot.distributed.checkpoint import full_model_state_dict, is_sharded_module
from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR
# The spawned children re-import this module by name, so this import must resolve there too:
# torch.multiprocessing propagates the parent's sys.path through the spawn preparation data.
from tests.fixtures.dummy_checkpoint_policy import DummyCheckpointConfig, DummyCheckpointPolicy
SEED = 20260712
HIDDEN = 8 # DummyCheckpointPolicy is one Linear(hidden, hidden): 4 ranks shard dim 0 evenly
BATCH_SIZE = 2
SAVE_STEP = 2 # optimizer steps run before saving in the round-trip workers
PARITY_STEPS = 3
GA_UPDATES = 3
SAMPLES_PER_UPDATE = 4 # per rank per optimizer update — the fixed effective batch of test 5
GRAD_CLIP_NORM = 100.0 # generous: exercises the clip call without perturbing parity
# Generous headroom for cold NCCL init plus the lerobot re-import in 4 spawned children, while
# still bounding a deadlocked collective to minutes instead of a hung CI job.
WATCHDOG_TIMEOUT_S = 240.0
_JOIN_POLL_S = 5.0
# -------------------------------------------------------------------------------------------
# Spawn infrastructure
# -------------------------------------------------------------------------------------------
def _find_free_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.bind(("127.0.0.1", 0))
return sock.getsockname()[1]
def _spawn(world_size: int, worker, *args, timeout_s: float = WATCHDOG_TIMEOUT_S) -> None:
"""Run ``worker(rank, world_size, port, *args)`` on ``world_size`` fresh processes.
Watchdog approach: ``mp.spawn(join=False)`` returns a ``ProcessContext`` whose ``join`` is
polled under a deadline. On timeout every surviving worker is SIGKILLed and the test fails
with ``TimeoutError`` — a deadlock can never hang CI. When a worker raises, ``join`` itself
kills the remaining ranks and re-raises the worker's exception into the test.
"""
port = _find_free_port()
context = mp.spawn(worker, args=(world_size, port, *args), nprocs=world_size, join=False)
deadline = time.monotonic() + timeout_s
while not context.join(timeout=_JOIN_POLL_S):
if time.monotonic() >= deadline:
for process in context.processes:
if process.is_alive():
process.kill()
for process in context.processes:
process.join(timeout=10)
raise TimeoutError(
f"{getattr(worker, '__name__', worker)}: {world_size} workers still running "
f"after {timeout_s}s — presumed deadlock; all workers killed."
)
def _init_worker_env(rank: int, world_size: int, port: int) -> None:
"""Give the worker the torchrun-equivalent env accelerate's ``env://`` init consumes."""
# The tests configure accelerate through the config mirrors only; drop any accelerate env
# fallbacks inherited from the launching shell (what guard_against_env_interference would
# reject in production — here the env is simply owned by the test).
for name in list(os.environ):
if name.startswith(("FSDP_", "PARALLELISM_CONFIG_", "ACCELERATE_")):
del os.environ[name]
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(port)
os.environ["RANK"] = str(rank)
os.environ["LOCAL_RANK"] = str(rank)
os.environ["WORLD_SIZE"] = str(world_size)
# The fp32 parity tolerances below assume true-fp32 matmuls.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
# -------------------------------------------------------------------------------------------
# Shared building blocks
# -------------------------------------------------------------------------------------------
def _make_cfg(
world_size: int,
*,
dp_replicate: int = 1,
dp_shard: int = 1,
checkpoint_format: CheckpointFormat = CheckpointFormat.SAFETENSORS,
grad_accum: int = 1,
) -> TrainPipelineConfig:
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"), batch_size=BATCH_SIZE)
cfg.checkpoint_format = checkpoint_format
cfg.parallelism.dp_replicate = dp_replicate
cfg.parallelism.dp_shard = dp_shard
cfg.accelerator.mixed_precision = "no" # fp32 end to end: the parity tests depend on it
cfg.accelerator.gradient_accumulation.steps = grad_accum
# The dummy policy declares no _fsdp_wrap_modules; the size-based wrap policy shards its
# Linear without needing class names (the set_fsdp_wrap_modules no-op branch).
cfg.accelerator.fsdp.min_num_params = 1
cfg.parallelism.resolve(world_size)
return cfg
def _build_accelerator(cfg: TrainPipelineConfig):
"""``cfg.accelerator.build(cfg.parallelism)`` — make_accelerator's construction path.
Deliberately not ``make_accelerator`` itself: the factory additionally derives ``cpu=`` from
``cfg.trainable_config`` (no policy config is attached to these synthetic cfgs) and re-runs
the env guard — both owned explicitly by the tests (see ``_init_worker_env``).
"""
return cfg.accelerator.build(cfg.parallelism)
def _make_policy(seed: int) -> DummyCheckpointPolicy:
"""Identically seeded on every rank, so shard/replicate starts from one common init."""
torch.manual_seed(seed)
return DummyCheckpointPolicy(DummyCheckpointConfig(hidden=HIDDEN, device="cpu"))
def _batch(step: int, rank: int, device: torch.device) -> dict[str, torch.Tensor]:
"""Deterministic per-(step, rank) batch: every dp worker sees distinct, reproducible data."""
generator = torch.Generator().manual_seed(SEED + 1000 * step + rank)
return {"observation.state": torch.randn(BATCH_SIZE, HIDDEN, generator=generator).to(device)}
def _gather_full(model, optimizer) -> tuple[dict, dict]:
"""Full (unsharded) model + optimizer state via torch's DCP state-dict API — a COLLECTIVE.
With ``cpu_offload=True`` the dicts materialize on the main rank only; every other rank
receives a literal ``{}``.
"""
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
get_optimizer_state_dict,
)
options = StateDictOptions(full_state_dict=True, cpu_offload=True)
return (
get_model_state_dict(model, options=options),
get_optimizer_state_dict(model, optimizer, options=options),
)
def _assert_tree_equal(reference, actual, path: str) -> None:
"""Exact (bitwise for tensors) equality of nested state dicts, with a failing path."""
if isinstance(reference, torch.Tensor):
assert isinstance(actual, torch.Tensor), f"{path}: {type(actual)} is not a tensor"
assert reference.dtype == actual.dtype, f"{path}: {reference.dtype} != {actual.dtype}"
assert reference.shape == actual.shape, f"{path}: {reference.shape} != {actual.shape}"
assert torch.equal(reference.cpu(), actual.cpu()), f"{path}: tensor values differ"
elif isinstance(reference, dict):
assert isinstance(actual, dict), f"{path}: {type(actual)} is not a dict"
assert set(reference) == set(actual), f"{path}: keys {set(reference) ^ set(actual)} differ"
for key in reference:
_assert_tree_equal(reference[key], actual[key], f"{path}.{key}")
elif isinstance(reference, list | tuple):
assert type(reference) is type(actual) and len(reference) == len(actual), path
for index, (ref_item, actual_item) in enumerate(zip(reference, actual, strict=True)):
_assert_tree_equal(ref_item, actual_item, f"{path}[{index}]")
else:
assert reference == actual, f"{path}: {reference!r} != {actual!r}"
# -------------------------------------------------------------------------------------------
# Workers (module-level: torch.multiprocessing.spawn pickles them by reference)
# -------------------------------------------------------------------------------------------
def _train_and_save_worker(rank: int, world_size: int, port: int, tmp_dir: str, fmt_value: str) -> None:
"""FSDP2 (dp_shard=world_size): train SAVE_STEP steps, save_checkpoint, store the gathered
full model/optimizer state as the rank-0 reference for the resume workers."""
_init_worker_env(rank, world_size, port)
tmp = Path(tmp_dir)
fmt = CheckpointFormat(fmt_value)
cfg = _make_cfg(world_size, dp_shard=world_size, checkpoint_format=fmt)
accelerator = _build_accelerator(cfg)
policy = _make_policy(SEED)
optimizer = torch.optim.Adam(policy.parameters(), lr=1e-2)
# FSDP2 requires model and optimizer in one prepare() call (accelerate rebinds param groups).
policy, optimizer = accelerator.prepare(policy, optimizer)
assert is_sharded_module(accelerator.unwrap_model(policy)), "prepare() did not shard the policy"
for step in range(SAVE_STEP):
loss, _ = policy(_batch(step, rank, accelerator.device))
accelerator.backward(loss)
optimizer.step()
optimizer.zero_grad()
checkpoint_dir = tmp / "checkpoint"
save_checkpoint(
checkpoint_dir, step=SAVE_STEP, cfg=cfg, policy=policy, optimizer=optimizer, accelerator=accelerator
)
model_state, optimizer_state = _gather_full(policy, optimizer)
if accelerator.is_main_process:
from accelerate.utils.constants import FSDP_MODEL_NAME, OPTIMIZER_NAME
pretrained_dir = checkpoint_dir / PRETRAINED_MODEL_DIR
assert (pretrained_dir / f"{FSDP_MODEL_NAME}_0").is_dir() == fmt.wants_dcp
assert (pretrained_dir / "model.safetensors").is_file() == fmt.wants_safetensors
assert (pretrained_dir / "config.json").is_file()
assert (pretrained_dir / "train_config.json").is_file()
# Sharded runs always use the DCP optimizer channel, never the safetensors one.
assert (checkpoint_dir / TRAINING_STATE_DIR / f"{OPTIMIZER_NAME}_0").is_dir()
assert not (checkpoint_dir / TRAINING_STATE_DIR / "optimizer_state.safetensors").exists()
torch.save({"model": model_state, "optimizer": optimizer_state}, tmp / "reference_state.pt")
accelerator.wait_for_everyone()
dist.destroy_process_group()
def _resume_and_verify_worker(rank: int, world_size: int, port: int, tmp_dir: str, fmt_value: str) -> None:
"""Two-phase resume at dp_shard=world_size; the gathered state must match the saved
reference exactly (DCP round-trips are bit-exact)."""
_init_worker_env(rank, world_size, port)
tmp = Path(tmp_dir)
cfg = _make_cfg(world_size, dp_shard=world_size, checkpoint_format=CheckpointFormat(fmt_value))
cfg.checkpoint_path = tmp / "checkpoint"
accelerator = _build_accelerator(cfg)
assert resume_before_prepare(cfg) == SAVE_STEP # phase 1: RNG + step counter only
# Deliberately different init: the DCP load must overwrite every parameter.
policy = _make_policy(SEED + 1)
optimizer = torch.optim.Adam(policy.parameters(), lr=1e-2)
policy, optimizer = accelerator.prepare(policy, optimizer)
resume_after_prepare(cfg, accelerator, policy, optimizer, None) # phase 2: DCP reshard-load
model_state, optimizer_state = _gather_full(policy, optimizer)
if accelerator.is_main_process:
reference = torch.load(tmp / "reference_state.pt", map_location="cpu", weights_only=True)
_assert_tree_equal(reference["model"], model_state, "model")
_assert_tree_equal(reference["optimizer"], optimizer_state, "optimizer")
accelerator.wait_for_everyone()
dist.destroy_process_group()
def _loss_parity_worker(
rank: int, world_size: int, port: int, tmp_dir: str, dp_replicate: int, dp_shard: int, tag: str
) -> None:
"""Train PARITY_STEPS fp32 steps on per-rank deterministic data; rank 0 records the
dp-mean loss of every step. Gradient averaging spans the same rank set in any (R, S)
factorization of the world, so the loss trajectory is topology-invariant."""
_init_worker_env(rank, world_size, port)
cfg = _make_cfg(world_size, dp_replicate=dp_replicate, dp_shard=dp_shard)
accelerator = _build_accelerator(cfg)
policy = _make_policy(SEED)
optimizer = torch.optim.SGD(policy.parameters(), lr=0.05)
policy, optimizer = accelerator.prepare(policy, optimizer)
assert is_sharded_module(accelerator.unwrap_model(policy)) == (dp_shard > 1)
per_step_losses = []
for step in range(PARITY_STEPS):
loss, _ = policy(_batch(step, rank, accelerator.device))
per_step_losses.append(accelerator.gather(loss.detach().reshape(1)).double().mean().item())
accelerator.backward(loss)
optimizer.step()
optimizer.zero_grad()
if accelerator.is_main_process:
(Path(tmp_dir) / f"losses_{tag}.json").write_text(json.dumps(per_step_losses))
accelerator.wait_for_everyone()
dist.destroy_process_group()
def _save_pretrained_all_ranks_worker(rank: int, world_size: int, port: int, tmp_dir: str) -> None:
"""The all-ranks contract: every rank calls save_pretrained, the
collective gather completes (watchdog proves no deadlock), and only rank 0 writes files."""
_init_worker_env(rank, world_size, port)
cfg = _make_cfg(world_size, dp_shard=world_size)
accelerator = _build_accelerator(cfg)
policy = _make_policy(SEED)
# FSDP2 prepare requires an optimizer alongside the model even though this test never steps it.
optimizer = torch.optim.SGD(policy.parameters(), lr=0.1)
policy, optimizer = accelerator.prepare(policy, optimizer)
unwrapped = accelerator.unwrap_model(policy)
assert is_sharded_module(unwrapped)
# Gather semantics: the full dict materializes on the main rank; every other rank
# receives the literal empty dict.
reference = full_model_state_dict(unwrapped)
if accelerator.is_main_process:
assert set(reference) == {"net.weight", "net.bias"}
else:
assert reference == {}
# Every rank targets its own directory so writes are attributable per rank.
target = Path(tmp_dir) / f"rank_{rank}"
unwrapped.save_pretrained(target)
accelerator.wait_for_everyone()
if accelerator.is_main_process:
weights = load_file(target / "model.safetensors")
assert set(weights) == set(reference)
for key, tensor in reference.items():
assert torch.equal(weights[key], tensor), key
assert (target / "config.json").is_file()
else:
assert list(target.rglob("*")) == [], f"rank {rank} wrote files despite the rank-0 gate"
dist.destroy_process_group()
def _grad_accum_worker(
rank: int, world_size: int, port: int, tmp_dir: str, micro_batch_size: int, grad_accum: int, tag: str
) -> None:
"""DDP fp32 with the exact accumulate/clip/step/zero_grad pattern of
``lerobot_train.update_policy``; rank 0 records the final weights."""
_init_worker_env(rank, world_size, port)
assert micro_batch_size * grad_accum == SAMPLES_PER_UPDATE # fixed effective batch
cfg = _make_cfg(world_size, dp_replicate=world_size, grad_accum=grad_accum)
accelerator = _build_accelerator(cfg)
# The GradientAccumulationPlugin wiring, un-overridden by any env fallback.
assert accelerator.gradient_accumulation_steps == grad_accum
policy = _make_policy(SEED)
optimizer = torch.optim.SGD(policy.parameters(), lr=0.05)
policy, optimizer = accelerator.prepare(policy, optimizer)
# One fixed per-rank sample stream, consumed in order by both variants: update k always
# covers rows [k * SAMPLES_PER_UPDATE, (k + 1) * SAMPLES_PER_UPDATE).
generator = torch.Generator().manual_seed(SEED + 7919 * rank)
stream = torch.randn(GA_UPDATES * SAMPLES_PER_UPDATE, HIDDEN, generator=generator)
updates_applied = 0
for micro_step in range(GA_UPDATES * grad_accum):
rows = stream[micro_step * micro_batch_size : (micro_step + 1) * micro_batch_size]
batch = {"observation.state": rows.to(accelerator.device)}
# update_policy's pattern: accumulate() suppresses grad sync and rescales the loss on
# non-final micro-batches, and AcceleratedOptimizer makes step()/zero_grad() no-ops
# until sync_gradients is True.
with accelerator.accumulate(policy):
loss, _ = policy(batch)
accelerator.backward(loss)
if accelerator.sync_gradients:
accelerator.clip_grad_norm_(policy.parameters(), GRAD_CLIP_NORM)
updates_applied += 1
optimizer.step()
optimizer.zero_grad()
assert updates_applied == GA_UPDATES # exactly one optimizer update per accumulation window
if accelerator.is_main_process:
state = {key: value.cpu() for key, value in accelerator.unwrap_model(policy).state_dict().items()}
torch.save(state, Path(tmp_dir) / f"weights_{tag}.pt")
accelerator.wait_for_everyone()
dist.destroy_process_group()
# -------------------------------------------------------------------------------------------
# Tests
# -------------------------------------------------------------------------------------------
@pytest.mark.multigpu
@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 GPUs")
def test_fsdp2_train_save_resume_round_trip(tmp_path):
"""FSDP2 dp_shard=4, checkpoint_format=safetensors_dcp: train -> save_checkpoint -> resume.
A second spawn resumes through the two-phase path and its gathered model weights and Adam
state tensors must match the pre-save gathered reference exactly (DCP round-trips are
bit-exact).
"""
fmt = CheckpointFormat.SAFETENSORS_AND_DCP.value
_spawn(4, _train_and_save_worker, str(tmp_path), fmt)
_spawn(4, _resume_and_verify_worker, str(tmp_path), fmt)
@pytest.mark.multigpu
@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 GPUs")
def test_hsdp_loss_parity_with_ddp(tmp_path):
"""Same seed and per-rank data: DDP (dp_replicate=4) vs HSDP (2x2), fp32, no AMP.
Both topologies average gradients over the same four ranks, so per-step dp-mean losses must
match within tolerance. Exact parity is not expected: DDP all-reduces where HSDP
reduce-scatters within the shard group and all-reduces across replicas, and the different
reduction orders accumulate fp32 rounding — rtol=1e-4 leaves orders of magnitude of headroom
over that noise while still catching any real divergence (wrong averaging, wrong data).
"""
_spawn(4, _loss_parity_worker, str(tmp_path), 4, 1, "ddp")
_spawn(4, _loss_parity_worker, str(tmp_path), 2, 2, "hsdp")
ddp_losses = json.loads((tmp_path / "losses_ddp.json").read_text())
hsdp_losses = json.loads((tmp_path / "losses_hsdp.json").read_text())
assert len(ddp_losses) == len(hsdp_losses) == PARITY_STEPS
for step, (ddp_loss, hsdp_loss) in enumerate(zip(ddp_losses, hsdp_losses, strict=True)):
assert hsdp_loss == pytest.approx(ddp_loss, rel=1e-4, abs=1e-6), f"step {step}"
@pytest.mark.multigpu
@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 GPUs")
def test_changed_topology_resume(tmp_path):
"""Save at dp_shard=4 (format=dcp), resume at dp_shard=2 on 2 ranks.
The DCP load reshards both the model weights and the optimizer state across the topology
change; the post-resume gathered state must equal the pre-save gathered reference exactly
(cross-topology resharding is runtime-verified).
"""
fmt = CheckpointFormat.DCP.value
_spawn(4, _train_and_save_worker, str(tmp_path), fmt)
_spawn(2, _resume_and_verify_worker, str(tmp_path), fmt)
@pytest.mark.multigpu
@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 GPUs")
def test_save_pretrained_all_ranks_no_deadlock(tmp_path):
"""dp_shard=4: save_pretrained on ALL ranks completes under the watchdog.
Rank 0 writes model.safetensors (+ config.json) whose tensors equal the gathered full state;
ranks 1-3 write nothing. A rank-gated call would deadlock in the collective gather and be
killed by :func:`_spawn`'s timeout — completing at all is half of what this test asserts.
"""
_spawn(4, _save_pretrained_all_ranks_worker, str(tmp_path))
@pytest.mark.multigpu
@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires 2 GPUs")
def test_gradient_accumulation_equivalence(tmp_path):
"""Fixed effective batch on 2-rank DDP fp32: (batch=4, GA=1) vs (batch=2, GA=2).
Both variants consume the identical per-rank sample stream in the same order for
GA_UPDATES optimizer updates, using update_policy's accumulate/clip/step pattern. The final
weights must agree: accumulate() rescales each micro-loss by 1/GA, so summed mean-of-2
gradients equal the mean-of-4 gradient up to fp32 summation order — hence allclose with
rtol=1e-5/atol=1e-6 (roughly 100x the observed associativity noise), not bitwise equality.
"""
_spawn(2, _grad_accum_worker, str(tmp_path), 4, 1, "ga1")
_spawn(2, _grad_accum_worker, str(tmp_path), 2, 2, "ga2")
ga1 = torch.load(tmp_path / "weights_ga1.pt", weights_only=True)
ga2 = torch.load(tmp_path / "weights_ga2.pt", weights_only=True)
assert set(ga1) == set(ga2) == {"net.weight", "net.bias"}
for key in ga1:
assert torch.allclose(ga1[key], ga2[key], rtol=1e-5, atol=1e-6), key
@@ -0,0 +1,121 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import pytest
from lerobot.configs.parallelism import ContextParallelConfig, ParallelismConfig
from lerobot.distributed import ParallelDims, guard_against_env_interference, is_main_process
from lerobot.distributed.factory import _ENV_OVERRIDE
class TestIsMainProcess:
def test_true_outside_distributed(self):
assert is_main_process() is True
class TestParallelDims:
def _resolved(self, world_size: int = 8, **kwargs) -> ParallelismConfig:
cfg = ParallelismConfig(**kwargs)
cfg.resolve(world_size)
return cfg
def test_from_resolved_config(self):
dims = ParallelDims.from_config(self._resolved(dp_replicate=2, dp_shard=4), 8, "cpu")
assert dims.dp_world_size == 8
assert dims.is_sharded
assert dims.cp_size == 1
assert dims.dp_rank == 0 # no process group in unit tests
def test_rejects_unresolved_config(self):
with pytest.raises(ValueError, match="resolve"):
ParallelDims.from_config(ParallelismConfig(dp_shard=-1), 8, "cpu")
def test_rejects_world_mismatch(self):
with pytest.raises(ValueError, match="world_size=4"):
ParallelDims.from_config(self._resolved(8), 4, "cpu")
def test_cp_mesh_reserved(self):
dims = ParallelDims(dp_replicate=1, dp_shard=2, ring=2, ulysses=2, world_size=8, device_type="cpu")
assert dims.dp_rank == 0 and dims.dp_world_size == 2
with pytest.raises(NotImplementedError):
dims.cp_mesh()
def test_cp_peers_share_dp_rank_arithmetic(self):
"""Row-major layout: cp is innermost, so dp_rank = global_rank // cp_size."""
dims = ParallelDims(dp_replicate=1, dp_shard=2, ring=1, ulysses=2, world_size=4, device_type="cpu")
# Without a process group the global rank is 0; the arithmetic contract is what matters.
assert dims.cp_size == 2
assert dims.dp_rank == 0 // dims.cp_size
def test_config_placeholder_degrees_flow_through(self):
cfg = ParallelismConfig(
dp_replicate=1,
dp_shard=2,
context_parallel=ContextParallelConfig(ring_degree=2, ulysses_degree=2),
)
# resolve() rejects cp>1 this round; ParallelDims math itself is already cp-aware.
dims = ParallelDims(
dp_replicate=cfg.dp_replicate,
dp_shard=cfg.dp_shard,
ring=cfg.context_parallel.ring_degree,
ulysses=cfg.context_parallel.ulysses_degree,
world_size=8,
device_type="cpu",
)
assert dims.dp_world_size == 2 and dims.cp_size == 4
class TestEnvGuard:
# Silent config overrides inside accelerate itself — the guard must catch them.
_POISON = (
"ACCELERATE_USE_FSDP",
"ACCELERATE_USE_PARALLELISM_CONFIG",
"ACCELERATE_GRADIENT_ACCUMULATION_STEPS",
)
def test_clean_env_passes(self, monkeypatch):
for name in self._POISON + (_ENV_OVERRIDE,):
monkeypatch.delenv(name, raising=False)
guard_against_env_interference()
@pytest.mark.parametrize("name", _POISON)
def test_accelerate_env_rejected_with_actionable_error(self, name, monkeypatch):
monkeypatch.delenv(_ENV_OVERRIDE, raising=False)
monkeypatch.setenv(name, "true")
with pytest.raises(RuntimeError, match=name):
guard_against_env_interference()
def test_override_acknowledges(self, monkeypatch):
monkeypatch.setenv("ACCELERATE_USE_FSDP", "true")
monkeypatch.setenv(_ENV_OVERRIDE, "1")
guard_against_env_interference()
def test_make_accelerator_rejects_format_after_sentinel_resolution(monkeypatch):
"""dp_shard=-1 counts as sharded at parse time but can resolve
to an unsharded run (world size 1), which would write a safetensors-only checkpoint whose
recorded checkpoint_format=dcp fails its own validation on resume."""
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import CheckpointFormat, TrainPipelineConfig
from lerobot.distributed.factory import make_accelerator
monkeypatch.delenv("WORLD_SIZE", raising=False)
cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"))
cfg.parallelism.dp_shard = -1
cfg.checkpoint_format = CheckpointFormat.DCP
cfg._validate_distributed() # passes: the sentinel is declared as sharded
with pytest.raises(ValueError, match="resolved to a non-sharded"):
make_accelerator(cfg)
+126
View File
@@ -0,0 +1,126 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""The declarative policy surface and its distributed-side consumers."""
from types import SimpleNamespace
import pytest
import torch
from torch import nn
from lerobot.configs.accelerator import FSDPConfig
from lerobot.distributed import set_fsdp_wrap_modules, strip_accelerate_cp_hooks
from lerobot.policies.pretrained import PreTrainedPolicy
class TestDeclarativeAttributes:
def test_base_defaults(self):
assert PreTrainedPolicy._fsdp_wrap_modules is None
assert PreTrainedPolicy._fsdp_forward_methods == ("select_action", "predict_action_chunk")
assert PreTrainedPolicy.supports_gradient_checkpointing is False
assert PreTrainedPolicy._cp_plan is None
def test_act_wrap_units_name_real_classes(self):
"""The declared class names must track the modeling code — this test pins the drift."""
from lerobot.policies.act import modeling_act
for name in modeling_act.ACTPolicy._fsdp_wrap_modules:
assert isinstance(getattr(modeling_act, name), type), name
def test_fastwam_wrap_units_name_real_classes(self):
from lerobot.policies.fastwam import modeling_fastwam
from lerobot.policies.fastwam.wan import modular
for name in modeling_fastwam.FastWAMPolicy._fsdp_wrap_modules:
assert isinstance(getattr(modular, name), type), name
class _SelfAttn(nn.Module):
def forward(self, x, attention_mask=None, is_causal=False):
return x, attention_mask, is_causal
class _TinyModel(nn.Module):
def __init__(self):
super().__init__()
self.self_attn = _SelfAttn()
class TestStripAccelerateCpHooks:
def test_strips_the_real_accelerate_hook_and_restores_mask_semantics(self):
"""Attach accelerate's actual mask-stripping hook, strip it, verify masks survive."""
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
from accelerate.big_modeling import _attach_context_parallel_hooks
model = _TinyModel()
mask = torch.ones(2, 2)
_attach_context_parallel_hooks(model)
_, hooked_mask, hooked_causal = model.self_attn(torch.zeros(1), attention_mask=mask)
assert hooked_mask is None and hooked_causal is True # the hazard is real
assert strip_accelerate_cp_hooks(model) == 1
_, clean_mask, clean_causal = model.self_attn(torch.zeros(1), attention_mask=mask)
assert clean_mask is mask and clean_causal is False
assert not model.self_attn._forward_pre_hooks
assert not model.self_attn._forward_pre_hooks_with_kwargs
def test_user_hooks_survive(self):
model = _TinyModel()
model.self_attn.register_forward_pre_hook(lambda m, args: None)
assert strip_accelerate_cp_hooks(model) == 0
assert len(model.self_attn._forward_pre_hooks) == 1
class _DeclaredPolicy:
_fsdp_wrap_modules = ["DeclaredBlock"]
class _UndeclaredPolicy:
_fsdp_wrap_modules = None
def _accelerator_with(plugin) -> SimpleNamespace:
return SimpleNamespace(state=SimpleNamespace(fsdp_plugin=plugin))
class TestSetFsdpWrapModules:
@pytest.fixture(autouse=True)
def _requires_accelerate(self):
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
def test_policy_declaration_fills_plugin(self):
plugin = FSDPConfig().build_plugin()
set_fsdp_wrap_modules(_accelerator_with(plugin), _DeclaredPolicy())
assert plugin.transformer_cls_names_to_wrap == ["DeclaredBlock"]
def test_user_override_wins(self):
plugin = FSDPConfig(wrap_modules=["UserBlock"]).build_plugin()
set_fsdp_wrap_modules(_accelerator_with(plugin), _DeclaredPolicy())
assert plugin.transformer_cls_names_to_wrap == ["UserBlock"]
def test_no_wrap_source_fails_loudly(self):
plugin = FSDPConfig().build_plugin()
with pytest.raises(ValueError, match="_fsdp_wrap_modules"):
set_fsdp_wrap_modules(_accelerator_with(plugin), _UndeclaredPolicy())
def test_size_based_policy_needs_no_names(self):
plugin = FSDPConfig(min_num_params=1024).build_plugin()
set_fsdp_wrap_modules(_accelerator_with(plugin), _UndeclaredPolicy())
assert plugin.transformer_cls_names_to_wrap is None
def test_non_sharded_run_is_noop(self):
set_fsdp_wrap_modules(_accelerator_with(None), _UndeclaredPolicy())
+87
View File
@@ -0,0 +1,87 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""A minimal real PreTrainedPolicy for checkpoint/publish unit tests (CPU, tiny)."""
from dataclasses import dataclass
import torch
from torch import Tensor, nn
from lerobot.configs.policies import PreTrainedConfig
from lerobot.optim.optimizers import AdamConfig, OptimizerConfig
from lerobot.policies.pretrained import PreTrainedPolicy
@PreTrainedConfig.register_subclass("dummy_checkpoint")
@dataclass
class DummyCheckpointConfig(PreTrainedConfig):
hidden: int = 4
@property
def observation_delta_indices(self) -> list | None:
return None
@property
def action_delta_indices(self) -> list | None:
return None
@property
def reward_delta_indices(self) -> list | None:
return None
def get_optimizer_preset(self) -> OptimizerConfig:
return AdamConfig(lr=1e-3)
def get_scheduler_preset(self) -> None:
return None
def validate_features(self) -> None:
pass
class DummyCheckpointPolicy(PreTrainedPolicy):
config_class = DummyCheckpointConfig
name = "dummy_checkpoint"
def __init__(self, config: DummyCheckpointConfig, **kwargs):
super().__init__(config)
self.net = nn.Linear(config.hidden, config.hidden)
def get_optim_params(self) -> dict:
return self.parameters()
def reset(self) -> None:
pass
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
out = self.net(batch["observation.state"])
return out.mean(), None
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
return self.net(batch["observation.state"])
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
return self.net(batch["observation.state"])
def make_dummy_policy(repo_id: str | None = None) -> DummyCheckpointPolicy:
config = DummyCheckpointConfig(device="cpu")
if repo_id is not None:
config.repo_id = repo_id
policy = DummyCheckpointPolicy(config)
with torch.no_grad():
policy.net.weight.fill_(0.5)
return policy
-39
View File
@@ -20,7 +20,6 @@ from lerobot.optim.optimizers import (
MultiAdamConfig, MultiAdamConfig,
SGDConfig, SGDConfig,
load_optimizer_state, load_optimizer_state,
load_optimizer_state_dict,
save_optimizer_state, save_optimizer_state,
) )
from lerobot.utils.constants import ( from lerobot.utils.constants import (
@@ -66,44 +65,6 @@ def test_save_and_load_optimizer_state(model_params, optimizer, tmp_path):
torch.testing.assert_close(optimizer.state_dict(), loaded_optimizer.state_dict()) torch.testing.assert_close(optimizer.state_dict(), loaded_optimizer.state_dict())
def test_save_and_load_fsdp_optimizer_state_dict_roundtrip(tmp_path):
"""The FSDP full optimizer state dict is keyed by parameter FQNs (dotted strings), not the
integer indices of the single-GPU path. Verify it survives the safetensors save -> read
round-trip used by the FSDP save/resume path (save_optimizer_state(optim_state_dict=...) then
load_optimizer_state_dict), which the flatten/unflatten "/" separator must not corrupt."""
full_osd = {
"state": {
"model.layers.0.weight": {
"step": torch.tensor(3.0),
"exp_avg": torch.randn(4, 4),
"exp_avg_sq": torch.randn(4, 4),
},
"model.layers.0.bias": {
"step": torch.tensor(3.0),
"exp_avg": torch.randn(4),
"exp_avg_sq": torch.randn(4),
},
},
"param_groups": [
{"lr": 1e-4, "betas": [0.9, 0.999], "eps": 1e-8, "weight_decay": 0.0, "params": [0, 1]}
],
}
save_optimizer_state(
torch.optim.Adam([torch.nn.Parameter(torch.randn(1))]), tmp_path, optim_state_dict=full_osd
)
assert (tmp_path / OPTIMIZER_STATE).is_file()
assert (tmp_path / OPTIMIZER_PARAM_GROUPS).is_file()
loaded = load_optimizer_state_dict(tmp_path)
# FQN keys must be preserved verbatim (not int-cast, not split on their dots).
assert set(loaded["state"].keys()) == set(full_osd["state"].keys())
for fqn, sub in full_osd["state"].items():
for k, v in sub.items():
torch.testing.assert_close(loaded["state"][fqn][k], v)
assert loaded["param_groups"] == full_osd["param_groups"]
@pytest.fixture @pytest.fixture
def base_params_dict(): def base_params_dict():
return { return {
+8 -4
View File
@@ -301,8 +301,12 @@ def test_save_and_load_pretrained(dummy_dataset_metadata, tmp_path, policy_name:
torch.testing.assert_close(list(policy.parameters()), list(loaded_policy.parameters()), rtol=0, atol=0) torch.testing.assert_close(list(policy.parameters()), list(loaded_policy.parameters()), rtol=0, atol=0)
def test_save_pretrained_with_state_dict(dummy_dataset_metadata, tmp_path): def test_save_pretrained_single_file_artifact(dummy_dataset_metadata, tmp_path):
"""Exercise the FSDP checkpoint path: save_pretrained with a pre-gathered state_dict.""" """The distributable checkpoint is one unsharded safetensors file.
The former `state_dict=` variant of this test died with the #3810 save override: the
kwarg would now be silently swallowed by HubMixin's **push_to_hub_kwargs.
"""
policy_cls = get_policy_class("act") policy_cls = get_policy_class("act")
policy_cfg = make_policy_config("act") policy_cfg = make_policy_config("act")
features = dataset_to_policy_features(dummy_dataset_metadata.features) features = dataset_to_policy_features(dummy_dataset_metadata.features)
@@ -313,8 +317,8 @@ def test_save_pretrained_with_state_dict(dummy_dataset_metadata, tmp_path):
policy = policy_cls(policy_cfg) policy = policy_cls(policy_cfg)
policy.to(policy_cfg.device) policy.to(policy_cfg.device)
save_dir = tmp_path / "fsdp_state_dict" save_dir = tmp_path / "single_file_artifact"
policy.save_pretrained(save_dir, state_dict=policy.state_dict()) policy.save_pretrained(save_dir)
# A single, unsharded safetensors file (no sharded set + index). # A single, unsharded safetensors file (no sharded set + index).
assert (save_dir / SAFETENSORS_SINGLE_FILE).is_file() assert (save_dir / SAFETENSORS_SINGLE_FILE).is_file()
@@ -16,7 +16,6 @@ from conftest import (
make_config, make_config,
set_seed_all, set_seed_all,
) # noqa: E402 ) # noqa: E402
from lerobot.policies.vla_jepa.action_head import ( # noqa: E402 from lerobot.policies.vla_jepa.action_head import ( # noqa: E402
VLAJEPAActionHead, VLAJEPAActionHead,
) )
@@ -3,8 +3,8 @@
from __future__ import annotations from __future__ import annotations
import pytest import pytest
from conftest import ACTION_DIM, ACTION_HORIZON, IMAGE_SIZE, NUM_VIDEO_FRAMES, STATE_DIM, make_config
from conftest import ACTION_DIM, ACTION_HORIZON, IMAGE_SIZE, NUM_VIDEO_FRAMES, STATE_DIM, make_config
from lerobot.configs.types import FeatureType, PolicyFeature from lerobot.configs.types import FeatureType, PolicyFeature
from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig
from lerobot.utils.constants import ACTION, OBS_IMAGES, OBS_STATE from lerobot.utils.constants import ACTION, OBS_IMAGES, OBS_STATE
-1
View File
@@ -32,7 +32,6 @@ from conftest import ( # noqa: E402
make_train_batch, make_train_batch,
set_seed_all, set_seed_all,
) )
from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig # noqa: E402 from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig # noqa: E402
from lerobot.policies.vla_jepa.modeling_vla_jepa import VLAJEPAPolicy # noqa: E402 from lerobot.policies.vla_jepa.modeling_vla_jepa import VLAJEPAPolicy # noqa: E402
from lerobot.utils.constants import ACTION # noqa: E402 from lerobot.utils.constants import ACTION # noqa: E402
+86 -26
View File
@@ -22,6 +22,7 @@ from types import SimpleNamespace
import pytest import pytest
import torch import torch
from lerobot.common.train_utils import generate_model_card
from lerobot.configs.rewards import RewardModelConfig from lerobot.configs.rewards import RewardModelConfig
from lerobot.optim.optimizers import AdamWConfig from lerobot.optim.optimizers import AdamWConfig
from lerobot.rewards.pretrained import PreTrainedRewardModel from lerobot.rewards.pretrained import PreTrainedRewardModel
@@ -326,7 +327,7 @@ def test_train_pipeline_config_from_pretrained_strips_legacy_rabc_when_disabled(
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# PreTrainedRewardModel hub upload: push_model_to_hub + generate_model_card. # PreTrainedRewardModel hub upload: publish_trained_model + generate_model_card.
# We test the generation side (offline) fully, and the upload side with HfApi # We test the generation side (offline) fully, and the upload side with HfApi
# mocked so nothing actually hits the network. # mocked so nothing actually hits the network.
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -336,6 +337,13 @@ def _make_dummy_reward_model(**config_kwargs):
return _DummyHubReward(_DummyHubRewardConfig(**config_kwargs)), _DummyHubRewardConfig return _DummyHubReward(_DummyHubRewardConfig(**config_kwargs)), _DummyHubRewardConfig
def _make_train_cfg(dataset_repo_id: str):
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import TrainPipelineConfig
return TrainPipelineConfig(dataset=DatasetConfig(repo_id=dataset_repo_id))
@pytest.fixture @pytest.fixture
def _offline_model_card(monkeypatch): def _offline_model_card(monkeypatch):
"""``ModelCard.validate`` does a live ``POST`` to huggingface.co — bypass it """``ModelCard.validate`` does a live ``POST`` to huggingface.co — bypass it
@@ -353,12 +361,7 @@ def test_reward_model_generate_model_card_renders_expected_fields(_offline_model
tags=["robot", "sim"], tags=["robot", "sim"],
) )
card = model.generate_model_card( card = generate_model_card(model.config, cfg=_make_train_cfg("user/my_dataset"))
dataset_repo_id="user/my_dataset",
model_type=model.config.type,
license=model.config.license,
tags=model.config.tags,
)
# Metadata (YAML header) — ModelCardData fields. # Metadata (YAML header) — ModelCardData fields.
assert card.data.license == "mit" assert card.data.license == "mit"
@@ -380,21 +383,16 @@ def test_reward_model_generate_model_card_uses_default_license(_offline_model_ca
"""When config.license is None the card falls back to apache-2.0.""" """When config.license is None the card falls back to apache-2.0."""
model, _ = _make_dummy_reward_model() model, _ = _make_dummy_reward_model()
card = model.generate_model_card( card = generate_model_card(model.config, cfg=_make_train_cfg("user/my_dataset"))
dataset_repo_id="user/my_dataset",
model_type=model.config.type,
license=model.config.license,
tags=None,
)
assert card.data.license == "apache-2.0" assert card.data.license == "apache-2.0"
def test_reward_model_push_model_to_hub_uploads_expected_files(monkeypatch, _offline_model_card): def test_publish_trained_model_uploads_expected_reward_files(monkeypatch, _offline_model_card):
"""``push_model_to_hub`` must: """Publishing a reward model through ``publish_trained_model`` must:
1. create the repo, 1. create the repo,
2. assemble a temp folder with weights + config.json + train_config.json + README.md, 2. push the model through ``HubMixin.push_to_hub`` (weights + config.json),
3. call ``api.upload_folder`` on that folder. 3. upload a bundle sidecar with train_config.json + the reward-specific README.md.
All network calls are mocked. All network calls are mocked.
""" """
from huggingface_hub.constants import CONFIG_NAME from huggingface_hub.constants import CONFIG_NAME
@@ -430,18 +428,80 @@ def test_reward_model_push_model_to_hub_uploads_expected_files(monkeypatch, _off
uploaded["files"] = sorted(p.name for p in Path(folder_path).iterdir()) uploaded["files"] = sorted(p.name for p in Path(folder_path).iterdir())
return fake_commit_info return fake_commit_info
from lerobot.rewards import pretrained as reward_pretrained import lerobot.common.train_utils as train_utils
import lerobot.utils.hub as hub_module
from lerobot.common.train_utils import publish_trained_model
monkeypatch.setattr(reward_pretrained, "HfApi", lambda *a, **kw: _FakeHfApi()) all_files: set[str] = set()
model.push_model_to_hub(train_cfg) class _RecordingFakeHfApi(_FakeHfApi):
def __init__(self, *args, **kwargs):
pass
def upload_folder(self, *, repo_id, repo_type, folder_path, commit_message, **_kwargs):
result = super().upload_folder(
repo_id=repo_id,
repo_type=repo_type,
folder_path=folder_path,
commit_message=commit_message,
)
all_files.update(uploaded["files"])
return result
monkeypatch.setattr(train_utils, "HfApi", _RecordingFakeHfApi)
monkeypatch.setattr(hub_module, "HfApi", _RecordingFakeHfApi)
publish_trained_model(train_cfg, model, None, None, dataset_meta=None)
assert uploaded["create_repo_id"] == "user/my_reward" assert uploaded["create_repo_id"] == "user/my_reward"
assert uploaded["upload_repo_id"] == "user/my_reward" assert uploaded["upload_repo_id"] == "user/my_reward"
assert uploaded["upload_repo_type"] == "model" assert uploaded["upload_repo_type"] == "model"
assert uploaded["commit_message"] == "Upload reward model weights, train config and readme" # Minimum required files across the publish commits.
# Minimum required files that must be uploaded with a reward model. assert CONFIG_NAME in all_files # config.json (model commit)
assert CONFIG_NAME in uploaded["files"] # config.json assert TRAIN_CONFIG_NAME in all_files # train_config.json (bundle commit)
assert TRAIN_CONFIG_NAME in uploaded["files"] # train_config.json assert "README.md" in all_files # reward-specific card (bundle commit)
assert "README.md" in uploaded["files"] assert any(name.endswith(".safetensors") for name in all_files) # weights (model commit)
assert any(name.endswith(".safetensors") for name in uploaded["files"])
def test_save_pretrained_writes_nothing_off_main_rank(tmp_path, monkeypatch):
"""save_checkpoint calls save_pretrained on every rank; the
reward serializer must gate its writes so DDP replicas do not race on the same files."""
import lerobot.distributed.utils as dist_utils
model, _ = _make_dummy_reward_model()
monkeypatch.setattr(dist_utils, "is_main_process", lambda: False)
model.save_pretrained(tmp_path)
assert not any(tmp_path.iterdir())
def test_reward_model_push_model_to_hub_shim_warns_and_publishes(monkeypatch, _offline_model_card):
"""The deprecated ``push_model_to_hub`` stays callable, delegating to the publisher."""
from huggingface_hub.constants import CONFIG_NAME
import lerobot.common.train_utils as train_utils
import lerobot.utils.hub as hub_module
from lerobot.configs.train import TRAIN_CONFIG_NAME
all_files: set[str] = set()
class _FakeHfApi:
def __init__(self, *args, **kwargs):
pass
def create_repo(self, repo_id, private=None, exist_ok=False, **kwargs):
return SimpleNamespace(repo_id=repo_id)
def upload_folder(self, *, repo_id, folder_path, **_kwargs):
all_files.update(p.name for p in Path(folder_path).iterdir())
return SimpleNamespace(repo_url=SimpleNamespace(url=f"https://huggingface.co/{repo_id}"))
monkeypatch.setattr(train_utils, "HfApi", _FakeHfApi)
monkeypatch.setattr(hub_module, "HfApi", _FakeHfApi)
model, _ = _make_dummy_reward_model(repo_id="user/my_reward")
with pytest.warns(FutureWarning, match="push_model_to_hub is deprecated"):
model.push_model_to_hub(_make_train_cfg("user/my_dataset"))
assert CONFIG_NAME in all_files
assert TRAIN_CONFIG_NAME in all_files
assert "README.md" in all_files
+123
View File
@@ -0,0 +1,123 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""lerobot-convert-dcp: locating, converting, and graceful-degradation publishing."""
import logging
import shutil
from pathlib import Path
from types import SimpleNamespace
import pytest
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
import lerobot.distributed.checkpoint as dist_checkpoint
from lerobot.scripts.lerobot_convert_dcp import (
ConvertDcpConfig,
_locate_pretrained_dir,
_publish_converted,
convert_checkpoint,
)
from lerobot.utils.constants import PRETRAINED_MODEL_DIR
@pytest.fixture
def fake_merge(monkeypatch):
"""Stand in for accelerate.utils.merge_fsdp_weights: writes a marker safetensors file."""
import accelerate.utils
def merge(checkpoint_dir, output_path, safe_serialization=True, remove_checkpoint_dir=False):
assert isinstance(checkpoint_dir, str) and isinstance(output_path, str) # str, not Path
(Path(output_path) / "model.safetensors").write_bytes(b"merged")
# Mirror accelerate: the shard directory is removed by the merge itself, when asked.
if remove_checkpoint_dir:
shutil.rmtree(checkpoint_dir)
monkeypatch.setattr(accelerate.utils, "merge_fsdp_weights", merge)
def make_dcp_checkpoint(tmp_path: Path) -> Path:
pretrained = tmp_path / PRETRAINED_MODEL_DIR
dcp_dir = pretrained / "pytorch_model_fsdp_0"
dcp_dir.mkdir(parents=True)
(dcp_dir / "__0_0.distcp").write_bytes(b"shard")
(pretrained / "config.json").write_text("{}")
return tmp_path
class TestConvert:
def test_locate_accepts_step_dir_or_pretrained_dir(self, tmp_path):
step_dir = make_dcp_checkpoint(tmp_path)
pretrained = step_dir / PRETRAINED_MODEL_DIR
assert _locate_pretrained_dir(step_dir) == pretrained
assert _locate_pretrained_dir(pretrained) == pretrained
def test_convert_keeps_dcp_by_default(self, tmp_path, fake_merge):
step_dir = make_dcp_checkpoint(tmp_path)
out = convert_checkpoint(ConvertDcpConfig(checkpoint_dir=step_dir))
assert out.read_bytes() == b"merged"
assert (step_dir / PRETRAINED_MODEL_DIR / "pytorch_model_fsdp_0").is_dir()
def test_convert_delete_dcp(self, tmp_path, fake_merge):
step_dir = make_dcp_checkpoint(tmp_path)
convert_checkpoint(ConvertDcpConfig(checkpoint_dir=step_dir, delete_dcp=True))
assert not (step_dir / PRETRAINED_MODEL_DIR / "pytorch_model_fsdp_0").exists()
def test_missing_shards_error_names_the_format(self, tmp_path):
with pytest.raises(FileNotFoundError, match="checkpoint_format=dcp"):
convert_checkpoint(ConvertDcpConfig(checkpoint_dir=tmp_path))
class TestPublishGracefulDegradation:
def _mock_api(self, monkeypatch):
calls = {}
class FakeApi:
def create_repo(self, repo_id, private=None, exist_ok=False):
return SimpleNamespace(repo_id=repo_id)
def upload_folder(self, *, repo_id, folder_path, allow_patterns, **kwargs):
calls["repo_id"] = repo_id
calls["files"] = sorted(p.name for p in Path(folder_path).iterdir())
calls["allow_patterns"] = allow_patterns
return SimpleNamespace(repo_url=SimpleNamespace(url=f"https://huggingface.co/{repo_id}"))
import lerobot.scripts.lerobot_convert_dcp as mod
monkeypatch.setattr(mod, "HfApi", FakeApi)
return calls
def test_missing_train_config_warns_and_uploads_core(self, tmp_path, monkeypatch, caplog):
calls = self._mock_api(monkeypatch)
pretrained = make_dcp_checkpoint(tmp_path) / PRETRAINED_MODEL_DIR
(pretrained / "model.safetensors").write_bytes(b"w")
with caplog.at_level(logging.WARNING):
_publish_converted(pretrained, "user/converted", private=None)
assert any("train_config.json missing" in m for m in caplog.messages)
assert "model.safetensors" in calls["files"]
# The DCP shard directory is still on disk (--delete_dcp defaults to False) but the
# allow list admits neither `.distcp` shards nor their `.metadata` sidecar.
assert set(calls["allow_patterns"]) == {"*.safetensors", "*.json", "*.yaml", "*.md"}
# config.json is not parseable as a policy config here -> card skipped with a warning
assert any("model card" in m for m in caplog.messages)
def test_dcp_to_safetensors_passes_str_paths(self, tmp_path, fake_merge):
"""accelerate 1.14's DCP helpers do string containment checks."""
dcp_dir = tmp_path / "pytorch_model_fsdp_0"
dcp_dir.mkdir()
out = dist_checkpoint.dcp_to_safetensors(dcp_dir, tmp_path, delete_dcp=True)
assert out == tmp_path / "model.safetensors"
assert not dcp_dir.exists()
+314
View File
@@ -0,0 +1,314 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for the opt-in EMA shadow maintained by the training pipeline (--ema.enable=true)."""
import draccus
import numpy as np
import pytest
import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.configs.default import EMAConfig
from lerobot.configs.train import TrainPipelineConfig
from lerobot.datasets.lerobot_dataset import LeRobotDataset
from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR
DUMMY_REPO_ID = "dummy/repo"
DUMMY_STATE_DIM = 6
DUMMY_ACTION_DIM = 6
IMAGE_SIZE = 32
N_EPISODES = 2
EPISODE_LENGTH = 12
def test_ema_config_defaults_match_the_reference():
cfg = EMAConfig()
assert not cfg.enable
assert cfg.inv_gamma == 1.0
assert cfg.power == 0.75
assert cfg.update_after_step == 0
@pytest.mark.parametrize(
"kwargs",
[
{"min_decay": 0.5, "max_decay": 0.1},
{"max_decay": 1.5},
{"min_decay": -0.1},
{"inv_gamma": 0.0},
{"power": -1.0},
{"update_after_step": -1},
{"decay": 1.5},
{"decay": -0.1},
{"decay": 0.99, "min_decay": 0.5},
{"decay": 0.99, "max_decay": 0.9},
],
)
def test_ema_config_rejects_invalid_values(kwargs):
with pytest.raises(ValueError):
EMAConfig(**kwargs)
def test_ema_config_cli_parsing():
cfg = draccus.parse(
TrainPipelineConfig,
None,
args=[
f"--dataset.repo_id={DUMMY_REPO_ID}",
"--ema.enable=true",
"--ema.power=0.8",
"--ema.update_after_step=10",
],
)
assert cfg.ema.enable
assert cfg.ema.power == 0.8
assert cfg.ema.update_after_step == 10
def test_ema_config_cli_parsing_constant_decay():
cfg = draccus.parse(
TrainPipelineConfig,
None,
args=[
f"--dataset.repo_id={DUMMY_REPO_ID}",
"--ema.enable=true",
"--ema.decay=0.99",
],
)
assert cfg.ema.enable
assert cfg.ema.decay == 0.99
def test_ema_constant_decay_pins_the_schedule():
"""min_decay == max_decay clamps the warmup curve to a constant (how --ema.decay is implemented)."""
pytest.importorskip("diffusers")
from diffusers.training_utils import EMAModel
model = torch.nn.Linear(4, 4)
ema = EMAModel(
model.parameters(), decay=0.99, min_decay=0.99, use_ema_warmup=True, inv_gamma=1.0, power=0.75
)
# The first update is a hard copy (decay 0); every one after uses the constant decay.
for step in range(1, 6):
ema.step(model.parameters())
if step > 1:
assert ema.cur_decay_value == 0.99
def test_ema_weights_context_swaps_and_restores():
pytest.importorskip("diffusers")
from diffusers.training_utils import EMAModel
from lerobot.scripts.lerobot_train import _ema_weights
torch.manual_seed(0)
model = torch.nn.Linear(4, 4)
ema = EMAModel(model.parameters(), decay=0.9999, use_ema_warmup=True, inv_gamma=1.0, power=0.75)
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
for _ in range(3):
model(torch.randn(2, 4)).sum().backward()
optimizer.step()
optimizer.zero_grad()
ema.step(model.parameters())
live = [p.detach().clone() for p in model.parameters()]
with _ema_weights(ema, model):
swapped = [p.detach().clone() for p in model.parameters()]
restored = list(model.parameters())
assert any(not torch.equal(a, b) for a, b in zip(live, swapped, strict=True))
assert all(torch.equal(a, b.detach()) for a, b in zip(live, restored, strict=True))
def make_dummy_dataset(tmp_path):
features = {
"action": {"dtype": "float32", "shape": (DUMMY_ACTION_DIM,), "names": None},
"observation.state": {"dtype": "float32", "shape": (DUMMY_STATE_DIM,), "names": None},
"observation.images.top": {
"dtype": "image",
"shape": (IMAGE_SIZE, IMAGE_SIZE, 3),
"names": ["height", "width", "channel"],
},
}
root = tmp_path / "_dataset"
dataset = LeRobotDataset.create(repo_id=DUMMY_REPO_ID, fps=30, features=features, root=root)
rng = np.random.default_rng(0)
for ep_idx in range(N_EPISODES):
for _ in range(EPISODE_LENGTH):
dataset.add_frame(
{
"action": rng.standard_normal(DUMMY_ACTION_DIM).astype(np.float32),
"observation.state": rng.standard_normal(DUMMY_STATE_DIM).astype(np.float32),
"observation.images.top": rng.integers(
0, 255, size=(IMAGE_SIZE, IMAGE_SIZE, 3), dtype=np.uint8
),
"task": f"task_{ep_idx}",
}
)
dataset.save_episode()
dataset.finalize()
return root
def make_train_config(root, output_dir, steps, ema_enable, ema_decay=None):
from lerobot.configs.default import DatasetConfig
from lerobot.policies.factory import make_policy_config
policy_config = make_policy_config(
"diffusion",
device="cpu",
push_to_hub=False,
n_obs_steps=2,
horizon=8,
n_action_steps=4,
drop_n_last_frames=0,
down_dims=(32, 64),
diffusion_step_embed_dim=32,
spatial_softmax_num_keypoints=8,
num_inference_steps=2,
pretrained_backbone_weights=None,
use_group_norm=True,
)
cfg = TrainPipelineConfig(
dataset=DatasetConfig(repo_id=DUMMY_REPO_ID, root=str(root)),
policy=policy_config,
output_dir=output_dir,
steps=steps,
batch_size=2,
num_workers=0,
seed=42,
log_freq=0,
env_eval_freq=0,
save_freq=2,
ema=EMAConfig(enable=ema_enable, decay=ema_decay),
)
cfg.optimizer = policy_config.get_optimizer_preset()
cfg.scheduler = policy_config.get_scheduler_preset()
# The config is built in-process, so skip the CLI-oriented validation.
cfg.validate = lambda: None
return cfg
def load_safetensors(path):
from safetensors.torch import load_file
return load_file(path)
def test_train_diffusion_with_ema_checkpoint_and_resume(tmp_path):
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])")
from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train
root = make_dummy_dataset(tmp_path)
output_dir = tmp_path / "_output"
cfg = make_train_config(root, output_dir, steps=4, ema_enable=True)
train(cfg)
checkpoint_dir = output_dir / "checkpoints" / "000004"
ema_state_path = checkpoint_dir / TRAINING_STATE_DIR / EMA_STATE_FILENAME
ema_model_dir = checkpoint_dir / f"{PRETRAINED_MODEL_DIR}_ema"
# The shadow state is saved for resume and tracks every optimizer step.
assert ema_state_path.exists()
ema_state = torch.load(ema_state_path, weights_only=True)
assert ema_state["optimization_step"] == 4
# A directly loadable EMA model is saved next to the live one, with different weights.
live_weights = load_safetensors(checkpoint_dir / PRETRAINED_MODEL_DIR / "model.safetensors")
ema_weights = load_safetensors(ema_model_dir / "model.safetensors")
assert set(live_weights) == set(ema_weights)
assert any(not torch.equal(live_weights[k], ema_weights[k]) for k in live_weights)
from lerobot.policies.diffusion.modeling_diffusion import DiffusionPolicy
policy = DiffusionPolicy.from_pretrained(str(ema_model_dir))
assert isinstance(policy, DiffusionPolicy)
# Resuming picks the shadow up where it left off instead of restarting it.
resume_cfg = make_train_config(root, output_dir, steps=6, ema_enable=True)
resume_cfg.resume = True
resume_cfg.checkpoint_path = checkpoint_dir
train(resume_cfg)
resumed_state = torch.load(
output_dir / "checkpoints" / "000006" / TRAINING_STATE_DIR / EMA_STATE_FILENAME,
weights_only=True,
)
assert resumed_state["optimization_step"] == 6
def test_train_with_constant_ema_decay(tmp_path):
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])")
from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train
root = make_dummy_dataset(tmp_path)
output_dir = tmp_path / "_output"
cfg = make_train_config(root, output_dir, steps=2, ema_enable=True, ema_decay=0.99)
train(cfg)
ema_state = torch.load(
output_dir / "checkpoints" / "000002" / TRAINING_STATE_DIR / EMA_STATE_FILENAME,
weights_only=True,
)
# The constant decay is implemented by pinning the schedule clamp to that value.
assert ema_state["decay"] == 0.99
assert ema_state["min_decay"] == 0.99
assert ema_state["optimization_step"] == 2
def test_train_with_ema_and_gradient_accumulation(tmp_path):
"""The shadow tracks optimizer steps, not micro-batches, under gradient accumulation."""
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])")
from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train
root = make_dummy_dataset(tmp_path)
output_dir = tmp_path / "_output"
cfg = make_train_config(root, output_dir, steps=4, ema_enable=True)
cfg.accelerator.gradient_accumulation.steps = 2
train(cfg)
ema_state = torch.load(
output_dir / "checkpoints" / "000004" / TRAINING_STATE_DIR / EMA_STATE_FILENAME,
weights_only=True,
)
# 4 micro-batches / 2 accumulation steps = 2 optimizer updates.
assert ema_state["optimization_step"] == 2
def test_train_without_ema_writes_no_ema_files(tmp_path):
pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])")
pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])")
from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train
root = make_dummy_dataset(tmp_path)
output_dir = tmp_path / "_output"
cfg = make_train_config(root, output_dir, steps=2, ema_enable=False)
train(cfg)
checkpoint_dir = output_dir / "checkpoints" / "000002"
assert (checkpoint_dir / PRETRAINED_MODEL_DIR / "model.safetensors").exists()
assert not (checkpoint_dir / TRAINING_STATE_DIR / EMA_STATE_FILENAME).exists()
assert not (checkpoint_dir / f"{PRETRAINED_MODEL_DIR}_ema").exists()
+34 -79
View File
@@ -21,8 +21,10 @@ This module tests multi-GPU training functionality with accelerate.
These tests are designed to run on machines with 2+ GPUs and are executed These tests are designed to run on machines with 2+ GPUs and are executed
in the nightly CI workflow. in the nightly CI workflow.
The tests automatically generate accelerate configs and launch training The tests launch `lerobot-train` through `accelerate launch` in a subprocess to properly test the
with subprocess to properly test the distributed training environment. distributed training environment. Accelerate is used as a plain launcher only: the topology comes
from `--parallelism.*` flags, never from an accelerate YAML config (see
`lerobot.distributed.factory.guard_against_env_interference`).
""" """
import os import os
@@ -58,73 +60,25 @@ def download_dataset(repo_id, episodes):
print(f"Dataset {repo_id} downloaded successfully") print(f"Dataset {repo_id} downloaded successfully")
def _write_multi_gpu_config(f, num_processes): def run_accelerate_training(config_args, num_processes=4):
f.write("compute_environment: LOCAL_MACHINE\n")
f.write("distributed_type: MULTI_GPU\n")
f.write("mixed_precision: 'no'\n")
f.write(f"num_processes: {num_processes}\n")
f.write("use_cpu: false\n")
f.write("gpu_ids: all\n")
f.write("downcast_bf16: 'no'\n")
f.write("machine_rank: 0\n")
f.write("main_training_function: main\n")
f.write("num_machines: 1\n")
f.write("rdzv_backend: static\n")
f.write("same_network: true\n")
def _write_fsdp_config(f, num_processes):
# FSDP1 with FULL_SHARD (ZeRO-3-equivalent) and FULL_STATE_DICT, matching
# docs/source/multi_gpu_training.mdx. ACT's repeated transformer blocks are the wrap units;
# fsdp_use_orig_params is required because LeRobot builds the optimizer before prepare().
f.write("compute_environment: LOCAL_MACHINE\n")
f.write("distributed_type: FSDP\n")
f.write("mixed_precision: 'no'\n")
f.write(f"num_processes: {num_processes}\n")
f.write("use_cpu: false\n")
f.write("gpu_ids: all\n")
f.write("machine_rank: 0\n")
f.write("main_training_function: main\n")
f.write("num_machines: 1\n")
f.write("rdzv_backend: static\n")
f.write("same_network: true\n")
f.write("fsdp_config:\n")
f.write(" fsdp_version: 1\n")
f.write(" fsdp_sharding_strategy: FULL_SHARD\n")
f.write(" fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP\n")
f.write(" fsdp_transformer_layer_cls_to_wrap: ACTEncoderLayer,ACTDecoderLayer\n")
f.write(" fsdp_use_orig_params: true\n")
f.write(" fsdp_state_dict_type: FULL_STATE_DICT\n")
def run_accelerate_training(config_args, num_processes=4, temp_dir=None, distributed_type="MULTI_GPU"):
""" """
Helper function to run training with accelerate launch. Helper function to run training with accelerate launch.
`accelerate launch` is used as a plain launcher (no `--config_file`): it only sets the
rendezvous env vars, and the layout — DDP by default, FSDP with `--parallelism.dp_shard` —
comes from `config_args`.
Args: Args:
config_args: List of config arguments to pass to lerobot_train.py config_args: List of config arguments to pass to lerobot_train.py
num_processes: Number of processes (GPUs) to use num_processes: Number of processes (GPUs) to use
temp_dir: Temporary directory for outputs
distributed_type: "MULTI_GPU" (DDP) or "FSDP" — selects the generated accelerate config.
Returns: Returns:
subprocess.CompletedProcess result subprocess.CompletedProcess result
""" """
config_path = Path(temp_dir) / "accelerate_config.yaml"
# Write YAML config
with open(config_path, "w") as f:
if distributed_type == "FSDP":
_write_fsdp_config(f, num_processes)
else:
_write_multi_gpu_config(f, num_processes)
cmd = [ cmd = [
"accelerate", "accelerate",
"launch", "launch",
"--config_file", f"--num_processes={num_processes}",
str(config_path),
"-m", "-m",
"lerobot.scripts.lerobot_train", "lerobot.scripts.lerobot_train",
] + config_args ] + config_args
@@ -173,7 +127,7 @@ class TestMultiGPUTraining:
"--num_workers=0", "--num_workers=0",
] ]
result = run_accelerate_training(config_args, num_processes=4, temp_dir=temp_dir) result = run_accelerate_training(config_args, num_processes=4)
# Check that training completed successfully # Check that training completed successfully
assert result.returncode == 0, ( assert result.returncode == 0, (
@@ -216,7 +170,7 @@ class TestMultiGPUTraining:
"--num_workers=0", "--num_workers=0",
] ]
result = run_accelerate_training(config_args, num_processes=2, temp_dir=temp_dir) result = run_accelerate_training(config_args, num_processes=2)
assert result.returncode == 0, ( assert result.returncode == 0, (
f"Training failed:\nSTDOUT:\n{result.stdout}\n\nSTDERR:\n{result.stderr}" f"Training failed:\nSTDOUT:\n{result.stdout}\n\nSTDERR:\n{result.stderr}"
@@ -246,11 +200,12 @@ class TestMultiGPUTraining:
def test_fsdp_optimizer_save_and_resume(self): def test_fsdp_optimizer_save_and_resume(self):
""" """
Test that FSDP saves the (gathered) optimizer state and can resume from it. Test that FSDP saves the sharded optimizer state and can resume from it.
Trains a few steps under FSDP, verifies the gathered optimizer state is written next to the Trains a few steps under FSDP2 (`--parallelism.dp_shard=2`), verifies the DCP optimizer
rest of the training state, then resumes from the checkpoint for more steps and checks it shards are written next to the rest of the training state, then resumes from the
completes without shape/key errors in the FSDP optimizer load path. checkpoint for more steps and checks it completes without shape/key errors in the
resharding optimizer load path.
""" """
# Pre-download dataset to avoid race conditions # Pre-download dataset to avoid race conditions
download_dataset("lerobot/pusht", episodes=[0]) download_dataset("lerobot/pusht", episodes=[0])
@@ -265,6 +220,7 @@ class TestMultiGPUTraining:
"--policy.device=cuda", "--policy.device=cuda",
"--policy.push_to_hub=false", "--policy.push_to_hub=false",
f"--output_dir={output_dir}", f"--output_dir={output_dir}",
"--parallelism.dp_shard=2",
"--batch_size=4", "--batch_size=4",
"--steps=10", "--steps=10",
"--env_eval_freq=-1", "--env_eval_freq=-1",
@@ -274,34 +230,33 @@ class TestMultiGPUTraining:
"--num_workers=0", "--num_workers=0",
] ]
result = run_accelerate_training( result = run_accelerate_training(config_args, num_processes=2)
config_args, num_processes=2, temp_dir=temp_dir, distributed_type="FSDP"
)
assert result.returncode == 0, ( assert result.returncode == 0, (
f"FSDP training failed:\nSTDOUT:\n{result.stdout}\n\nSTDERR:\n{result.stderr}" f"FSDP training failed:\nSTDOUT:\n{result.stdout}\n\nSTDERR:\n{result.stderr}"
) )
# The gathered optimizer state must be written under FSDP (proves the save collective ran), # Under sharding the optimizer state is written as DCP shards (proves the save
# in the same safetensors format as single-GPU training. # collective ran); the model artifact stays a gathered model.safetensors at the
training_state_dir = output_dir / "checkpoints" / "last" / "training_state" # default --checkpoint_format=safetensors.
optimizer_state = training_state_dir / "optimizer_state.safetensors" checkpoint_dir = output_dir / "checkpoints" / "last"
optimizer_param_groups = training_state_dir / "optimizer_param_groups.json" training_state_dir = checkpoint_dir / "training_state"
assert optimizer_state.exists(), f"FSDP optimizer state not saved in {training_state_dir}" optimizer_shards = training_state_dir / "optimizer_0"
assert optimizer_param_groups.exists(), ( assert optimizer_shards.is_dir(), f"FSDP optimizer shards not saved in {training_state_dir}"
f"FSDP optimizer param groups not saved in {training_state_dir}" assert any(optimizer_shards.iterdir()), f"FSDP optimizer shard dir is empty: {optimizer_shards}"
assert (checkpoint_dir / "pretrained_model" / "model.safetensors").exists(), (
f"Gathered model weights not saved in {checkpoint_dir}"
) )
# Resume from the checkpoint for more steps. A successful run proves load_fsdp_optimizer # Resume from the checkpoint for more steps. A successful run proves the DCP optimizer
# accepts the saved state and reshards it without shape/key errors. # load accepts the saved state and reshards it without shape/key errors. The topology
resume_config = output_dir / "checkpoints" / "last" / "pretrained_model" / "train_config.json" # is restored from train_config.json, so --parallelism.* is not repeated here.
resume_config = checkpoint_dir / "pretrained_model" / "train_config.json"
resume_args = [ resume_args = [
f"--config_path={resume_config}", f"--config_path={resume_config}",
"--resume=true", "--resume=true",
"--steps=20", "--steps=20",
] ]
resume_result = run_accelerate_training( resume_result = run_accelerate_training(resume_args, num_processes=2)
resume_args, num_processes=2, temp_dir=temp_dir, distributed_type="FSDP"
)
assert resume_result.returncode == 0, ( assert resume_result.returncode == 0, (
f"FSDP resume failed:\nSTDOUT:\n{resume_result.stdout}\n\nSTDERR:\n{resume_result.stderr}" f"FSDP resume failed:\nSTDOUT:\n{resume_result.stdout}\n\nSTDERR:\n{resume_result.stderr}"
) )
+111
View File
@@ -0,0 +1,111 @@
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import doctest
from lerobot.utils.doctest_utils import LeRobotDocTestParser, preprocess_string
# An example with expected output, formatted the way ruff's `docstring-code-format` leaves it: no blank
# line between the last output line and the closing fence. This is the exact shape that breaks stdlib.
FORMATTED_EXAMPLE = """Summary.
Example:
```python
>>> 1 + 1
2
```
"""
def test_stdlib_parser_swallows_the_closing_fence():
"""Guards the premise of the port: without the patch, the fence lands in the expected output.
Uses the base class rather than `doctest.DocTestParser`, which the root `conftest.py` has already
replaced with ours by the time this runs.
"""
stdlib_parser = LeRobotDocTestParser.__bases__[0]()
(example,) = (e for e in stdlib_parser.parse(FORMATTED_EXAMPLE) if isinstance(e, doctest.Example))
assert "```" in example.want
def test_parser_stops_at_the_closing_fence():
"""The whole reason `LeRobotDocTestParser` exists: `want` must be the output and nothing else."""
(example,) = (
e for e in LeRobotDocTestParser().parse(FORMATTED_EXAMPLE) if isinstance(e, doctest.Example)
)
assert example.source == "1 + 1\n"
assert example.want == "2\n"
def test_example_with_output_passes_end_to_end():
"""A formatted example with output should actually run green."""
runner = doctest.DocTestRunner()
test = LeRobotDocTestParser().get_doctest(FORMATTED_EXAMPLE, {}, "formatted", None, 0)
results = runner.run(test, out=lambda _: None)
assert results.failed == 0
assert results.attempted == 1
def test_noisy_calls_get_ignore_result():
string = """
```python
>>> ds = load_dataset("lerobot/pusht")
```
"""
assert "# doctest: +IGNORE_RESULT" in preprocess_string(string, False, False)
def test_ignore_result_is_not_added_twice():
string = """
```python
>>> ds = load_dataset("lerobot/pusht") # doctest: +IGNORE_RESULT
```
"""
assert preprocess_string(string, False, False).count("# doctest: +IGNORE_RESULT") == 1
def test_cuda_examples_are_dropped_when_requested():
string = """
```python
>>> model.to("cuda")
```
"""
assert preprocess_string(string, True, False) == ""
assert preprocess_string(string, False, False) != ""
def test_hardware_examples_are_dropped_when_requested():
"""Serial ports, connect calls and Hub downloads all need real resources."""
for source in [
'>>> robot = SO101Follower(SO101FollowerConfig(port="/dev/ttyACM0"))',
">>> robot.connect()",
'>>> policy = ACTPolicy.from_pretrained("lerobot/act")',
]:
string = f"""
```python
{source}
```
"""
assert preprocess_string(string, False, True) == "", source
assert preprocess_string(string, False, False) != "", source
def test_plain_examples_survive_both_skips():
string = """
```python
>>> 1 + 1
2
```
"""
assert preprocess_string(string, True, True) == string
+18 -46
View File
@@ -17,6 +17,7 @@
import pytest import pytest
import torch import torch
import lerobot.utils.logging_utils as logging_utils
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
@@ -25,19 +26,6 @@ def mock_metrics():
return {"loss": AverageMeter("loss", ":.3f"), "accuracy": AverageMeter("accuracy", ":.2f")} return {"loss": AverageMeter("loss", ":.3f"), "accuracy": AverageMeter("accuracy", ":.2f")}
class MockAccelerator:
def __init__(self, num_processes: int, reduce_fn=None):
self.num_processes = num_processes
self.device = torch.device("cpu")
self._reduce_fn = reduce_fn
def reduce(self, tensor, reduction="mean"):
# In single-process tests we just want a deterministic stand-in for accelerate's reduce.
if self._reduce_fn is not None:
return self._reduce_fn(tensor, reduction)
return tensor
def test_average_meter_initialization(): def test_average_meter_initialization():
meter = AverageMeter("loss", ":.2f") meter = AverageMeter("loss", ":.2f")
assert meter.name == "loss" assert meter.name == "loss"
@@ -96,14 +84,14 @@ def test_metrics_tracker_step(mock_metrics):
assert tracker.epochs == tracker.samples / 1000 assert tracker.epochs == tracker.samples / 1000
def test_metrics_tracker_initialization_with_accelerator(mock_metrics): def test_metrics_tracker_initialization_with_dp_world(mock_metrics):
tracker = MetricsTracker( tracker = MetricsTracker(
batch_size=32, batch_size=32,
num_frames=1000, num_frames=1000,
num_episodes=50, num_episodes=50,
metrics=mock_metrics, metrics=mock_metrics,
initial_step=10, initial_step=10,
accelerator=MockAccelerator(num_processes=2), dp_world_size=2,
) )
assert tracker.steps == 10 assert tracker.steps == 10
assert tracker.samples == 10 * 32 * 2 assert tracker.samples == 10 * 32 * 2
@@ -111,14 +99,14 @@ def test_metrics_tracker_initialization_with_accelerator(mock_metrics):
assert tracker.epochs == tracker.samples / 1000 assert tracker.epochs == tracker.samples / 1000
def test_metrics_tracker_step_with_accelerator(mock_metrics): def test_metrics_tracker_step_with_dp_world(mock_metrics):
tracker = MetricsTracker( tracker = MetricsTracker(
batch_size=32, batch_size=32,
num_frames=1000, num_frames=1000,
num_episodes=50, num_episodes=50,
metrics=mock_metrics, metrics=mock_metrics,
initial_step=5, initial_step=5,
accelerator=MockAccelerator(num_processes=2), dp_world_size=2,
) )
tracker.step() tracker.step()
assert tracker.steps == 6 assert tracker.steps == 6
@@ -178,54 +166,38 @@ def test_average_meter_reduction_stored():
assert meter.reduction == "max" assert meter.reduction == "max"
def test_metrics_tracker_reduce_across_ranks_no_accelerator(): def test_metrics_tracker_reduce_across_ranks_outside_distributed():
metrics = {"update_s": AverageMeter("update_s", reduction="max")} metrics = {"update_s": AverageMeter("update_s", reduction="max")}
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics=metrics) tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics=metrics)
tracker.update_s = 0.5 tracker.update_s = 0.5
tracker.reduce_across_ranks() # no-op without accelerator tracker.reduce_across_ranks() # no-op without an initialized process group
assert tracker.update_s.avg == 0.5 assert tracker.update_s.avg == 0.5
def test_metrics_tracker_reduce_across_ranks_single_process(): def test_metrics_tracker_reduce_across_ranks_invokes_all_reduce(monkeypatch):
metrics = {"update_s": AverageMeter("update_s", reduction="max")}
tracker = MetricsTracker(
batch_size=32,
num_frames=1000,
num_episodes=50,
metrics=metrics,
accelerator=MockAccelerator(num_processes=1),
)
tracker.update_s = 0.5
tracker.reduce_across_ranks() # no-op when world size is 1
assert tracker.update_s.avg == 0.5
def test_metrics_tracker_reduce_across_ranks_invokes_reduce():
captured = {} captured = {}
def fake_reduce(tensor, reduction): def fake_all_reduce(tensor, op):
captured["reduction"] = reduction captured["op"] = op
captured["values"] = tensor.clone() captured["values"] = tensor.clone()
# Pretend the slowest rank reported 0.9 instead of this rank's 0.4. # Pretend the slowest rank reported 0.9 instead of this rank's 0.4.
return torch.tensor([0.9], dtype=tensor.dtype, device=tensor.device) tensor.fill_(0.9)
monkeypatch.setattr(logging_utils.dist, "is_initialized", lambda: True)
monkeypatch.setattr(logging_utils.dist, "get_world_size", lambda: 4)
monkeypatch.setattr(logging_utils.dist, "all_reduce", fake_all_reduce)
metrics = { metrics = {
"loss": AverageMeter("loss"), # reduction="none" -> not touched "loss": AverageMeter("loss"), # reduction="none" -> not touched
"update_s": AverageMeter("update_s", reduction="max"), "update_s": AverageMeter("update_s", reduction="max"),
} }
tracker = MetricsTracker( tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics=metrics)
batch_size=32,
num_frames=1000,
num_episodes=50,
metrics=metrics,
accelerator=MockAccelerator(num_processes=4, reduce_fn=fake_reduce),
)
tracker.loss = 1.0 tracker.loss = 1.0
tracker.update_s = 0.4 tracker.update_s = 0.4
tracker.reduce_across_ranks() tracker.reduce_across_ranks()
assert captured["reduction"] == "max" assert captured["op"] == logging_utils.dist.ReduceOp.MAX
assert torch.allclose(captured["values"], torch.tensor([0.4])) assert torch.allclose(captured["values"], torch.tensor([0.4], device=captured["values"].device))
assert tracker.update_s.avg == pytest.approx(0.9) assert tracker.update_s.avg == pytest.approx(0.9)
# Metrics without a reduction stay untouched. # Metrics without a reduction stay untouched.
assert tracker.loss.avg == 1.0 assert tracker.loss.avg == 1.0
+25 -81
View File
@@ -15,24 +15,22 @@
# limitations under the License. # limitations under the License.
from pathlib import Path from pathlib import Path
from unittest.mock import MagicMock, Mock, patch from unittest.mock import MagicMock
import pytest import pytest
from lerobot.common.train_utils import ( from lerobot.common.train_utils import (
get_step_checkpoint_dir, get_step_checkpoint_dir,
get_step_identifier, get_step_identifier,
load_training_batch_size, load_training_metadata,
load_training_num_processes,
load_training_state,
load_training_step,
push_checkpoint_to_hub, push_checkpoint_to_hub,
save_checkpoint, save_training_metadata,
save_training_state, save_training_state,
save_training_step,
should_save_checkpoint, should_save_checkpoint,
update_last_checkpoint, update_last_checkpoint,
) )
from lerobot.configs.default import DatasetConfig
from lerobot.configs.train import TrainPipelineConfig
from lerobot.utils.constants import ( from lerobot.utils.constants import (
CHECKPOINTS_DIR, CHECKPOINTS_DIR,
LAST_CHECKPOINT_LINK, LAST_CHECKPOINT_LINK,
@@ -69,38 +67,23 @@ def test_get_step_checkpoint_dir():
assert step_dir == output_dir / CHECKPOINTS_DIR / "000005" assert step_dir == output_dir / CHECKPOINTS_DIR / "000005"
def test_save_load_training_step(tmp_path): def make_cfg(batch_size: int = 32) -> TrainPipelineConfig:
save_training_step(5000, tmp_path) cfg = TrainPipelineConfig(dataset=DatasetConfig(repo_id="lerobot/dummy"), batch_size=batch_size)
cfg.parallelism.resolve(1)
return cfg
def test_save_training_metadata_writes_the_step_file(tmp_path):
save_training_metadata(5000, tmp_path, make_cfg())
assert (tmp_path / TRAINING_STEP).is_file() assert (tmp_path / TRAINING_STEP).is_file()
def test_load_training_step(tmp_path): def test_save_training_state_records_topology(tmp_path, optimizer, scheduler):
step = 5000 save_training_state(tmp_path, 10, make_cfg(batch_size=32), optimizer, scheduler)
save_training_step(step, tmp_path) metadata = load_training_metadata(tmp_path / TRAINING_STATE_DIR)
loaded_step = load_training_step(tmp_path) assert metadata["step"] == 10
assert loaded_step == step assert metadata["dp_world_size"] == 1
assert metadata["batch_size"] == 32
def test_save_training_state_records_num_processes(tmp_path, optimizer, scheduler):
save_training_state(tmp_path, 10, optimizer, scheduler, num_processes=4)
assert load_training_num_processes(tmp_path) == 4
def test_load_training_num_processes_absent_returns_none(tmp_path, optimizer, scheduler):
# Checkpoints written before the world size was recorded must still load (back-compat).
save_training_state(tmp_path, 10, optimizer, scheduler)
assert load_training_num_processes(tmp_path) is None
def test_save_training_state_records_batch_size(tmp_path, optimizer, scheduler):
save_training_state(tmp_path, 10, optimizer, scheduler, batch_size=32)
assert load_training_batch_size(tmp_path) == 32
def test_load_training_batch_size_absent_returns_none(tmp_path, optimizer, scheduler):
# Checkpoints written before the batch size was recorded must still load (back-compat).
save_training_state(tmp_path, 10, optimizer, scheduler)
assert load_training_batch_size(tmp_path) is None
def test_update_last_checkpoint(tmp_path): def test_update_last_checkpoint(tmp_path):
@@ -112,32 +95,12 @@ def test_update_last_checkpoint(tmp_path):
assert last_checkpoint.resolve() == checkpoint assert last_checkpoint.resolve() == checkpoint
@patch("lerobot.common.train_utils.save_training_state") # save_checkpoint round-trips (all formats, real policies) live in
def test_save_checkpoint(mock_save_training_state, tmp_path, optimizer): # tests/common/test_checkpoint_save_resume.py.
policy = Mock()
cfg = Mock()
save_checkpoint(tmp_path, 10, cfg, policy, optimizer)
policy.save_pretrained.assert_called_once()
cfg.save_pretrained.assert_called_once()
mock_save_training_state.assert_called_once()
@patch("lerobot.common.train_utils.save_training_state") def test_save_training_state_layout(tmp_path, optimizer, scheduler):
def test_save_checkpoint_peft(mock_save_training_state, tmp_path, optimizer): save_training_state(tmp_path, 10, make_cfg(), optimizer, scheduler)
policy = Mock()
policy.config = Mock()
policy.config.save_pretrained = Mock()
cfg = Mock()
cfg.use_peft = True
save_checkpoint(tmp_path, 10, cfg, policy, optimizer)
policy.save_pretrained.assert_called_once()
cfg.save_pretrained.assert_called_once()
policy.config.save_pretrained.assert_called_once()
mock_save_training_state.assert_called_once()
def test_save_training_state(tmp_path, optimizer, scheduler):
save_training_state(tmp_path, 10, optimizer, scheduler)
assert (tmp_path / TRAINING_STATE_DIR).is_dir() assert (tmp_path / TRAINING_STATE_DIR).is_dir()
assert (tmp_path / TRAINING_STATE_DIR / TRAINING_STEP).is_file() assert (tmp_path / TRAINING_STATE_DIR / TRAINING_STEP).is_file()
assert (tmp_path / TRAINING_STATE_DIR / RNG_STATE).is_file() assert (tmp_path / TRAINING_STATE_DIR / RNG_STATE).is_file()
@@ -146,27 +109,8 @@ def test_save_training_state(tmp_path, optimizer, scheduler):
assert (tmp_path / TRAINING_STATE_DIR / SCHEDULER_STATE).is_file() assert (tmp_path / TRAINING_STATE_DIR / SCHEDULER_STATE).is_file()
def test_save_load_training_state(tmp_path, optimizer, scheduler): # The two-phase resume (resume_before_prepare / resume_after_prepare) is covered in
save_training_state(tmp_path, 10, optimizer, scheduler) # tests/common/test_checkpoint_save_resume.py with real policies and optimizer state.
loaded_step, loaded_optimizer, loaded_scheduler = load_training_state(tmp_path, optimizer, scheduler)
assert loaded_step == 10
assert loaded_optimizer is optimizer
assert loaded_scheduler is scheduler
def test_load_training_state_skip_optimizer(tmp_path, optimizer, scheduler):
# FSDP loads optimizer separately (after accelerator.prepare)
# load_training_state(load_optimizer=False) must restore step + scheduler but leave the
# optimizer untouched and never touch the on-disk optimizer state.
save_training_state(tmp_path, 10, optimizer, scheduler)
with patch("lerobot.common.train_utils.load_optimizer_state") as mock_load_optimizer_state:
loaded_step, loaded_optimizer, loaded_scheduler = load_training_state(
tmp_path, optimizer, scheduler, load_optimizer=False
)
mock_load_optimizer_state.assert_not_called()
assert loaded_step == 10
assert loaded_optimizer is optimizer
assert loaded_scheduler is scheduler
def test_push_checkpoint_to_hub_creates_repo_and_uploads(tmp_path, monkeypatch): def test_push_checkpoint_to_hub_creates_repo_and_uploads(tmp_path, monkeypatch):
+153
View File
@@ -0,0 +1,153 @@
# 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.
"""Check that every registered hardware config documents the fields users have to get right.
Modelled on `transformers/utils/check_config_docstrings.py`, which checks that every model config links a
checkpoint. LeRobot's equivalent question is the one every new user hits: which port is the device on, and
what happens on calibration. A config that leaves those undocumented sends people to the source.
Only fields the config actually declares are required a config without a `port` is not asked to document
one.
```bash
python utils/check_config_docstrings.py
```
"""
import inspect
import re
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from check_docstrings import _re_args, _re_parse_arg, find_indent, iter_objects_to_check # noqa: E402
# Fields whose semantics are not obvious from the name and that a user must set correctly on first run.
REQUIRED_FIELDS = ["port"]
# A config must say something about calibration if it participates in it at all.
CALIBRATION_PATTERN = re.compile(r"calibrat", re.IGNORECASE)
MODULES_TO_CHECK = ["lerobot.robots"]
# Configs that document their fields with `#` comments above each field, which doc-builder cannot see.
# Each entry is removed as that config's comments are converted to an `Args:` block.
OBJECTS_TO_IGNORE: set[str] = {
"BiOpenArmFollowerConfig",
"BiRebotB601FollowerConfig",
"BiSOFollowerConfig",
"EarthRoverMiniPlusConfig",
"HopeJrArmConfig",
"HopeJrHandConfig",
"KochFollowerConfig",
"LeKiwiConfig",
"OmxFollowerConfig",
"OpenArmFollowerConfig",
"Reachy2RobotConfig",
"RebotB601FollowerRobotConfig",
"SOFollowerRobotConfig",
}
def documented_args(obj: object) -> set[str]:
"""Return the argument names documented in an object's `Args:` block.
Args:
obj (`object`):
The class to inspect.
Returns:
`set[str]`: The documented argument names, empty if there is no `Args:` section.
"""
doc = getattr(obj, "__doc__", None)
if not doc:
return set()
lines = doc.split("\n")
idx = 0
while idx < len(lines) and _re_args.search(lines[idx]) is None:
idx += 1
if idx == len(lines):
return set()
indent = find_indent(lines[idx])
names = set()
idx += 1
while idx < len(lines) and (len(lines[idx].strip()) == 0 or find_indent(lines[idx]) > indent):
if find_indent(lines[idx]) == indent + 4:
match = _re_parse_arg.search(lines[idx])
if match is not None:
names.add(match.groups()[1])
idx += 1
return names
def check_config_docstrings() -> list[str]:
"""Check every registered config in `MODULES_TO_CHECK`.
Returns:
`list[str]`: One message per config that is missing a required field or calibration semantics.
"""
from lerobot.robots import RobotConfig
failures = []
for module_name in MODULES_TO_CHECK:
for obj in iter_objects_to_check(module_name):
if not inspect.isclass(obj) or not issubclass(obj, RobotConfig) or obj is RobotConfig:
continue
if inspect.isabstract(obj) or obj.__qualname__ in OBJECTS_TO_IGNORE:
continue
try:
fields = set(inspect.signature(obj).parameters)
except (TypeError, ValueError):
continue
doc = getattr(obj, "__doc__", "") or ""
documented = documented_args(obj)
name = f"{obj.__module__}.{obj.__qualname__}"
for field in REQUIRED_FIELDS:
if field in fields and field not in documented:
failures.append(f"{name}: does not document `{field}`")
if "calibration_dir" in fields and CALIBRATION_PATTERN.search(doc) is None:
failures.append(f"{name}: says nothing about calibration")
return failures
def main() -> int:
"""Run the check.
Returns:
`int`: `0` when every registered config is documented, `1` otherwise.
"""
failures = check_config_docstrings()
if failures:
print(
"The following robot configs are missing documentation a user needs on first run. See "
"docs/source/writing_docstrings.mdx:",
file=sys.stderr,
)
for failure in failures:
print(f"- {failure}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())
+566
View File
@@ -0,0 +1,566 @@
# 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.
"""Check that documented arguments match the real signature.
Adapted from the core of `transformers/utils/check_docstrings.py`. The parts of that file bound to
transformers internals the `@auto_docstring` decorator system, modular-file propagation, `ModelArgs`,
GitPython are deliberately not ported.
What this enforces, for every public object in `MODULES_TO_CHECK`:
- every parameter in the signature has an `Args:` entry, in signature order;
- no `Args:` entry names a parameter that does not exist;
- the `*optional*, defaults to `X`` clause matches the real default.
That last one is why the clause is not decorative. See docs/source/writing_docstrings.mdx.
Check, as CI does:
```bash
python utils/check_docstrings.py
```
Rewrite the `Args:` blocks to match the signatures, inserting `<fill_docstring>` placeholders for
parameters that are missing entirely:
```bash
python utils/check_docstrings.py --fix_and_overwrite
```
`MODULES_TO_CHECK` is the ratchet: add a module once its docstrings are converted.
"""
import argparse
import ast
import enum
import importlib
import inspect
import operator as op
import pkgutil
import re
import sys
from pathlib import Path
from typing import Any
PATH_TO_REPO = Path(__file__).resolve().parent.parent
PATH_TO_LEROBOT = PATH_TO_REPO / "src" / "lerobot"
# Modules whose public objects are checked. Add a module here once its docstrings follow the standard.
MODULES_TO_CHECK = [
"lerobot.robots",
]
# Objects that do not yet follow the standard, so the check can be green from day one. Removing an entry
# and running `--fix_and_overwrite` is how a module gets converted.
#
# Every entry below has a bare `Attributes:` section, which this checker reads as an argument section (the
# same aliasing doc-builder does) and therefore compares against the signature. They are converted in the
# docstring PR that follows this one, which empties this set.
OBJECTS_TO_IGNORE: set[str] = {
"ChannelFactoryInitialize",
"EarthRoverMiniPlus",
"EarthRoverMiniPlusConfig",
"EEBoundsAndSafety",
"EEReferenceAndDelta",
"ForwardKinematicsJointsToEEAction",
"ForwardKinematicsJointsToEEObservation",
"GripperVelocityToJoint",
"InverseKinematicsEEToJoints",
"Robot",
}
OPTIONAL_KEYWORD = "*optional*"
_re_args = re.compile(r"^\s*(Args?|Arguments?|Attributes?|Params?|Parameters?):\s*$")
_re_parse_arg = re.compile(r"^(\s*)(\S+)\s+\((.+)\)(?:\:|$)")
_re_parse_description = re.compile(r"\*optional\*, defaults to (.*)$")
MATH_OPERATORS = {
ast.Add: op.add,
ast.Sub: op.sub,
ast.Mult: op.mul,
ast.Div: op.truediv,
ast.Pow: op.pow,
ast.BitXor: op.xor,
ast.USub: op.neg,
}
def find_indent(line: str) -> int:
"""Return the number of spaces a line is indented by.
Args:
line (`str`):
The line to measure.
Returns:
`int`: The indentation width.
"""
search = re.search(r"^(\s*)(?:\S|$)", line)
return 0 if search is None else len(search.groups()[0])
def is_dataclass_factory_default(default: Any) -> bool:
"""Whether a signature default came from a dataclass `field(default_factory=...)`.
`inspect.signature` renders those as a `<factory>` sentinel, which must not be written into a
docstring as a literal default.
Args:
default (`Any`):
The default value taken from the signature.
Returns:
`bool`: `True` for the factory sentinel.
"""
return repr(default) == "<factory>"
def stringify_default(default: Any) -> str:
"""Render a default value the way a docstring should show it.
Args:
default (`Any`):
The default value to process.
Returns:
`str`: Numbers are left bare, everything else is wrapped in backticks.
"""
if isinstance(default, bool):
# Must precede the int check: a bool passes isinstance(x, int).
return f"`{default}`"
elif isinstance(default, enum.Enum):
# Must also precede the int check: an IntEnum passes isinstance(x, int).
return f"`{str(default)}`"
elif isinstance(default, int):
return str(default)
elif isinstance(default, float):
result = str(default)
return str(round(default, 2)) if len(result) > 6 else result
elif isinstance(default, str):
return str(default) if default.isnumeric() else f'`"{default}"`'
elif isinstance(default, type):
return f"`{default.__name__}`"
else:
return f"`{default}`"
def eval_node(node):
"""Evaluate one node of a arithmetic-only AST.
Args:
node (`ast.AST`):
The node to evaluate.
Returns:
`float | int | complex`: The node's value.
Raises:
TypeError: If the node is not a number or a supported arithmetic operation.
"""
if isinstance(node, ast.Constant) and type(node.value) in (int, float, complex):
return node.value
elif isinstance(node, ast.BinOp):
return MATH_OPERATORS[type(node.op)](eval_node(node.left), eval_node(node.right))
elif isinstance(node, ast.UnaryOp):
return MATH_OPERATORS[type(node.op)](eval_node(node.operand))
else:
raise TypeError(node)
def eval_math_expression(expression: str) -> float | int | None:
"""Safely evaluate an arithmetic expression found in a docstring.
Docstrings often document a default as an expression (`1 / 255` is the classic), which should be left
alone rather than replaced by its computed value.
Args:
expression (`str`):
The expression to evaluate.
Returns:
`float | int | None`: The value, or `None` if it is not a plain arithmetic expression.
"""
try:
return eval_node(ast.parse(expression, mode="eval").body)
except (TypeError, SyntaxError, KeyError, ZeroDivisionError):
return None
def replace_default_in_arg_description(description: str, default: Any) -> str:
"""Rewrite the `*optional*, defaults to X` clause of one argument description.
Args:
description (`str`):
The argument description from the docstring, without the name.
default (`Any`):
The real default from the signature, or `inspect._empty` if the argument is required.
Returns:
`str`: The description with its optional/default clause matching the signature.
"""
# Plenty of docstrings use `optional` or **optional** instead of *optional*.
description = description.replace("`optional`", OPTIONAL_KEYWORD)
description = description.replace("**optional**", OPTIONAL_KEYWORD)
if default is inspect._empty:
# Required: the description must not claim otherwise.
idx = description.find(OPTIONAL_KEYWORD)
if idx != -1:
description = description[:idx].rstrip().removesuffix(",").rstrip()
elif default is None or is_dataclass_factory_default(default):
# A `None` default is not spelled out, and a `default_factory` has no literal value to show.
idx = description.find(OPTIONAL_KEYWORD)
if idx == -1:
description = f"{description}, {OPTIONAL_KEYWORD}"
elif re.search(r"defaults to `?None`?", description) is not None:
description = description[: idx + len(OPTIONAL_KEYWORD)]
else:
str_default = None
documented_match = re.search("defaults to `?(.*?)(?:`|$)", description)
if isinstance(default, (int, float)) and documented_match is not None:
documented = documented_match.groups()[0]
if default == eval_math_expression(documented):
try:
# Directly convertible means it was a plain literal.
str_default = str(type(default)(documented))
except (TypeError, ValueError):
# Otherwise it was an expression; keep it as written.
str_default = f"`{documented}`"
if str_default is None:
str_default = stringify_default(default)
if OPTIONAL_KEYWORD not in description:
description = f"{description}, {OPTIONAL_KEYWORD}, defaults to {str_default}"
elif _re_parse_description.search(description) is None:
idx = description.find(OPTIONAL_KEYWORD)
description = f"{description[: idx + len(OPTIONAL_KEYWORD)]}, defaults to {str_default}"
else:
description = _re_parse_description.sub(f"*optional*, defaults to {str_default}", description)
return description
def get_default_description(arg: inspect.Parameter) -> str:
"""Build the parenthesised type-and-default part for an undocumented parameter.
Args:
arg (`inspect.Parameter`):
The parameter to describe.
Returns:
`str`: Something like ``` `int`, *optional*, defaults to 3 ```.
"""
if arg.annotation is inspect._empty:
arg_type = "<fill_type>"
elif hasattr(arg.annotation, "__name__"):
arg_type = arg.annotation.__name__
else:
arg_type = str(arg.annotation)
if arg.default is inspect._empty:
return f"`{arg_type}`"
elif arg.default is None or is_dataclass_factory_default(arg.default):
return f"`{arg_type}`, {OPTIONAL_KEYWORD}"
else:
return f"`{arg_type}`, {OPTIONAL_KEYWORD}, defaults to {stringify_default(arg.default)}"
def find_source_file(obj: Any) -> Path:
"""Locate the file an object is defined in.
Args:
obj (`Any`):
The object to locate.
Returns:
`Path`: The source file.
"""
obj_file = PATH_TO_LEROBOT
for part in obj.__module__.split(".")[1:]:
obj_file = obj_file / part
return obj_file.with_suffix(".py")
def match_docstring_with_signature(obj: Any) -> tuple[str, str] | None:
"""Compare an object's documented arguments against its signature.
Dataclasses need no special handling: `inspect.signature` resolves the generated `__init__`, inherited
fields included, which is exactly the set a reader sees on the rendered page.
Args:
obj (`Any`):
The class or function to check.
Returns:
`tuple[str, str] | None`: The current `Args:` block and the one matching the signature, or `None`
when there is nothing to compare no docstring, no documented arguments, or an unsupported
signature.
"""
if not getattr(obj, "__doc__", None):
return None
try:
source, _ = inspect.getsourcelines(obj)
except (OSError, TypeError):
source = []
idx = 0
while idx < len(source) and '"""' not in source[idx]:
idx += 1
ignore_order = False
if idx < len(source) and idx > 0:
line_before_docstring = source[idx - 1]
if re.search(r"^\s*#\s*no-format\s*$", line_before_docstring):
return None
elif re.search(r"^\s*#\s*ignore-order\s*$", line_before_docstring):
ignore_order = True
try:
signature = inspect.signature(obj).parameters
except (ValueError, TypeError):
return None
obj_doc_lines = obj.__doc__.split("\n")
idx = 0
while idx < len(obj_doc_lines) and _re_args.search(obj_doc_lines[idx]) is None:
idx += 1
if idx == len(obj_doc_lines):
# No arguments documented; coverage is interrogate's job, not this check's.
return None
if "kwargs" in signature and signature["kwargs"].annotation != inspect._empty:
# Typed **kwargs are not introspectable in a useful way here.
return None
indent = find_indent(obj_doc_lines[idx])
arguments: dict[str, Any] = {}
current_arg = None
idx += 1
start_idx = idx
# Consume until a non-empty line returns to the section's own indent, or the docstring ends.
while idx < len(obj_doc_lines) and (
len(obj_doc_lines[idx].strip()) == 0 or find_indent(obj_doc_lines[idx]) > indent
):
if find_indent(obj_doc_lines[idx]) == indent + 4:
re_search_arg = _re_parse_arg.search(obj_doc_lines[idx])
if re_search_arg is not None:
_, name, description = re_search_arg.groups()
current_arg = name
if name in signature:
default = signature[name].default
if signature[name].kind is inspect._ParameterKind.VAR_KEYWORD:
default = None
new_description = replace_default_in_arg_description(description, default)
else:
new_description = description
arguments[current_arg] = [
_re_parse_arg.sub(rf"\1\2 ({new_description}):", obj_doc_lines[idx])
]
elif current_arg is not None:
arguments[current_arg].append(obj_doc_lines[idx])
idx += 1
# Walk back over the trailing blank lines we consumed.
idx -= 1
if current_arg:
while len(obj_doc_lines[idx].strip()) == 0:
arguments[current_arg] = arguments[current_arg][:-1]
idx -= 1
idx += 1
old_doc_arg = "\n".join(obj_doc_lines[start_idx:idx])
old_arguments = list(arguments.keys())
arguments = {name: "\n".join(doc) for name, doc in arguments.items()}
for name in set(signature.keys()) - set(arguments.keys()):
arg = signature[name]
# Private parameters and *args/**kwargs are only documented if the author chose to.
if name.startswith("_") or arg.kind in [
inspect._ParameterKind.VAR_KEYWORD,
inspect._ParameterKind.VAR_POSITIONAL,
]:
arguments[name] = ""
else:
arguments[name] = (
" " * (indent + 4) + f"{name} ({get_default_description(arg)}): <fill_docstring>"
)
if ignore_order:
new_param_docs = [arguments[name] for name in old_arguments if name in signature]
missing = set(signature.keys()) - set(old_arguments)
new_param_docs.extend([arguments[name] for name in missing if len(arguments[name]) > 0])
else:
new_param_docs = [arguments[name] for name in signature if len(arguments[name]) > 0]
return old_doc_arg, "\n".join(new_param_docs)
def fix_docstring(obj: Any, old_doc_args: str, new_doc_args: str) -> None:
"""Rewrite an object's `Args:` block in its source file.
Args:
obj (`Any`):
The object whose docstring is being fixed.
old_doc_args (`str`):
The current `Args:` block, as returned by [`match_docstring_with_signature`].
new_doc_args (`str`):
The replacement block, as returned by [`match_docstring_with_signature`].
Raises:
ValueError: If the block found in the source does not match the one parsed from `__doc__`, which
means the boundaries were identified wrongly and rewriting would corrupt the file.
"""
source, line_number = inspect.getsourcelines(obj)
idx = 0
while idx < len(source) and _re_args.search(source[idx]) is None:
idx += 1
if idx == len(source):
# Inherited docstring: do not rewrite it on the child.
return
indent = find_indent(source[idx])
idx += 1
start_idx = idx
while idx < len(source) and (len(source[idx].strip()) == 0 or find_indent(source[idx]) > indent):
idx += 1
idx -= 1
while len(source[idx].strip()) == 0:
idx -= 1
idx += 1
# `old_doc_args` comes from `__doc__`, whose indentation differs from the raw source lines.
source_args_as_str = "".join(source[start_idx:idx])
if inspect.cleandoc(source_args_as_str) != inspect.cleandoc(old_doc_args):
raise ValueError(
f"Cannot fix the docstring of {obj.__name__} in {find_source_file(obj)}: the argument section "
f"in the source does not match the one parsed from __doc__, so the block boundaries are "
f"wrong and rewriting it would corrupt the file.\n\n"
f"Parsed:\n{old_doc_args!r}\n\nFound in source:\n{source_args_as_str.rstrip()!r}\n"
)
obj_file = find_source_file(obj)
lines = obj_file.read_text(encoding="utf-8").split("\n")
# `new_doc_args` is built from `__doc__`, and Python keeps every line after the first at its exact
# source indentation, so the block is already correctly indented for the file. transformers re-indents
# here because its docstrings are often assembled by decorators and no longer match the source.
lines = lines[: line_number + start_idx - 1] + [new_doc_args] + lines[line_number + idx - 1 :]
print(f"Fixing the docstring of {obj.__name__} in {obj_file}.")
obj_file.write_text("\n".join(lines), encoding="utf-8")
def iter_objects_to_check(module_name: str):
"""Yield the public classes and functions defined in a package.
Args:
module_name (`str`):
An importable package name, e.g. `"lerobot.robots"`.
Yields:
`Any`: Each public class or function whose `__module__` is inside the package, deduplicated so
that aliases (`SO101Follower = SOFollower`) are visited once.
"""
package = importlib.import_module(module_name)
module_names = [module_name]
if hasattr(package, "__path__"):
module_names += [
name for _, name, _ in pkgutil.walk_packages(package.__path__, prefix=f"{module_name}.")
]
seen = set()
for name in module_names:
try:
module = importlib.import_module(name)
except Exception as error: # An optional extra is missing; not this check's problem.
print(f"Skipping {name}: {type(error).__name__}: {error}", file=sys.stderr)
continue
for attr_name, obj in vars(module).items():
if attr_name.startswith("_") or not (inspect.isclass(obj) or inspect.isfunction(obj)):
continue
if not getattr(obj, "__module__", "").startswith(module_name):
continue
key = f"{obj.__module__}.{obj.__qualname__}"
if key in seen or obj.__qualname__ in OBJECTS_TO_IGNORE or key in OBJECTS_TO_IGNORE:
continue
seen.add(key)
yield obj
def check_docstrings(overwrite: bool = False) -> list[str]:
"""Check every object in `MODULES_TO_CHECK`.
Args:
overwrite (`bool`, *optional*, defaults to `False`):
Whether to rewrite mismatched `Args:` blocks in place.
Returns:
`list[str]`: The names of objects whose documented arguments do not match their signature. Empty
when everything is consistent.
"""
failures = []
hard_failures = []
for module_name in MODULES_TO_CHECK:
for obj in iter_objects_to_check(module_name):
try:
result = match_docstring_with_signature(obj)
except Exception as error:
hard_failures.append(f"{obj.__qualname__}: {type(error).__name__}: {error}")
continue
if result is None:
continue
old_doc, new_doc = result
if old_doc == new_doc:
continue
if overwrite:
fix_docstring(obj, old_doc, new_doc)
else:
failures.append(f"{obj.__module__}.{obj.__qualname__}")
if hard_failures:
print("The following objects could not be processed:", file=sys.stderr)
for failure in hard_failures:
print(f"- {failure}", file=sys.stderr)
return failures
def main() -> int:
"""Run the check.
Returns:
`int`: `0` when every documented argument matches its signature, `1` otherwise.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--fix_and_overwrite", action="store_true", help="Whether to fix inconsistencies.")
args = parser.parse_args()
failures = check_docstrings(overwrite=args.fix_and_overwrite)
if failures:
print(
"The docstrings of the following objects do not match their signature. Run "
"`make fix-docstrings` to rewrite them, then fill in any `<fill_docstring>` placeholders:",
file=sys.stderr,
)
for failure in failures:
print(f"- {failure}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())
+109
View File
@@ -0,0 +1,109 @@
# 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.
"""Keep the doctest list honest: every path exists, and the file stays sorted.
Adapted from `transformers/utils/check_doctest_list.py`. It is agnostic to whether the list is an allowlist
(what we have now) or a denylist (where transformers ended up), so it survives that inversion unchanged.
Check, as CI does:
```bash
python utils/check_doctest_list.py
```
Sort in place:
```bash
python utils/check_doctest_list.py --fix_and_overwrite
```
"""
import argparse
import sys
from pathlib import Path
REPO_PATH = Path(__file__).resolve().parent.parent
DOCTEST_FILE_PATHS = ["documentation_tests.txt"]
def split_header(lines: list[str]) -> tuple[list[str], list[str]]:
"""Split a list file into its leading comment header and its path entries.
Args:
lines (`list[str]`):
The file's lines, without trailing newlines.
Returns:
`tuple[list[str], list[str]]`: The leading comment/blank lines, and the remaining lines.
"""
for i, line in enumerate(lines):
if line.strip() and not line.lstrip().startswith("#"):
return lines[:i], lines[i:]
return lines, []
def clean_doctest_list(doctest_file: Path, overwrite: bool = False) -> None:
"""Check, and optionally fix, one doctest list file.
Args:
doctest_file (`Path`):
The list file to check or clean.
overwrite (`bool`, *optional*, defaults to `False`):
Whether to fix problems in place. When `False`, raises instead.
Raises:
ValueError: If the file lists a path that does not exist, or is not alphabetically sorted and
`overwrite` is `False`.
"""
lines = doctest_file.read_text(encoding="utf-8").splitlines()
header, entries = split_header(lines)
paths = [line.strip().split(" ")[0] for line in entries if line.strip()]
non_existent = [p for p in paths if not (REPO_PATH / p).exists()]
if non_existent:
listed = "\n".join(f"- {p}" for p in non_existent)
raise ValueError(f"`{doctest_file.name}` contains non-existent paths:\n{listed}")
if paths != sorted(paths):
if not overwrite:
raise ValueError(
f"Files in `{doctest_file.name}` are not in alphabetical order, run "
"`make fix-docstrings` to fix this automatically."
)
doctest_file.write_text("\n".join(header + sorted(paths)) + "\n", encoding="utf-8")
def main() -> int:
"""Run the check over every doctest list file.
Returns:
`int`: A process exit code `0` when every file is clean, `1` otherwise.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--fix_and_overwrite", action="store_true", help="Whether to fix inconsistencies.")
args = parser.parse_args()
failed = False
for name in DOCTEST_FILE_PATHS:
try:
clean_doctest_list(REPO_PATH / "utils" / name, args.fix_and_overwrite)
except ValueError as error:
print(error, file=sys.stderr)
failed = True
return 1 if failed else 0
if __name__ == "__main__":
raise SystemExit(main())
+12
View File
@@ -0,0 +1,12 @@
# Files whose docstring examples are executed by `make doctest`.
#
# This is an ALLOWLIST: only the paths below are collected. transformers started the same way and has since
# inverted to a denylist (`utils/not_doctested.txt`), which is the better end state — it makes a new file
# tested by default. LeRobot cannot start there: at the time of writing, public docstring coverage is under
# 50% and only a handful of files carry any example at all, so a denylist would need hundreds of entries on
# day one and would say nothing about what is actually verified.
#
# Invert once coverage is high enough that the exclusions are the short list.
# `utils/check_doctest_list.py` does not care which way round it is.
#
# Keep alphabetically sorted: `make check-doctest-list` enforces it, `make fix-docstrings` sorts it.