使用 PyTorch Lightning Fabric 与 FSDP 训练十亿参数级大模型:完整实战指南
2026/9/19 22:14:10 网站建设 项目流程

使用 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:单卡装不下的大模型

训练大模型的显存开销通常由四部分组成:

  1. 模型参数(weights);
  2. 前向传播产生的层激活(layer activations);
  3. 反向传播计算的梯度(gradients);
  4. 优化器状态(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.backwardoptimizer.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 环境下测得的数据:

指标DDPFSDP
显存(MB)26,95311,578
迭代时间(秒)0.260.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_policyprocess_groupdevice_mesh之一使用,否则会在构造时报RuntimeError(见 fsdp.py_init_sharding_strategy)。device_mesh接受(replication size, sharding size)元组,乘积须等于 world size;
  • process_groupdevice_mesh互斥,不能同时传入。

5.2 选择分片策略的推荐配方

  1. 先试默认的FULL_SHARD:最慢,但最省显存;
  2. 再试SHARD_GRAD_OP:若 OOM 就退回默认;否则你会看到迭代速度提升;
  3. 多机训练时试HYBRID_SHARD:把跨机通信降到最低。

5.3 各策略的实测数据

以下数据同样产自 A100 40GB、Lightning 2.1、PyTorch 2.1:

指标DDPNO_SHARDSHARD_GRAD_OPFULL_SHARD
显存(MB)26,95323,18111,81511,578
迭代时间(秒)0.260.300.310.36

可以看到:NO_SHARD只节省少量显存(主要来自 activation 的布局差异),SHARD_GRAD_OPFULL_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):

指标DDPFSDPFSDP + CPU offload
显存(MB)26,95311,5782,825
迭代时间(秒)0.260.363.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 训练调优路径如下:

  1. 启用strategy="fsdp"或显式FSDPStrategy()
  2. 分片粒度:用auto_wrap_policy只包装大层(transformer block 等);
  3. 初始化:用with fabric.init_module(empty_init=True)快速创建超大模型;
  4. 分片策略:默认FULL_SHARD→ 显存允许时尝试SHARD_GRAD_OP→ 多机用HYBRID_SHARD(需配device_mesh/process_group/auto_wrap_policy);
  5. 进一步省显存activation_checkpointing_policy开启激活检查点,最后才考虑cpu_offload=True
  6. 持久化:预训练大模型用state_dict_type="sharded",中小模型/微调用"full",用fabric.save/fabric.load传入对象引用;
  7. 微调:必要时关foreach、开limit_all_gathers=True应对显存压力。

每一步都可以通过torch.cuda.memory_summary()与迭代耗时进行量化对比,从而在「显存」与「吞吐」之间找到适合你硬件与模型规模的平衡点。仓库内的单元测试(tests/tests_fabric/strategies/test_fsdp.py)覆盖了cpu_offloadsharding_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询