使用 PyTorch Lightning Fabric 与 FSDP 训练十亿参数级大模型:完整实战指南
【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning
导读
本文是 PyTorch Lightning Fabric 官方指南 fsdp.rst 的深度解读与实战扩展,围绕Fully Sharded Data Parallel(FSDP,全分片数据并行)展开:从一行代码启用 FSDP,到通过 auto-wrap 策略、sharding strategy、激活检查点、CPU offload 等配置在「显存占用」与「训练吞吐」之间做精细权衡,再到大规模 checkpoint 的保存与恢复。读完本文,你将掌握用 Fabric 在多卡、多机环境下训练数十亿参数模型的完整配置方法、底层原理与排错思路,并能直接复现文中提供的 Transformer 示例。
1. 为什么需要 FSDP:单卡装不下的大模型
训练大模型的显存开销通常由四部分组成:
- 模型参数(weights);
- 前向传播产生的层激活(layer activations);
- 反向传播计算的梯度(gradients);
- 优化器状态(optimizer states,例如 Adam 为每个参数额外维护两个指数滑动平均)。
当这四者之和超过单张 GPU 的显存时,常规的数据并行(DDP)便无法工作。一个直观的参照:即便使用目前最大的 H100 80GB 显存 GPU,在 batch size 为 1、16 位精度的情况下,也不足以训练一个 30B 参数的模型。
FSDP 正是为了解决这一问题而生:它将模型参数、梯度和优化器状态分片(shard)到所有 GPU 上,每个 GPU 只保存全量状态的一个分片,从而把单卡显存需求降到原来的约 1/N(N 为 GPU 数量)。其思想与 ZeRO-Stage 3 类似(见 FSDPStrategy 源码 docstring),并且不需要修改任何模型代码。
Fabric 通过 PyTorch 原生支持 FSDP,实现代码集中在 src/lightning/fabric/strategies/fsdp.py。
使用 FSDP 的前置清单
在动手之前,请确认满足以下条件:
- ✅ 拥有多张 GPU;
- ✅ 已经尝试过普通 DDP 训练(batch size 1),但仍然显存不足;
- ✅ 安装了PyTorch 2.0 或更新版本。
注意:FSDP 对网络带宽要求较高。单卡被 gather 出来的一层在前后向传播时必须能放进该卡显存,且多机场景下 GPU 间的数据传输常常成为瓶颈(参见 model_parallel/index.rst 中 FSDP 的适用性对比)。
2. 在 Fabric 中启用 FSDP
2.1 一行代码启用
在 Fabric 中,启用 FSDP 只需要把strategy参数改为"fsdp":
fabric = L.Fabric(accelerator="cuda", devices=2, strategy="fsdp")字符串"fsdp"由策略注册表自动映射到FSDPStrategy(见 fsdp.py 的register_strategies,注册表同时注册了"fsdp"与"fsdp_cpu_offload"两个别名,后者等价于开启了 CPU offload 的 FSDP)。
2.2 显式构造策略对象
如果后续需要配置更多参数,则显式构造FSDPStrategy:
from lightning.fabric.strategies import FSDPStrategy fabric = L.Fabric(accelerator="cuda", devices=2, strategy=FSDPStrategy())2.3 完整可运行示例
下面是一个完整的 Transformer 训练示例(本文后续所有优化都会基于它展开,并与 DDP 对比):
import torch import torch.nn as nn import torch.nn.functional as F import lightning as L from lightning.fabric.strategies import FSDPStrategy from lightning.pytorch.demos import Transformer, WikiText2 fabric = L.Fabric(accelerator="cuda", devices=2, strategy=FSDPStrategy()) fabric.launch() fabric.seed_everything(42) with fabric.rank_zero_first(): dataset = WikiText2() # 1B 参数的 Transformer model = Transformer(vocab_size=dataset.vocab_size, nlayers=32, nhid=4096, ninp=1024, nhead=64) model = fabric.setup(model) optimizer = torch.optim.Adam(model.parameters(), lr=0.1) optimizer = fabric.setup_optimizers(optimizer) for i in range(10): input, target = fabric.to_device(dataset[i]) output = model(input.unsqueeze(0), target.unsqueeze(0)) loss = F.nll_loss(output, target.view(-1)) fabric.backward(loss) optimizer.step() optimizer.zero_grad() fabric.print(loss.item()) fabric.print(torch.cuda.memory_summary())代码要点说明:
fabric.launch()负责启动多进程分布式环境;fabric.rank_zero_first()保证数据集只在 rank 0 上下载/预处理,其余 rank 等待;fabric.setup(model)会在内部把模型包装成torch.distributed.fsdp.FullyShardedDataParallel(见 fsdp.pysetup_module),并完成参数分片;- 训练循环中的
fabric.backward、optimizer.step()都是标准写法,FSDP 的梯度同步由策略自动管理; - 结尾的
torch.cuda.memory_summary()用于观察 CUDA 显存分配情况,是后续验证优化效果的依据。
从源码实现看,Fabric 会默认向 PyTorch FSDP 传入
use_orig_params=True(见 fsdp.py 第 174-175 行),这使得模型与优化器可以联合设置(setup_module_and_optimizers),并支持多个优化器参数组以及torch.compile()。
3. 识别大层:用 auto_wrap_policy 指定分片单元
3.1 为什么要控制分片粒度
FSDP 受益最大的场景是模型中存在大量大层——例如 LLM、ViT 中的线性层,单层参数超过 1 亿。这些层的参数、激活和优化器状态可以被均匀地分片到所有 GPU 上。
反过来,不要分片只有几千参数的小层:分片后的 gather 通信开销会主导训练,反而拖慢速度。因此 FSDP 引入了wrapping policy(包装策略),用来告诉 FSDP 哪些层需要被单独分片管理。
3.2 集合式策略(推荐,Lightning 2.1+)
只需传入一个包含层类的 set,Fabric 会将其转换为 PyTorch 的ModuleWrapPolicy(见 fsdp.py_auto_wrap_policy_kwargs):
# 1. 定义 FSDP 应该管理的层集合,这里选择大的 encoder/decoder 层 policy = {nn.TransformerEncoderLayer, nn.TransformerDecoderLayer} # 2. 传给 FSDPStrategy strategy = FSDPStrategy(auto_wrap_policy=policy) fabric = L.Fabric(..., strategy=strategy)3.3 旧式函数策略(Lightning < 2.1)
auto_wrap_policy也接受旧式的函数型策略,例如 PyTorch 提供的size_based_auto_wrap_policy:
from functools import partial # 1. 从 PyTorch 导入合适的包装策略 from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy # 2. 配置策略:参数数超过 min_num_params 的层被自动包装 policy = partial(size_based_auto_wrap_policy, min_num_params=10000) # 3. 传给 FSDPStrategy strategy = FSDPStrategy(auto_wrap_policy=policy)PyTorch 在torch.distributed.fsdp.wrap下还提供了其他函数式策略可供选用。
经验法则:典型的做法是把「大块头」层(如 transformer block:attention + feed-forward)放进 policy,让 FSDP 以这些层为单位进行分片;小层(如 embedding、layer norm)保持不包装,减少通信开销。
3.4 验证 FSDP 是否生效
用 2.1 节示例中打印的 CUDA 显存摘要与普通 DDP 训练对比。正确配置后,你应该看到分配的显存下降、单次迭代时间略有上升。以下是作者在 A100 40GB GPU、Lightning 2.1、PyTorch 2.1 环境下测得的数据:
| 指标 | DDP | FSDP |
|---|---|---|
| 显存(MB) | 26,953 | 11,578 |
| 迭代时间(秒) | 0.26 | 0.36 |
FSDP 将显存从约 27GB 降到约 11.6GB,代价是迭代时间从 0.26s 增加到 0.36s——这正是「显存换速度」的典型 trade-off。
4. 加速模型初始化:init_module 与 empty_init
4.1 默认初始化方式的瓶颈
PyTorch 的标准做法是:先在 CPU 内存中创建全部参数,第二步再搬到 GPU。模型越大,这两步耗时越长,而且会瞬间产生巨大的 CPU 内存峰值——对 10B 以上模型,这一步甚至可能直接 OOM。
4.2 用init_module直接建在 GPU 上
Fabric 的fabric.init_module()上下文管理器可以让模型在创建时就落到目标设备与目标精度上(其实现是调用策略的module_init_context,见 fabric.py 的init_module):
# 慢:先在 CPU 上创建模型 model = Transformer(vocab_size=dataset.vocab_size) # 快:直接在 GPU 上创建模型 with fabric.init_module(): model = Transformer(vocab_size=dataset.vocab_size)4.3 FSDP 推荐:empty_init=True
对 FSDP 而言,官方建议设置empty_init=True,这样可以初始化更大的模型:
with fabric.init_module(empty_init=True): model = Transformer(vocab_size=dataset.vocab_size)原理(可从 fsdp.pymodule_init_context的源码得到印证):empty_init=True会让参数创建发生在torch.device("meta")上下文中,即产生不分配任何内存的假参数(meta 设备参数);真正的参数初始化被推迟到fabric.setup(model),此时 FSDP 会先物化参数、调用reset_parameters()、再完成分片,从而避免在任何时刻持有完整模型的真实参数。
使用注意:
empty_init为分布式训练所必需,它要求所有自定义管理参数的模块实现reset_parameters()方法(PyTorch 内置模块都有);- 如果配合加载 checkpoint(微调场景),只要 checkpoint 包含全部参数就是安全的;若以
strict=False加载部分 checkpoint,需自行处理未初始化参数;- 更多
empty_init的使用场景(半精度初始化、加载 checkpoint 做推理/微调等)可参考 model_init.rst。
5. 优化分片策略:四档 sharding strategy 的取舍
5.1 四种 sharding strategy
默认情况下,FSDP 会对被 auto-wrap policy 选中的层,将 1)模型权重、2)反向传播的梯度、3)优化器状态全部分片到所有 GPU。你可以通过sharding_strategy参数调整,在显存与速度之间做交易:
strategy = FSDPStrategy( # 默认:分片权重 + 梯度 + 优化器状态(1 + 2 + 3) sharding_strategy="FULL_SHARD", # 只分片梯度 + 优化器状态(2 + 3) sharding_strategy="SHARD_GRAD_OP", # 机器内 FULL_SHARD,跨机器复制 sharding_strategy="HYBRID_SHARD", # 不分片任何东西(类似 DDP) sharding_strategy="NO_SHARD", ) fabric = L.Fabric(..., strategy=strategy)每种策略的含义(与 FSDPStrategy docstring 完全一致):
| 取值 | 分片内容 | 适用场景 |
|---|---|---|
FULL_SHARD(默认) | 参数 + 梯度 + 优化器状态 | 最省显存,速度最慢 |
SHARD_GRAD_OP | 仅梯度 + 优化器状态(参数复制) | 显存充裕时提速 |
HYBRID_SHARD | 机器内全分片、跨机器复制 | 多机训练,减少跨机通信 |
NO_SHARD | 不分片 | 等价 DDP,仅作对比 |
源码还支持直接传入torch.distributed.fsdp.ShardingStrategy枚举值,字符串不区分大小写(测试见 tests/tests_fabric/strategies/test_fsdp.py)。
两个实现细节需要注意:
HYBRID_SHARD必须配合auto_wrap_policy、process_group或device_mesh之一使用,否则会在构造时报RuntimeError(见 fsdp.py_init_sharding_strategy)。device_mesh接受(replication size, sharding size)元组,乘积须等于 world size;process_group与device_mesh互斥,不能同时传入。
5.2 选择分片策略的推荐配方
- 先试默认的
FULL_SHARD:最慢,但最省显存; - 再试
SHARD_GRAD_OP:若 OOM 就退回默认;否则你会看到迭代速度提升; - 多机训练时试
HYBRID_SHARD:把跨机通信降到最低。
5.3 各策略的实测数据
以下数据同样产自 A100 40GB、Lightning 2.1、PyTorch 2.1:
| 指标 | DDP | NO_SHARD | SHARD_GRAD_OP | FULL_SHARD |
|---|---|---|---|---|
| 显存(MB) | 26,953 | 23,181 | 11,815 | 11,578 |
| 迭代时间(秒) | 0.26 | 0.30 | 0.31 | 0.36 |
可以看到:NO_SHARD只节省少量显存(主要来自 activation 的布局差异),SHARD_GRAD_OP与FULL_SHARD的显存几乎相同,而速度上SHARD_GRAD_OP略快。
6. 用速度换显存:激活检查点与 CPU offload
当模型超过 100 亿参数或需要极大 batch size 时,如果前面几档策略仍不够省显存,可以考虑以下两种「以时间换空间」的手段。
6.1 Activation checkpointing(激活检查点)
前向传播期间,各层的激活值(中间输出)会被保存下来,供反向传播计算梯度时使用。激活检查点的思路是:丢弃选定层的激活,在反向传播需要时重新计算。
开启方法——把需要检查点的层列表传进去,通常就是你的 transformer block(含 attention 与 feed-forward):
strategy = FSDPStrategy( # 在这些层上启用激活检查点 activation_checkpointing_policy={ nn.TransformerEncoderLayer, nn.TransformerDecoderLayer, }, ) fabric = L.Fabric(..., strategy=strategy)要点:
- 如示例所示,
activation_checkpointing_policy通常与auto_wrap_policy保持一致; - 该参数接受 set(内部转换为
ModuleWrapPolicy),也接受函数式策略;旧的activation_checkpointing参数已弃用(见 fsdp.py_activation_checkpointing_kwargs,两者不能同时设置); - 底层通过
torch.distributed.algorithms._checkpoint.checkpoint_wrapper.apply_activation_checkpointing实现(见 fsdp.py_setup_activation_checkpointing,测试见 tests/tests_fabric/strategies/test_fsdp.py); - 代价是训练速度略降,但省下的显存可以用于扩大模型容量或增大 batch size,最终可能反而带来整体性能提升。
6.2 把参数 offload 到 CPU
最激进的显存节省方式是参数 CPU offload:
# 设置 cpu_offload=True strategy = FSDPStrategy(..., cpu_offload=True) fabric = L.Fabric(..., strategy=strategy)代价非常直接:每个 forward pass 都需要在 CPU 与 GPU 之间搬运参数,训练速度显著下降。因此:
- 仅在你有足够的 CPU 内存、且其他手段都无法满足显存需求时才使用;
- 源码层面,
cpu_offload接受布尔值或torch.distributed.fsdp.CPUOffload配置对象,布尔值会被转换为CPUOffload(offload_params=...)(见 fsdp.py_init_cpu_offload,测试见 tests/tests_fabric/strategies/test_fsdp.py)。
作者实测数据(A100 40GB、Lightning 2.1、PyTorch 2.1):
| 指标 | DDP | FSDP | FSDP + CPU offload |
|---|---|---|---|
| 显存(MB) | 26,953 | 11,578 | 2,825 |
| 迭代时间(秒) | 0.26 | 0.36 | 3.24 |
CPU offload 相比纯 FSDP 又把显存压低了约 4 倍(11.6GB → 2.8GB),但迭代时间暴增约 10 倍(0.36s → 3.24s)——这是最极端的 trade-off,请务必谨慎评估。
7. 保存 checkpoint:sharded 与 full 两种格式
大模型训练成本高昂,周期性地保存 checkpoint 是必备的最佳实践,以防训练意外中断导致前功尽弃。
7.1 推荐做法:保存对象引用而非 state_dict
Fabric 提供了高效便捷的保存接口。只需在 state dict 中放入对象本身,而不是手动调state_dict():
# 1. 定义模型、优化器及其他训练循环状态 state = {"model": model, "optimizer": optimizer, "iter": iteration} # ✅ 推荐:使用 Fabric 的方法保存 fabric.save("path/to/checkpoint/file", state) # ❌ 不要这样(低效): # state = {"model": model.state_dict(), "optimizer": optimizer.state_dict(), ...} # torch.save("path/to/checkpoint/file", state)从源码看,FSDPStrategy.save_checkpoint会自动在 state dict 上下文中把模块和优化器对象转换为本地分片的 state dict(模型用module.state_dict(),优化器用FSDP.optim_state_dict),并把其他非模块/优化器的条目作为元数据单独保存(见 fsdp.pysave_checkpoint)。
7.2 sharded 格式的目录结构
默认情况下(state_dict_type="sharded"),每个进程/GPU 各保存自己的分片文件到一个文件夹中,以降低保存时的内存峰值并加快落盘速度。生成的目录结构如下:
path/to/checkpoint/file ├── .metadata ├── __0_0.distcp ├── __1_0.distcp ... └── meta.pt其中.distcp文件包含各进程的张量分片,meta.pt保存除模型/优化器之外的用户元数据(仅由 rank 0 写入)。更多细节可参见 distributed_checkpoint.rst。
7.3 切换为单一文件格式
如果你希望得到一个单一的合并 checkpoint 文件,可通过state_dict_type配置:
# 默认:每个进程保存各自状态的独立文件 strategy = FSDPStrategy(state_dict_type="sharded") # 保存单个合并的 checkpoint 文件 strategy = FSDPStrategy(state_dict_type="full")两种格式的行为差异(源码依据见 fsdp.py 第 134-138 行 docstring):
"full":所有权重和优化器状态在rank 0上汇总,保存为单个文件;"sharded":每个 rank 保存自己的权重/优化器分片,checkpoint 是一个包含与 world size 相同数量文件的目录。
7.4 该选哪种格式?
state_dict_type="sharded":适合预训练超大规模模型。快、省内存,但可移植性差,需要额外步骤将分片 checkpoint 转换为常规文件,参见 分布式 checkpoint 转换指南;state_dict_type="full":适合预训练中小规模模型(<100 亿参数)、微调以及需要可移植性的场景。
另外注意:sharded 格式下暂不支持
filter参数(保存端被显式禁用),而storage_options在 FSDP 策略下不被支持(见 fsdp.py 第 440-449 行)。
8. 加载 checkpoint 恢复训练
加载由 Fabric 保存的 checkpoint 同样简单,且必须传入对象引用:
# 1. 定义模型、优化器及其他训练循环状态 state = {"model": model, "optimizer": optimizer, "iter": iteration} # 2. 使用 Fabric 的方法加载 fabric.load("path/to/checkpoint/file", state) # ❌ 不要这样(低效): # model.load_state_dict(torch.load("path/to/checkpoint/file"))关键行为:
- Fabric自动识别路径中是
state_dict_type="full"还是state_dict_type="sharded"的 checkpoint:full 是单文件,sharded 是包含meta.pt的目录(源码判定逻辑见 fsdp.py_is_sharded_checkpoint/_is_full_checkpoint); "full"格式的 checkpoint 可以被所有策略加载,而"sharded"格式只能被 FSDP 加载;- sharded 格式加载时,优化器状态通过
torch.distributed.checkpoint.optimizer.load_sharded_optimizer_state_dict单独恢复(见 fsdp.pyload_checkpoint); - 若
state中未包含任何被 FSDP 包装的模型,加载会直接报错——请确保传入的是fabric.setup之后的模型对象。
更多的 checkpoint 功能(过滤、迁移、目录结构等)可阅读 checkpoints 指南。
9. 进阶性能优化
当你已经理解前面各参数对显存和速度的影响后,还有两个「锦上添花」的旋钮可以尝试。它们的效果高度依赖具体场景,需要实际开启/关闭对比验证。
9.1 关闭优化器的 foreach
PyTorch 常见优化器都有一个foreach=True|False开关:开启时参数与状态更新会被加速。但代价是可能出现轻微的内存峰值,且模型越大越明显。若观察到异常的内存增长,考虑关闭:
optimizer = torch.optim.AdamW(model.parameters(), foreach=False)支持该参数的全部优化器列表见 PyTorch 官方优化器文档。
9.2 限制 all-gather 调度(limit_all_gathers)
当训练接近显存上限时,你可能在日志中看到CUDA malloc retries:这是 GPU 在即将 OOM 前尝试回收未使用或缓存内存的行为。retry 频繁发生时对速度影响显著。
常规做法是略微减小 batch size,而 FSDP 额外提供了limit_all_gathers旋钮:
strategy = FSDPStrategy( # 默认:CPU 按需调度 GPU 间的权重传输,有时过于激进 limit_all_gathers=False, # 接近显存上限时开启 limit_all_gathers=True, ) fabric = L.Fabric(..., strategy=strategy)你可以通过torch.cuda.memory_summary()或 PyTorch profiler 的输出监控 CUDA malloc retries 的发生频率,据此决定是否开启该选项。
10. 小结:一套完整的 FSDP 调参流程
结合本文内容,推荐的大模型 FSDP 训练调优路径如下:
- 启用:
strategy="fsdp"或显式FSDPStrategy(); - 分片粒度:用
auto_wrap_policy只包装大层(transformer block 等); - 初始化:用
with fabric.init_module(empty_init=True)快速创建超大模型; - 分片策略:默认
FULL_SHARD→ 显存允许时尝试SHARD_GRAD_OP→ 多机用HYBRID_SHARD(需配device_mesh/process_group/auto_wrap_policy); - 进一步省显存:
activation_checkpointing_policy开启激活检查点,最后才考虑cpu_offload=True; - 持久化:预训练大模型用
state_dict_type="sharded",中小模型/微调用"full",用fabric.save/fabric.load传入对象引用; - 微调:必要时关
foreach、开limit_all_gathers=True应对显存压力。
每一步都可以通过torch.cuda.memory_summary()与迭代耗时进行量化对比,从而在「显存」与「吞吐」之间找到适合你硬件与模型规模的平衡点。仓库内的单元测试(tests/tests_fabric/strategies/test_fsdp.py)覆盖了cpu_offload、sharding_strategy、激活检查点与 checkpoint 保存/加载等核心行为,可作为理解各参数语义的补充参考。
【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考