mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
ef88d4e52b
* 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.
119 lines
4.8 KiB
Python
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()
|