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.
This commit is contained in:
Haoming Song
2026-08-06 19:16:41 +08:00
committed by GitHub
parent 64b23178d5
commit ef88d4e52b
46 changed files with 4898 additions and 1132 deletions
@@ -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."
)