Transformers 分布式训练实战:使用 Fully Sharded Data Parallel (FSDP) 在 Accelerate 与 Trainer 中训练超大模型
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
本文是 🤗 Transformers 官方文档中 FSDP(Fully Sharded Data Parallel,完全分片数据并行)指南的深度实战解读。它面向需要在多卡 GPU 或 TPU 上训练超出单卡显存容量的大模型(如数十亿参数的 LLM)的开发者,完整覆盖从accelerate config交互式配置、五种分片策略选型、CPU 卸载与自动包装策略,到检查点保存/恢复、PyTorch/XLA TPU 支持及accelerate launch启动训练的全流程。读完本文,你将掌握基于 Trainer 快速搭建 FSDP 训练环境的具体方法,并能结合仓库源码理解每个配置项背后的真实语义。
FSDP 是什么:把"整卡复制"变成"按卡分片"
Fully Sharded Data Parallel (FSDP) 是一种数据并行训练方式,其核心思想是将模型的参数(parameters)、梯度(gradients)和优化器状态(optimizer states)按照可用的 GPU 数量(即 worker 或rank)进行切分。与传统的 DistributedDataParallel (DDP) 在每个 GPU 上都维护一份完整模型副本不同,FSDP 让每张 GPU 只持有模型的一部分,从而显著降低单卡显存占用,让开发者能够用较少的 GPU 训练远大于单卡容量的模型。
从源码与官方英文指南(docs/source/en/fsdp.md)中可以进一步确认其工作原理:在每次前向计算之前,每个 GPU 会通过all-gather通信从所有分片中聚合出当前层所需的完整参数;前向结束后立即释放这些参数(即"重新分片"),以省出下一层使用的显存;反向传播阶段则通过reduce-scatter将梯度聚合回各自的梯度分片。也就是说,FSDP 用更高的通信开销换取了更低的显存峰值——这是它与 DDP 最本质的权衡,因此官方建议:当模型或优化器状态无法装进单卡显存时,才选择 FSDP。
FSDP 与 Accelerate 深度集成(Accelerate 是简化分布式训练环境管理的库),因此可以直接在Trainer中开箱使用。
环境准备
开始之前,请确认以下依赖已就绪:
- 已安装Accelerate;
- 已安装PyTorch 2.1.0 及以上版本(英文版指南 docs/source/en/fsdp.md 面向 FSDP2 进一步要求较新的 PyTorch,仓库源码 src/transformers/distributed/fsdp.py 中 FSDP2 的
fully_shardAPI 在 torch ≥ 2.6 时导入)。
pip install accelerate第一步:用accelerate config生成训练环境配置
要开始配置 FSDP 训练,首先运行交互式命令:
accelerate configaccelerate config会依次弹出若干问题,用于生成训练环境的配置文件;Accelerate 会依据该文件中选定的训练选项自动搭建正确的分布式训练环境。在交互过程中会出现多个 FSDP 相关选项,下面逐一讲解其中最关键的几项。其余可用选项可查阅TrainingArguments.fsdp_config参数的完整说明(training_args.py 源码)。
核心配置项详解
1. 分片策略(Sharding Strategy)
FSDP 提供多种分片粒度,在accelerate config中通过数字选择,对应fsdp_sharding_strategy标志:
| 选项 | 选择值 | 分片范围 | 语义 |
|---|---|---|---|
FULL_SHARD | 1 | 参数 + 梯度 + 优化器状态 | 三者全部在 worker 间分片,显存最省(对应 ZeRO-3) |
SHARD_GRAD_OP | 2 | 梯度 + 优化器状态 | 参数保持完整副本,显存省一半(对应 ZeRO-2) |
NO_SHARD | 3 | 不分片 | 与 DDP 行为一致 |
HYBRID_SHARD | 4 | 节点内分片参数/梯度/优化器 | 每个节点保留完整副本,节点内分片,跨节点复制 |
HYBRID_SHARD_ZERO2 | 5 | 节点内分片梯度/优化器 | 混合分片的 ZeRO-2 变体 |
在旧版 API(FSDPOption枚举,见 src/transformers/trainer_utils.py)中这些策略对应full_shard、shard_grad_op、no_shard、hybrid_shard、hybrid_shard_zero2字符串。而当前仓库已默认升级到FSDP2(fsdp_config["version"]默认值为2),分片语义由 training_args.py 中的_apply_legacy_fsdp_to_config自动转换:旧的"full_shard"映射为reshard_after_forward: true,旧的"shard_grad_op"映射为reshard_after_forward: false。
在 FSDP2 中,reshard_after_forward是控制"显存 ↔ 吞吐"取舍的关键开关:
true:前向结束后立即重新分片参数,进一步省显存(默认);false:前向与反向之间保持参数处于聚合状态,避免再次 all-gather,但峰值显存更高。
2. CPU 卸载(CPU Offload)
当显存仍然吃紧时,可以把暂时不用的参数与梯度卸载到 CPU 内存,从而加载即使 FSDP 分片后也无法塞进 GPU 的更大模型。在accelerate config中将其置为:
fsdp_offload_params: true对应到TrainingArguments.fsdp_config的键为cpu_offload(默认false)。仓库源码 src/transformers/distributed/fsdp.py 中,当fsdp_cpu_offload打开时会为fully_shard注入CPUOffloadPolicy(),实现参数与梯度的 CPU 卸载。
3. 包装策略(Wrapping Policy)
FSDP 通过逐层包装(wrapping)模型网络的每一层来生效。包装通常以嵌套方式应用:每层前向通过后即丢弃其完整权重,为下一层腾出内存。实现这一机制最省事的方式是自动包装(auto wrap),无需修改任何模型代码:
- 选择
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP,并配合fsdp_transformer_layer_cls_to_wrap指定需要包装的 Transformer 层类名(例如BertLayer)。每个包装单元各自管理自己的 all-gather/reduce-scatter 操作,前向过程中只聚合当前单元的参数,前一个单元的参数随即被释放; - 或者选择基于尺寸的包装策略
fsdp_wrap_policy: SIZE_BASED_WRAP,并设置min_num_param为期望的阈值——当某个模块的参数数量超过该阈值时,FSDP 便对该模块应用包装。
英文版指南与源码给出了两条重要的工程建议:
- 不要只包装顶层模型——那将得不到任何显存收益(整个模型变成一个 FSDP 单元);也不要包装每一个
Linear层——单元间通信会变得极其昂贵; transformer_layer_cls_to_wrap通常可以留空,因为当auto_wrap_policy为TRANSFORMER_BASED_WRAP时,FSDP 会回退读取模型定义中的_no_split_modules,覆盖绝大多数 Transformers 模型(见 training_args.py 中_process_fsdp_args的实现)。
NO_WRAP策略则完全不包装,不推荐用于追求显存收益的场景。源码中FSDPOption还保留了一个AUTO_WRAP选项(旧版 accelerate 的自动包装开关),现已统一由auto_wrap_policy管理。
4. 更多fsdp_config键:来自源码的完整参数表
除上述交互式配置外,fsdp_config还支持直接以 JSON 文件或 dict 形式传入更多参数。结合 training_args.py 第 679–713 行的文档与_process_fsdp_args(约第 2734 行起)的解析逻辑,整理如下:
| 键 | 默认值 | 说明 |
|---|---|---|
version | 2 | FSDP 版本:2为 FSDP2(默认),1为旧版 FSDP1(已弃用,将在 v5.20 移除) |
reshard_after_forward | true | 前向后是否重新分片;FSDP2 下控制显存/吞吐取舍 |
cpu_offload | false | 将不用的参数与梯度卸载到 CPU |
activation_checkpointing | false | 反向时重算激活而非保存,节省显存 |
cpu_ram_efficient_loading | false | 仅在 rank 0 上从磁盘加载检查点,其他进程以空权重启动并靠广播接收权重,避免多进程同时把大模型读入 CPU 内存 |
state_dict_type | "FULL_STATE_DICT" | 检查点格式:单份兼容 Transformers 的完整权重,或"SHARDED_STATE_DICT"(每个 rank 一份,大模型更快) |
auto_wrap_policy | "TRANSFORMER_BASED_WRAP" | 可选TRANSFORMER_BASED_WRAP/SIZE_BASED_WRAP/NO_WRAP |
transformer_layer_cls_to_wrap | 无 | 要包装的层类名(区分大小写),如LlamaDecoderLayer;通常可留空 |
min_num_params | 0 | 尺寸包装的每模块最少参数数(配合SIZE_BASED_WRAP) |
forward_prefetch/backward_prefetch | false/"NO_PREFETCH" | 前向/反向的预取策略(FSDP1),如BACKWARD_PRE |
use_orig_params | true | 是否保留原始参数对象(FSDP1),对参数名可见性与某些优化器兼容性有影响 |
sync_module_states | true | 是否同步各 rank 的模块初始状态 |
limit_all_gathers | 无 | 限制并发 all-gather 数量,降低通信峰值 |
xla/xla_fsdp_v2/xla_fsdp_grad_ckpt/xla_fsdp_settings | false | TPU(PyTorch/XLA)相关,见下文 TPU 一节 |
值得注意的源码细节:_process_fsdp_args会先剔除fsdp_前缀(兼容 accelerate 风格键名),再对transformer_layer_cls_to_wrap做字符串 → 列表归一化;若xla开启,xla_fsdp_settings中的compute_dtype、buffer_dtype会被转换为真实的torch数据类型对象。
检查点(Checkpointing):中间断点与最终权重分开处理
训练过程中的中间检查点必须以fsdp_state_dict_type: SHARDED_STATE_DICT(每个 rank 各存一份分片状态)保存。原因在于:开启 CPU 卸载时,rank 0 上汇聚完整状态字典会非常耗时,且广播期间可能无限期等待,最终触发NCCL Timeout错误。
使用 Accelerate 的~accelerate.Accelerator.load_state方法可以从分片状态字典恢复训练:
# 恢复路径中隐含的检查点 accelerator.load_state("ckpt")但训练结束后必须保存一份完整状态字典(FULL_STATE_DICT),因为分片状态字典只能被 FSDP 自己加载,无法被普通流程读取:
if trainer.is_fsdp_enabled: trainer.accelerator.state.fsdp_plugin.set_state_dict_type("FULL_STATE_DICT") trainer.save_model(script_args.output_dir)源码层面的佐证(src/transformers/trainer.py):
- 恢复检查点时,Trainer 会检测目录中是否含有 FSDP 专属文件(常量
FSDP_MODEL_NAME = "pytorch_model_fsdp",约第 250 行),并通过save_fsdp_model/save_fsdp_optimizer/load_fsdp_model/load_fsdp_optimizer(来自 Accelerate)读写分片检查点; - 若目录是 FSDP 检查点但当前未启用 FSDP,Trainer 会直接报错(约第 3446 行);
- 两个与检查点相关的互斥约束:
save_only_model与SHARDED_STATE_DICT不兼容(约第 873–878 行);save_only_model与load_best_model_at_end在 FSDP/DeepSpeed 下同时使用会报错(约第 855–862 行); - 仓库测试 tests/trainer/distributed/test_trainer_distributed_fsdp.py 中
resume_params覆盖了FULL_STATE_DICT(仅 FSDP1)与SHARDED_STATE_DICT(FSDP1 与 FSDP2)的恢复组合,可直接作为回归验证参考。
在 TPU 上使用 FSDP(PyTorch/XLA)
PyTorch XLA 支持在 TPU 上进行 FSDP 训练。通过修改accelerate config生成的 FSDP 配置文件即可启用——除上文的分片策略与包装选项外,在文件中追加以下参数:
xla: True # 必须设为 True 以启用 PyTorch/XLA xla_fsdp_settings: # XLA 专属的 FSDP 参数 xla_fsdp_grad_ckpt: True # 使用梯度检查点(gradient checkpointing)xla_fsdp_settings用于配置额外的 XLA 专属 FSDP 参数(完整选项见 PyTorch/XLA 的xla_fully_sharded_data_parallel.py源码,仓库侧在 training_args.py 中负责解析与类型转换)。
仓库中 XLA 路径的实现细节:Trainer 在 trainer.py 中通过wrap_model_xla_fsdp包装模型(约第 2555 行),并支持更新的xla_fsdp_v2模式——该模式下会调用xs.set_global_mesh建立形如("fsdp", "tensor")的二维设备网格(约第 620–626 行);TPU 检查点则通过save_tpu_checkpoint处理。
启动训练:完整配置示例与 launch 命令
一份典型的 FSDP 配置文件示例如下(由accelerate config生成,可按需手改):
compute_environment: LOCAL_MACHINE debug: false distributed_type: FSDP downcast_bf16: 'no' fsdp_config: fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP fsdp_backward_prefetch_policy: BACKWARD_PRE fsdp_cpu_ram_efficient_loading: true fsdp_forward_prefetch: false fsdp_offload_params: true fsdp_sharding_strategy: 1 fsdp_state_dict_type: SHARDED_STATE_DICT fsdp_sync_module_states: true fsdp_transformer_layer_cls_to_wrap: BertLayer fsdp_use_orig_params: true machine_rank: 0 main_training_function: main mixed_precision: bf16 num_machines: 1 num_processes: 2 rdzv_backend: static same_network: true tpu_env: [] tpu_use_cluster: false tpu_use_sudo: false use_cpu: false该配置的要点:distributed_type: FSDP声明分布式类型;fsdp_sharding_strategy: 1即 FULL_SHARD 全分片;开启 CPU 卸载(fsdp_offload_params: true)与 BF16 混合精度;按BertLayer做 Transformer 层自动包装;中间断点采用分片状态字典;num_processes: 2表示使用 2 个进程(GPU)。真实仓库测试的最小 FSDP2 配置更为精简——见 tests/trainer/distributed/accelerate_configs/fsdp2.yaml,仅需distributed_type: FSDP、fsdp_config.fsdp_version: 2与num_processes: 2三行,其余模型相关设置通过 launch 参数传入。
随后使用accelerate launch启动训练脚本(Trainer 脚本无需改动),它会自动读取之前由accelerate config生成的配置文件:
accelerate launch my-trainer-script.py也可以不用交互式配置,直接在命令行指定 FSDP 选项与配置文件路径:
accelerate launch --fsdp="full shard" --fsdp_config="path/to/fsdp_config/" my-trainer-script.py两种方式等价:前者读取缓存中的default_config.yaml,后者显式传入。使用 FSDP2 时,也可以完全绕过 accelerate 配置文件,在TrainingArguments中直接开启(英文版指南 docs/source/en/fsdp.md):
from transformers import TrainingArguments TrainingArguments( ..., fsdp=True, fsdp_config="path/to/fsdp.json", # {"version": 2, "reshard_after_forward": true, ...} )fsdp_config既接受 JSON 文件路径,也接受已加载的 dict;源码中通过_process_fsdp_args(training_args.py 第 2734 行起)统一归一化后交给 Accelerate 的FullyShardedDataParallelPlugin使用(trainer.py 约第 817–841 行)。
Trainer 与 FSDP 的集成细节
将 FSDP 集成到Trainer后,以下几个源码级行为值得了解(均位于 src/transformers/trainer.py):
- 冲突检测:当 FSDP 配置中开启了
activation_checkpointing而TrainingArguments同时开启了gradient_checkpointing时,Trainer 会直接抛出ValueError(约第 844–850 行)。原因在于 FSDP 下应优先使用 FSDP 的激活检查点——gradient_checkpointing会在反向传播中引入冗余的 all-gather; - 自动注册生成方法:启用 FSDP2 后,Trainer 会自动调用
dist.fsdp.register_fsdp_forward_method(self.model, "generate")(约第 1724–1727 行),保证 FSDP 包装后的模型仍可正常执行generate推理/生成; - PEFT 兼容:FSDP 与 PEFT LoRA/QLoRA 组合时,Trainer 会调用
update_fsdp_plugin_peft(src/transformers/distributed/fsdp.py)更新自动包装策略与混合精度策略(QLoRA 使用量化存储 dtype); - 生成场景验证:仓库提供了独立的 FSDP 生成测试脚本 tests/trainer/distributed/scripts/fsdp_generate.py,分别演示 FSDP1(
FullyShardedDataParallel+summon_full_params包裹generate)与 FSDP2(fully_shard逐层包装 +register_fsdp_forward_method直接generate)两种写法; - 底层新机制:对于原生 FSDP2 路径,仓库 src/transformers/distributed/fsdp.py 中的
apply_fully_sharded_data_parallelism依据模型声明的_fsdp_plan分片计划,把模块划分为free_full_weight(前向后重分片)与keep_full_weight(保持完整权重,如最终 norm 与 lm_head 组合会被合并为一个不重分片单元以减少反向 all-gather),最后对根模块整体执行fully_shard——这也是理解"Transformer 层包装收益"的最底层实现。
下一步
FSDP 是训练超大规模模型的有力工具,能够充分利用多卡 GPU 或 TPU:通过切分模型参数、优化器状态与梯度,并在空闲时卸载到 CPU,FSDP 可以有效摊薄大规模训练的高昂算力成本。如果想继续深入,仓库内还有以下资料可供参考:
- 英文版 FSDP 指南:面向 FSDP2 的最新讲解(含分片示意图与 JSON 配置示例);
- DDP 数据并行指南:模型可装入单卡时的轻量选择;
- DeepSpeed 指南:ZeRO 优化与 NVMe 卸载的替代方案;
- FSDP 分布式训练测试 与 FSDP2 accelerate 配置:可复现的配置与回归验证样例。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考