Files
lerobot/tests/configs/test_train_config_distributed.py
T
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

119 lines
4.8 KiB
Python

#!/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()