ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch Lightning 分布式检查点完全指南:FSDP 分片保存、弹性加载与单文件转换(Expert 级)

PyTorch Lightning 分布式检查点完全指南:FSDP 分片保存、弹性加载与单文件转换(Expert 级) PyTorch Lightning 分布式检查点完全指南FSDP 分片保存、弹性加载与单文件转换Expert 级【免费下载链接】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 提供的**分布式检查点distributed checkpoint又称 sharded checkpoint**机制允许在多个 GPU 或节点上分片保存与加载训练状态既规避了内存峰值又显著提升保存速度。本文是面向资深用户的分片检查点实战手册覆盖在FSDPStrategy下启用分片格式、保存与加载的完整流程、跨 world size 弹性恢复以及把分片检查点合并转换为标准单文件的方法并穿插仓库源码级原理说明帮你把大规模预训练的中断恢复与模型导出做得又快又稳。本文聚焦 docs/source-pytorch/common/checkpointing_expert.rst 的核心内容如需了解 FSDP 策略本身的更多配置sharding strategy、activation checkpointing、CPU offload 等请参阅 FSDP 训练指南通用检查点机制见 checkpointing.rst。为什么需要分布式检查点在 DDP 或 FSDP 的 full 检查点模式下所有进程的张量状态最终都要汇集到 rank 0再写成一个单一文件。对于拥有数十亿乃至上万亿参数的大模型这一过程会带来两个难以忍受的问题内存峰值rank 0 进程需要同时持有全部权重、优化器状态极易触发 CPU/GPU OOM保存耗时单文件序列化与落盘随模型体积线性增长可能从分钟级恶化到小时级。分布式检查点把问题反过来解决每个进程/GPU 只保存自己负责的那一份张量分片各自写入独立的文件并配套一个轻量的元数据文件。这样既没有汇聚带来的内存峰值也天然并行写盘、缩短了保存时间。从源码结构看PyTorch Lightning 的分布式检查点建立在 PyTorch 的torch.distributed.checkpointAPI 之上保存侧调用torch.distributed.checkpoint.save加载侧调用torch.distributed.checkpoint.load见 src/lightning/fabric/strategies/fsdp.py#L955-L964Lightning 在此基础上负责路径管理、元数据写入与跨框架的状态格式转换。前置条件启用 FSDP 策略分布式检查点格式目前仅在FSDP 策略下可用。只需在FSDPStrategy中把state_dict_type设为sharded再传入Trainer即可import lightning as L from lightning.pytorch.strategies import FSDPStrategy # 1. 选择 FSDP 策略并设置分片分布式检查点格式 strategy FSDPStrategy(state_dict_typesharded) # 2. 把策略传给 Trainer trainer L.Trainer(devices2, strategystrategy, ...) # 3. 运行训练 trainer.fit(model)FSDPStrategy的state_dict_type参数只有两个合法取值见 src/lightning/pytorch/strategies/fsdp.py#L143-L147取值行为适用场景full默认完整权重与优化器状态在 rank 0 汇总后写入单个文件中小规模模型预训练、微调以及需要高可移植性的场景sharded每个 rank 保存自己的权重与优化器状态分片到独立文件检查点是一个包含与 world size 相同数量文件的目录超大规模模型预训练追求更快的保存速度与更低的内存峰值提示默认值full与 FSDP 的 Sharding 策略相互独立——即使训练时参数被 FSDP 切分只要state_dict_typefull检查点仍会在 rank 0 汇聚成单文件对应_get_full_state_dict_context设为sharded后才走_get_sharded_state_dict_context的逐分片路径见 src/lightning/pytorch/strategies/fsdp.py#L521-L558。完整可运行示例使用仓库内置的LightningTransformerdemo 模型定义于 src/lightning/pytorch/demos/transformer.py#L192-L214包含nn.Transformer语言模型与WikiText2数据集的训练 stepimport lightning as L from lightning.pytorch.strategies import FSDPStrategy from lightning.pytorch.demos import LightningTransformer model LightningTransformer() strategy FSDPStrategy(state_dict_typesharded) trainer L.Trainer( acceleratorcuda, devices4, strategystrategy, max_steps3, ) trainer.fit(model)保存分布式检查点检查点目录剖析训练结束后检查点会以目录形式落在默认路径lightning_logs/version_0/checkpoints/epoch0-step3.ckpt/。用如下命令查看其内容ls -a lightning_logs/version_0/checkpoints/epoch0-step3.ckpt/典型的目录结构如下epoch0-step3.ckpt/ ├── __0_0.distcp ├── __1_0.distcp ├── __2_0.distcp ├── __3_0.distcp ├── .metadata └── meta.pt逐项解读__rank_shard.distcp来自每个进程/GPU 的张量分片文件。示例中脚本把模型分布在 4 块 GPU 上因此每个.distcp文件的大小约为整个检查点总大小的 1/4——分片效果一目了然.metadataPyTorchtorch.distributed.checkpoint生成的分片元数据用于在加载时把各分片还原到正确的张量布局meta.ptLightning 额外写入的元数据文件常量_METADATA_FILENAME meta.pt见 src/lightning/fabric/utilities/load.py#L37保存除权重与优化器状态之外的用户数据如 hparams、epoch、global step、callback 状态等。保存流程的源码实现FSDPStrategy.save_checkpoint在分片模式下做了两件事见 src/lightning/pytorch/strategies/fsdp.py#L579-L591把 Lightning 内部约定的state_dict/optimizer_states键转换为分片格式所需的键名——模型权重放进model键多个优化器按optimizer_0、optimizer_1……顺序展开调用_distributed_checkpoint_save即torch.distributed.checkpoint.save写入分片同时由 rank 0 通过_atomic_save把其余元数据原子化写入meta.pt。这也解释了meta.pt只由全局 rank 0 写入、而.distcp文件由各进程并行写出的不对称结构。远程文件系统支持保存路径不必局限在本地磁盘。FSDP 指南fsdp.rst指出检查点路径可以是 fsspec 支持的远程文件系统 URL例如s3://my-bucket/checkpoint、gs://my-bucket/checkpoint或abfs://my-container/checkpoint前提是安装了对应的实现s3fs、gcsfs、adlfs。这为大规模集群训练直接向对象存储写检查点提供了便利。加载分布式检查点如果你的训练脚本同样使用 FSDP加载分布式检查点与保存一样简单——只需在fit时通过ckpt_path传入检查点目录import lightning as L from lightning.pytorch.strategies import FSDPStrategy # 1. 选择 FSDP 策略并设置分片检查点格式 strategy FSDPStrategy(state_dict_typesharded) # 2. 把策略传给 Trainer trainer L.Trainer(devices2, strategystrategy, ...) # 3. 指定要加载的检查点路径 trainer.fit(model, ckpt_pathpath/to/checkpoint)完整示例以 2 块 GPU 恢复之前用 4 块 GPU 保存的检查点import lightning as L from lightning.pytorch.strategies import FSDPStrategy from lightning.pytorch.demos import LightningTransformer model LightningTransformer() strategy FSDPStrategy(state_dict_typesharded) trainer L.Trainer( acceleratorcuda, devices2, strategystrategy, max_steps5, ) trainer.fit(model, ckpt_pathlightning_logs/version_0/checkpoints/epoch0-step3.ckpt)跨 world size 弹性恢复一个经常被忽视的重要特性是即使 world size 发生了变化即当前运行的 GPU 数量与保存检查点时不同也依然可以正常加载。上例中 4 卡保存、2 卡恢复即是直接证据。这得益于torch.distributed.checkpoint以张量分片元数据.metadata为索引的加载方式——加载器依据元数据把每个 shard 的权重搬运到当前各进程对应的设备而不是要求世界拓扑完全一致。加载流程的源码实现FSDPStrategy.load_checkpoint对分片目录的处理见 src/lightning/pytorch/strategies/fsdp.py#L600-L644由 rank 0 广播路径保证所有进程从同一路径加载通过_is_sharded_checkpoint判定路径是分片目录还是普通文件判定条件是是目录且存在meta.pt见 src/lightning/fabric/strategies/fsdp.py#L922-L928分片模式下在_get_sharded_state_dict_context上下文中用torch.distributed.checkpoint.load恢复模型权重再用load_sharded_optimizer_state_dict逐个优化器恢复状态最后读取meta.pt还原训练进度等元数据。重要限制⚠️如果你想把分布式检查点加载到不使用 FSDP甚至不用 Trainer的脚本中必须先把分片检查点转换为单文件格式见下文转换分布式检查点。Trainer 能自动识别路径是full还是sharded格式fsdp.rst但分片检查点只能由 FSDP 加载这一约束不会改变——普通 DDP 或纯 PyTorch 脚本并不理解.distcp分片布局。转换分布式检查点为单文件分片格式虽然高效但可移植性差。当你需要在不使用 FSDP 的脚本中加载检查点把检查点导出为部署、评测等场景需要的其他格式就得先把分片检查点合并成普通单文件。PyTorch Lightning 提供了现成的命令行工具python -m lightning.pytorch.utilities.consolidate_checkpoint path/to/my/checkpoint使用细节与限制纯 CPU 操作无需 GPU转换后检查点中的所有张量都会被转为 CPU 张量执行转换命令本身不需要任何 GPU内存前提该工具假定你有足够的空闲 CPU 内存容纳整个检查点因为需要把所有分片在内存中拼装成完整 state dict输出路径若不指定输出路径转换产物默认保存在输入检查点目录旁、同名加.consolidated后缀也可通过--output_file显式指定且该文件不能已存在见 src/lightning/fabric/utilities/consolidate_checkpoint.py#L63-L72格式校验工具会校验输入必须是 Lightning 保存的 FSDP 分片目录目录内存在元数据文件否则报错退出见 src/lightning/fabric/utilities/consolidate_checkpoint.py#L41-L61。完整示例假设用前面的例子保存了epoch0-step3.ckpt分片目录执行cd lightning_logs/version_0/checkpoints python -m lightning.pytorch.utilities.consolidate_checkpoint epoch0-step3.ckpt工具会在分片检查点旁边生成新文件epoch0-step3.ckpt.consolidated它就是一个标准 PyTorch 检查点可以像普通文件一样加载import torch checkpoint torch.load(epoch0-step3.ckpt.consolidated) print(list(checkpoint.keys())) print(checkpoint[state_dict][model.transformer.decoder.layers.31.norm1.weight])转换背后的实现入口模块 src/lightning/pytorch/utilities/consolidate_checkpoint.py 在__main__中依次完成四步_parse_cli_args()/_process_cli_args()解析并校验命令行参数fabric 侧实现见 src/lightning/fabric/utilities/consolidate_checkpoint.py_load_distributed_checkpoint(config.checkpoint_folder)通过torch.distributed.checkpoint的分片加载器把分片目录还原为完整 state dict并合并meta.pt中的附加数据见 src/lightning/fabric/utilities/load.py#L244-L268_format_checkpoint(checkpoint)把 FSDP 分片格式转换为 Lightning Trainer 可加载的标准格式——model键重命名为state_dictoptimizer_0、optimizer_1等键合并回optimizer_states列表见 src/lightning/pytorch/utilities/consolidate_checkpoint.py#L9-L21_atomic_save(checkpoint, config.output_file)原子化写出合并后的单文件。其中第 3 步的键名重组逻辑有专门的单元测试覆盖tests/tests_pytorch/utilities/test_consolidate_checkpoint.py模型键改名、多个优化器状态含乱序的optimizer_1、optimizer_0正确合并为optimizer_states列表、无关键如optimizer_abc原样保留。这也提示了一个易错点——如果你的检查点里自定义了optimizer_*命名键转换时需留意命名冲突。实战建议与常见误区大模型预训练选sharded中小模型与微调选full分片格式快且省内存但可移植性差full格式单文件、随处可加载。取舍依据是模型规模与是否需要跨策略复用fsdp.rst 对两者适用场景有明确划分。恢复训练时保持策略一致分片检查点只能被 FSDP 加载恢复脚本中务必同样配置FSDPStrategy(state_dict_typesharded)state_dict_type本身是保存格式开关加载分片目录时无需也无法切换为full。world size 可变放心在不同 GPU 数量间迁移分片检查点如 8 卡中断、4 卡续训元数据驱动的加载器会处理分片重排。转换前预估 CPU 内存合并过程需要完整 state dict 常驻内存超大模型请确认单机内存充足转换全程不需要 GPU可放到纯 CPU 节点上执行。避免手写torch.save(model.state_dict())式保存在 FSDP 下这会绕过 Lightning 的上下文管理器_get_sharded_state_dict_context大概率保存出残缺或错误布局的状态务必使用 Trainer 的检查点机制见 fsdp.rst 的 Save a checkpoint 小节。总结分布式检查点是 PyTorch Lightning 支撑超大模型训练的关键能力之一FSDPStrategy(state_dict_typesharded)一行开启后检查点以目录形式分片落盘保存更快、内存峰值更低加载时支持跨 world size 弹性恢复需要可移植性时python -m lightning.pytorch.utilities.consolidate_checkpoint一行命令即可把分片目录合并为标准单文件供非 FSDP 脚本、部署与评测流程直接torch.load。理解.distcp、.metadata、meta.pt三类文件的职责掌握 src/lightning/pytorch/strategies/fsdp.py、src/lightning/fabric/utilities/load.py 与 src/lightning/pytorch/utilities/consolidate_checkpoint.py 中的实现细节你就能在万卡规模的训练任务中设计出稳健、可恢复、可迁移的检查点流水线。【免费下载链接】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),仅供参考
返回列表