尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

使用 PyTorch Lightning Fabric 与 FSDP 训练十亿参数级大模型:完整实战指南

使用 PyTorch Lightning Fabric 与 FSDP 训练十亿参数级大模型:完整实战指南 使用 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 ParallelFSDP全分片数据并行展开从一行代码启用 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/NN 为 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 中启用 FSDP2.1 一行代码启用在 Fabric 中启用 FSDP 只需要把strategy参数改为fsdpfabric L.Fabric(acceleratorcuda, devices2, strategyfsdp)字符串fsdp由策略注册表自动映射到FSDPStrategy见 fsdp.py 的register_strategies注册表同时注册了fsdp与fsdp_cpu_offload两个别名后者等价于开启了 CPU offload 的 FSDP。2.2 显式构造策略对象如果后续需要配置更多参数则显式构造FSDPStrategyfrom lightning.fabric.strategies import FSDPStrategy fabric L.Fabric(acceleratorcuda, devices2, strategyFSDPStrategy())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(acceleratorcuda, devices2, strategyFSDPStrategy()) fabric.launch() fabric.seed_everything(42) with fabric.rank_zero_first(): dataset WikiText2() # 1B 参数的 Transformer model Transformer(vocab_sizedataset.vocab_size, nlayers32, nhid4096, ninp1024, nhead64) model fabric.setup(model) optimizer torch.optim.Adam(model.parameters(), lr0.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_paramsTrue见 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只需传入一个包含层类的 setFabric 会将其转换为 PyTorch 的ModuleWrapPolicy见 fsdp.py_auto_wrap_policy_kwargs# 1. 定义 FSDP 应该管理的层集合这里选择大的 encoder/decoder 层 policy {nn.TransformerEncoderLayer, nn.TransformerDecoderLayer} # 2. 传给 FSDPStrategy strategy FSDPStrategy(auto_wrap_policypolicy) fabric L.Fabric(..., strategystrategy)3.3 旧式函数策略Lightning 2.1auto_wrap_policy也接受旧式的函数型策略例如 PyTorch 提供的size_based_auto_wrap_policyfrom 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_params10000) # 3. 传给 FSDPStrategy strategy FSDPStrategy(auto_wrap_policypolicy)PyTorch 在torch.distributed.fsdp.wrap下还提供了其他函数式策略可供选用。经验法则典型的做法是把「大块头」层如 transformer blockattention feed-forward放进 policy让 FSDP 以这些层为单位进行分片小层如 embedding、layer norm保持不包装减少通信开销。3.4 验证 FSDP 是否生效用 2.1 节示例中打印的 CUDA 显存摘要与普通 DDP 训练对比。正确配置后你应该看到分配的显存下降、单次迭代时间略有上升。以下是作者在 A100 40GB GPU、Lightning 2.1、PyTorch 2.1 环境下测得的数据指标DDPFSDP显存MB26,95311,578迭代时间秒0.260.36FSDP 将显存从约 27GB 降到约 11.6GB代价是迭代时间从 0.26s 增加到 0.36s——这正是「显存换速度」的典型 trade-off。4. 加速模型初始化init_module 与 empty_init4.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_sizedataset.vocab_size) # 快直接在 GPU 上创建模型 with fabric.init_module(): model Transformer(vocab_sizedataset.vocab_size)4.3 FSDP 推荐empty_initTrue对 FSDP 而言官方建议设置empty_initTrue这样可以初始化更大的模型with fabric.init_module(empty_initTrue): model Transformer(vocab_sizedataset.vocab_size)原理可从 fsdp.pymodule_init_context的源码得到印证empty_initTrue会让参数创建发生在torch.device(meta)上下文中即产生不分配任何内存的假参数meta 设备参数真正的参数初始化被推迟到fabric.setup(model)此时 FSDP 会先物化参数、调用reset_parameters()、再完成分片从而避免在任何时刻持有完整模型的真实参数。使用注意empty_init为分布式训练所必需它要求所有自定义管理参数的模块实现reset_parameters()方法PyTorch 内置模块都有如果配合加载 checkpoint微调场景只要 checkpoint 包含全部参数就是安全的若以strictFalse加载部分 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_strategyFULL_SHARD, # 只分片梯度 优化器状态2 3 sharding_strategySHARD_GRAD_OP, # 机器内 FULL_SHARD跨机器复制 sharding_strategyHYBRID_SHARD, # 不分片任何东西类似 DDP sharding_strategyNO_SHARD, ) fabric L.Fabric(..., strategystrategy)每种策略的含义与 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 sizeprocess_group与device_mesh互斥不能同时传入。5.2 选择分片策略的推荐配方先试默认的FULL_SHARD最慢但最省显存再试SHARD_GRAD_OP若 OOM 就退回默认否则你会看到迭代速度提升多机训练时试HYBRID_SHARD把跨机通信降到最低。5.3 各策略的实测数据以下数据同样产自 A100 40GB、Lightning 2.1、PyTorch 2.1指标DDPNO_SHARDSHARD_GRAD_OPFULL_SHARD显存MB26,95323,18111,81511,578迭代时间秒0.260.300.310.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-forwardstrategy FSDPStrategy( # 在这些层上启用激活检查点 activation_checkpointing_policy{ nn.TransformerEncoderLayer, nn.TransformerDecoderLayer, }, ) fabric L.Fabric(..., strategystrategy)要点如示例所示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_offloadTrue strategy FSDPStrategy(..., cpu_offloadTrue) fabric L.Fabric(..., strategystrategy)代价非常直接每个 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显存MB26,95311,5782,825迭代时间秒0.260.363.24CPU offload 相比纯 FSDP 又把显存压低了约 4 倍11.6GB → 2.8GB但迭代时间暴增约 10 倍0.36s → 3.24s——这是最极端的 trade-off请务必谨慎评估。7. 保存 checkpointsharded 与 full 两种格式大模型训练成本高昂周期性地保存 checkpoint 是必备的最佳实践以防训练意外中断导致前功尽弃。7.1 推荐做法保存对象引用而非 state_dictFabric 提供了高效便捷的保存接口。只需在 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_typesharded每个进程/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_typesharded) # 保存单个合并的 checkpoint 文件 strategy FSDPStrategy(state_dict_typefull)两种格式的行为差异源码依据见 fsdp.py 第 134-138 行 docstringfull所有权重和优化器状态在rank 0上汇总保存为单个文件sharded每个 rank 保存自己的权重/优化器分片checkpoint 是一个包含与 world size 相同数量文件的目录。7.4 该选哪种格式state_dict_typesharded适合预训练超大规模模型。快、省内存但可移植性差需要额外步骤将分片 checkpoint 转换为常规文件参见 分布式 checkpoint 转换指南state_dict_typefull适合预训练中小规模模型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_typefull还是state_dict_typesharded的 checkpointfull 是单文件sharded 是包含meta.pt的目录源码判定逻辑见 fsdp.py_is_sharded_checkpoint/_is_full_checkpointfull格式的 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 关闭优化器的 foreachPyTorch 常见优化器都有一个foreachTrue|False开关开启时参数与状态更新会被加速。但代价是可能出现轻微的内存峰值且模型越大越明显。若观察到异常的内存增长考虑关闭optimizer torch.optim.AdamW(model.parameters(), foreachFalse)支持该参数的全部优化器列表见 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_gathersFalse, # 接近显存上限时开启 limit_all_gathersTrue, ) fabric L.Fabric(..., strategystrategy)你可以通过torch.cuda.memory_summary()或 PyTorch profiler 的输出监控 CUDA malloc retries 的发生频率据此决定是否开启该选项。10. 小结一套完整的 FSDP 调参流程结合本文内容推荐的大模型 FSDP 训练调优路径如下启用strategyfsdp或显式FSDPStrategy()分片粒度用auto_wrap_policy只包装大层transformer block 等初始化用with fabric.init_module(empty_initTrue)快速创建超大模型分片策略默认FULL_SHARD→ 显存允许时尝试SHARD_GRAD_OP→ 多机用HYBRID_SHARD需配device_mesh/process_group/auto_wrap_policy进一步省显存activation_checkpointing_policy开启激活检查点最后才考虑cpu_offloadTrue持久化预训练大模型用state_dict_typesharded中小模型/微调用full用fabric.save/fabric.load传入对象引用微调必要时关foreach、开limit_all_gathersTrue应对显存压力。每一步都可以通过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),仅供参考
返回列表