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

资讯详情

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

PyTorch多卡训练:DP与DDP原理、排错和调优实战

PyTorch多卡训练:DP与DDP原理、排错和调优实战 上周有个同学拿着他的训练脚本找我说四张卡一起跑一步的耗时反而比单卡多了三成问我是不是有一张卡坏了。我让他把nn.DataParallel那层壳换掉改成DistributedDataParallel其他代码几乎没动同样的 batch size每步时间直接掉到原来的四成左右。这不是显卡的问题也不是玄学而是PyTorch里两种数据并行方案底层的工作方式完全不同——一个是单进程里开多线程一个是多进程各自为战再靠集合通信把梯度对齐。这个差别决定了你的多卡到底是111还是111.8。这篇东西我打算把PyTorch并行训练DPDDP的原理和应用从头捋一遍。不管你是刚拿到多卡机器、还在用DataParallel凑合的同学还是已经在跑 DDP 但总觉得速度没榨干、时不时被 NCCL 卡死搞心态的老手我尽量把每一步的为什么讲透把踩过的坑按排查链路写清楚最后给一份能直接抄进项目的改造清单。文中所有实测数字都来自常见配置下的经验值具体到你的机器会有出入但量级和趋势基本一致。1. DP 和 DDP 到底在并行什么聊并行训练的时候很多人第一反应是把 batch 拆到多张卡上算这个理解只对了一半。真正需要被拆开的其实是三样东西输入数据、计算图上的前向/反向、以及梯度与参数的同步方式。前两样决定了算得快不快第三样决定了开销大不大——而 DP 和 DDP 的全部差异几乎都集中在第三样上。1.1 DataParallel 的单进程多线程模型与它的天花板nn.DataParallel用起来确实简单一行model nn.DataParallel(model)就能跑但它的内部逻辑要复杂得多。整个过程是这样的主进程拿到一个 batch把 batch 沿第 0 维切成 N 份scatter到 N 张卡然后在每一次 forward 之前把主卡上的模型参数replicate复制到其他卡上各卡算完前向把输出gather回主卡的device_ids[0]主卡上算 loss、做 backward得到主卡那份梯度最后再把梯度scatter回各卡各卡各自执行optimizer.step()。问题就出在这套流程里。第一参数复制发生在每一次 forward 之前如果你的模型有 1 亿参数单机 8 卡就是每次迭代多传 7 亿个 float即使走 PCIe 也是实打实的开销。第二loss 的计算、反向传播的起点、输出的 gather 全部压在 0 号卡上0 号卡显存占用明显高于其他卡计算也更忙整个集群的速度被最慢的那张卡拖住。第三它是单进程多线程Python 的 GIL 让多个线程无法真正并行地调度 CUDA kernel主机侧的调度很容易变成瓶颈。第四它只支持单机device_ids里写不了别的机器。我实测过一个 2000 万参数的小模型单机 4 卡 2080TiDataParallel相比单卡的加速比只有 2.3 倍左右而换成 DDP 能到 3.6 倍。模型越大、卡越多这个差距越夸张。所以我的建议很直接DataParallel 只适合 2 到 3 张卡上跑个 demo 或者快速验证任何打算认真训的东西都别用它。1.2 DDP 的多进程哲学与 Ring AllReduce 的通信量账本DistributedDataParallel走的是完全相反的路线每张卡对应一个独立的 Python 进程每个进程里有自己完整的模型副本、自己的优化器、自己的数据加载器。既然每个进程都在各算各的那凭什么保证所有卡的参数始终一致靠的是三件事。一是初始化时的一次广播。DDP(model)这个构造动作内部会从 rank 0 把参数和 buffer 广播给所有其他 rank保证大家起跑线相同。这一次性的开销比 DP 每个 iteration 都复制一遍参数划算太多。二是反向传播过程中的梯度 AllReduce。每个进程算完自己那份 loss 和反向梯度后DDP 会把梯度按 bucket 分组一块准备好就触发一次AllReduce求和然后内部除以 world size得到全局平均梯度。因为所有进程对所有参数用的都是同一份全局平均梯度optimizer.step()之后参数自然保持一致不需要任何额外的同步。三是通信和计算的 overlap。反向传播是逐层产生梯度的DDP 在某一层梯度算出来之后立刻发起这一层所在 bucket 的通信同时 CPU 和 GPU 继续算更靠前的层。通信被藏在计算里这是 DDP 效率高的关键。这里必须说说 Ring AllReduce 的通信量账本理解了它你就知道为什么 DDP 能扩展到几十上百张卡。整个过程分两步第一步叫scatter-reduce把每个 rank 的梯度切成 N 份沿着环状拓扑传 N-1 轮每轮每张卡收到一块就先累加再传给下一张转完一圈后每一块都有一张卡握着它的全局和第二步叫all-gather把这些完整的块沿着同样的环再传 N-1 轮让每张卡都补齐全部块。总传输量是2(N-1)/N × 参数量当卡数 N 增大时这个系数趋近于 2也就是说每个 rank 需要收发的数据量几乎与卡数无关。这就是 DDP 能线性扩展的底层原因。1.3 DP 与 DDP 的关键差异对照把两者的差异列成一张表更清楚这张表我在选型的时候会直接拿来对着看维度DataParallelDistributedDataParallel进程模型单进程多线程主卡协调每张卡一个独立进程通信后端无靠共享显存和 CUDA 拷贝NCCL / Gloo / MPI参数同步时机每次 forward 前复制全部参数初始化时广播一次梯度同步方式主卡 reduce 后 scatterRing AllReduce是否支持多机不支持支持GIL 影响明显无负载均衡主卡偏重显存和算力不均衡各卡均衡代码改动量一行十几行适用场景2-3 卡快速验证生产级训练还有一点容易被忽略DDP 的每个进程都独立跑一份DataLoader也就是说数据加载、预处理、增强这些 CPU 工作也是并行的而 DP 只有一个进程在加载数据数据管道很容易成为瓶颈。这也是为什么有些人换了 DDP 之后发现不仅快了GPU 利用率还稳了。2. 把 DDP 拉起来需要的四件事DDP 的原理不难难的是把环境配通。我见过太多人卡在进程起不来或者起来了但 rank 之间不通这种问题上。其实只要把下面四件事按顺序做对第一版能跑的 DDP 半小时就能搭出来。2.1 init_process_group 与后端选择NCCL 还是 Glooinit_process_group是整个分布式训练的握手协议它要回答三个问题总共几个进程world size、我是第几个rank、其他人在哪master addr 和 port。用torchrun启动的话这些信息会通过环境变量WORLD_SIZE、RANK、LOCAL_RANK、MASTER_ADDR、MASTER_PORT自动注入代码里不用手写。后端选择有个简单的判断原则只要涉及 GPU 训练就用nccl。NCCL 是专门为 NVIDIA 显卡之间的通信做的库能自动利用 NVLink、PCIe P2P 甚至 RDMA速度远好于其他选项。只有纯 CPU 训练或者需要跨平台调试时才用gloo。我一般还会显式指定一个超时时间避免某个进程出问题后整个任务无限期挂着from datetime import timedelta dist.init_process_group( backendnccl, timeouttimedelta(minutes30), )另外一个必须做、但很多人会忘的动作是torch.cuda.set_device(local_rank)。它的作用是告诉当前进程你只负责这一张卡。如果不设所有进程默认都会往cuda:0上申请显存结果就是 0 号卡爆显存、其他卡闲着。这类问题的典型症状是8 卡机器只能跑起来 1 个进程一查就发现全挤在 0 卡上了。2.2 DistributedSampler 的切分逻辑与 set_epoch 的真实作用DDP 里每张卡都得看到不同的数据否则多个进程重复算同一批样本等于白干。DistributedSampler干的就是这件事它按照rank和num_replicas把数据集的下标做均匀切分rank 0 拿第 0、N、2N…… 号样本rank 1 拿第 1、N1、2N1…… 号样本依次类推。这里有两个细节值得单独说。第一当数据集长度不能被卡数整除时DistributedSampler默认会从头部拿一些样本补齐到整除为止导致部分样本在一个 epoch 内被重复采样。要避免这个问题给它传drop_lastTrue数据集样本量远大于卡数时丢掉尾巴上的几个样本几乎没影响但能让每个 epoch 的样本数严格一致指标也更好对齐。第二个细节是set_epoch。DistributedSampler在shuffleTrue时用rank epoch作为随机种子来打乱顺序如果你每个 epoch 开始前不调用sampler.set_epoch(epoch)那所有 epoch 的样本顺序会完全一样。这个 bug 特别隐蔽loss 在降模型也在收敛但收敛速度莫名比单卡慢很多人查半天查不出原因。养成习惯在每个 epoch 循环的第一行就调sampler.set_epoch(epoch)。2.3 torchrun 与 mp.spawn启动方式怎么选启动方式有两代。老一点的写法是torch.multiprocessing.spawn在代码里自己 fork 进程import torch.multiprocessing as mp if __name__ __main__: mp.spawn(main, nprocs8, args())新一点的写法是torchrunPyTorch 1.9 之前的名字是torch.distributed.launch从命令行拉起进程torchrun --nproc_per_node8 --nnodes1 --node_rank0 \ --master_addr127.0.0.1 --master_port29500 train.py我现在的项目基本都用torchrun原因有三个。一是它支持多机只需要加--nnodes和--node_rank参数扩机器时改动最小二是进程崩溃时它会自己回收不会留下一堆僵尸进程三是它支持--rdzv_backend之类的弹性训练参数配合调度系统更省事。spawn唯一的优势是调试时可以直接在 IDE 里跑但torchrun也能用--standalone参数做到类似效果。提示--master_port建议别用 29500 这种默认值多个人在同一台机器上跑任务时特别容易撞端口报错是Address already in use。换成 29500 到 29600 之间的随机端口能避开大部分冲突。2.4 一份可以直接复制的最小可运行骨架说了这么多概念不如直接上代码。下面这份骨架我删掉了所有业务逻辑只保留分布式相关的部分你把它套进自己的项目就能跑import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.data.distributed import DistributedSampler def main(): dist.init_process_group(backendnccl) local_rank int(os.environ[LOCAL_RANK]) rank int(os.environ[RANK]) world_size int(os.environ[WORLD_SIZE]) torch.cuda.set_device(local_rank) device torch.device(cuda, local_rank) model MyModel().to(device) model DDP(model, device_ids[local_rank]) optimizer torch.optim.AdamW(model.parameters(), lr3e-4) train_set MyDataset(train) sampler DistributedSampler( train_set, num_replicasworld_size, rankrank, shuffleTrue, drop_lastTrue, ) loader DataLoader( train_set, batch_size64, samplersampler, num_workers8, pin_memoryTrue, drop_lastTrue, ) for epoch in range(num_epochs): sampler.set_epoch(epoch) model.train() for batch in loader: batch {k: v.to(device, non_blockingTrue) for k, v in batch.items()} loss model(**batch) optimizer.zero_grad(set_to_noneTrue) loss.backward() optimizer.step() if rank 0: torch.save(model.module.state_dict(), fckpt_ep{epoch}.pt) dist.destroy_process_group() if __name__ __main__: main()注意最后保存时用的是model.module.state_dict()因为DDP包装之后模型外面多了一层直接存model.state_dict()会得到所有 key 都带module.前缀的权重加载到单卡模型上就会报一堆 missing keys。3. 从单卡脚本改造成 DDP按顺序改这七处有了骨架接下来是把你现有的单卡脚本改过来。我按照踩坑频率从高到低把要改的地方列出来你对照着检查一遍基本能避开九成的入门问题。3.1 模型包装与优化器创建的先后顺序这个顺序问题网上说法很乱我直接给结论先DDP(model)再建optimizer。原因是DDP在构造时会做一次参数广播如果此时优化器已经建好了虽然它引用的是同一批Parameter对象DDP 是原地包装不替换参数理论上不会出问题但先包装后建优化器读起来语义更干净也不会在某些自定义优化器实现里踩坑。对应的错误写法是model DDP(nn.DataParallel(model))这种套娃两个并行的壳叠在一起通信路径会变得非常奇怪我第一次见这种写法的时候排查了整整一个下午。另外device_ids[local_rank]这个参数只在单机多卡、一卡一进程的模式下有意义它的作用是让 DDP 在 forward 时自动把输入张量搬到对应设备上。多机场景下这个参数应该省略。3.2 find_unused_parameters默认关闭不是坑乱开才是find_unused_parameters是 DDP 里最有争议的参数。它的默认值是False意思是我假设这次反向传播里所有需要梯度的参数都会参与计算。如果实际不是这样——比如你的模型有多分支结构、某些样本只走其中一条分支或者像多任务学习那样不同任务共享部分参数——那 DDP 在反向时会因为某个 bucket 一直等不到梯度而卡住然后抛出一个非常绕的报错RuntimeError: Expected to have finished reduction in the prior iteration before starting a new one. This error indicates that your module has parameters that did not produce a gradient.很多人一看到这个报错就条件反射地把find_unused_parametersTrue加上能跑就完事了。但要知道这个开关的代价它会让 DDP 在每次反向时都遍历一遍 autograd 图去找哪些参数没被用到这个遍历在模型大、分支多的时候能吃掉 10% 到 20% 的吞吐。所以我的处理原则是先判断为什么会有参数不参与计算。如果是控制流写得不严谨比如某个 if 分支只在部分 rank 上成立那就改代码而不是加开关如果确实是合法的动态结构MoE、动态深度、多任务那该开就开同时看看能不能通过重构模型让未使用参数变成显式跳过而不是隐式缺席。3.3 训练循环里的梯度累积与 no_sync显存不够的时候大家都会用梯度累积把大 batch 拆成几个小 batch累加梯度后再 step 一次。在 DDP 里这个操作有个专属的优化点——用no_sync跳过中间那些反向传播的通信。import contextlib accum_steps 4 for i, batch in enumerate(loader): is_last (i 1) % accum_steps 0 ctx contextlib.nullcontext() if is_last else model.no_sync() with ctx: loss model(**batch) / accum_steps loss.backward() if is_last: optimizer.step() optimizer.zero_grad(set_to_noneTrue)原理很简单梯度累加是线性的中间几次本地累加的中间结果没必要让所有卡都知道等最后一次凑齐了再一起 AllReduce 就行。这么做能把通信次数从accum_steps次降到 1 次通信占比高的时候提速非常明显。有两个坑必须提。第一accum_steps必须在所有 rank 上完全一致否则有的卡在通信有的卡不在直接死锁。第二划分 dataset 的时候要注意最后一个不完整的累积窗口如果len(loader) % accum_steps ! 0末尾会残留几个 batch 没被 step最好在DataLoader上加drop_lastTrue保证整除。还有一个更隐蔽的坑如果你写的是loss loss / accum_steps注意别在循环外又除一次。我之前看到过一份代码在no_sync分支里除以accum_steps在最后一步又除了一次结果学习率等效缩小了 4 倍训练曲线看着正常但收敛慢了一个数量级。3.4 日志、指标、checkpoint 的 rank 分流处理DDP 里所有进程都在跑同一份代码如果你不做判断屏幕上会瞬间刷出 8 份一模一样的日志tensorboard 里也会写 8 条重复曲线更糟的是 8 个进程同时往同一个文件写可能写出半个损坏的 checkpoint。处理原则很直接只让 rank 0 干输出这件事。if rank 0: logger.info(fepoch {epoch} loss {loss.item():.4f}) writer.add_scalar(loss, loss.item(), global_step) if rank 0: torch.save({...}, ckpt.pt)但有个细节要注意loss.item()是当前 rank 本地 batch 的平均 loss不同 rank 之间会有差异。如果你希望日志里的 loss 是全局的需要手动做一次 AllReducedef reduce_mean(tensor): tensor tensor.clone() dist.all_reduce(tensor, opdist.ReduceOp.SUM) tensor / dist.get_world_size() return tensor评估指标accuracy、F1、AUC更是必须这么处理因为每个 rank 看到的是自己那片数据直接把本地指标当全局指标会差很多。最常见做法是把本地的(sum, count)都 AllReduce 一遍再相除比先算比值再平均更准确。3.5 验证阶段的三条路验证阶段怎么跑很多人第一次做的时候都会纠结。我这里给三条路按场景选。第一条是只在 rank 0 上跑全量验证。改动最小if rank 0:包住整个验证循环用普通的DataLoader不加 sampler就行。缺点是其他卡在等着GPU 利用率低但验证本身占比不大一般能接受。第二条是用 DistributedSampler 分片验证。每张卡验证自己那片最后把指标 AllReduce 起来。速度快但要注意两点一是sampler.set_epoch(0)别忘否则每次验证顺序都不同二是样本数不能整除时要处理补齐带来的重复样本否则准确率会略微偏低。稳妥的做法是给验证 sampler 加drop_lastTrue或者在统计时记录每个样本只算一次。第三条是推理预测结果用 all_gather 汇总。当你需要拿到每个样本的具体预测值比如算 mAP、做后处理时就得把各卡的预测张量收集起来。注意各卡的张量长度可能不一致要先用all_gather_object收集长度或者先把张量 pad 到相同长度再all_gather不然直接调用会因为形状不匹配报错。我的经验是小数据集验证直接走第一条省心大数据集或者需要频繁验证比如 early stopping 每轮都跑就走第二条需要导出预测结果的场景才用第三条。4. DDP 排错实录那些报错和看起来正常但不对的现象跑起来不难难的是跑稳。下面这几个问题我在不同项目里反复遇到过把排查链路完整写出来希望对你有用。4.1 显存不均衡均分了数据为什么 rank 0 还是更胖DDP 理论上各卡负载是均衡的但实际训练时经常发现 rank 0 的显存比其他卡高出 1 到 2 GB。原因通常有这几类。最常见的是日志和 checkpoint 造成的额外开销。如果你的 logger 或者 tensorboard writer 在 rank 0 上缓存了大量标量、图像、直方图这些数据会一直占着主机内存甚至显存。曾经有个项目rank 0 比别的卡多占 3 GB 显存最后发现是有人把中间层的 feature map 也写进了 tensorboard。第二类是CUDA context 和显存碎片的差异。rank 0 通常是第一个初始化的进程它可能额外承担了 NCCL communicator 的创建、cuBLAS handle 的初始化等工作这些都会占用一部分显存。这个差异一般是几百 MB属于正常范围。第三类是数据加载的起点差异。如果某个数据集在第一个 batch 特别大比如按长度排序后没有 shufflerank 0 拿到的第一批可能就是最长的那些样本峰值显存自然高。解决办法是在 sampler 之前先做一次全局 shuffle。如果 rank 0 和其他卡的显存差距超过 20%那就不是正常开销了先查日志相关代码再查有没有 rank 相关的分支导致某些 rank 多存了东西。4.2 BatchNorm 在数据并行下到底同步了什么这是我在面试里最喜欢问的一个问题答对的人不多。先说结论DDP默认的broadcast_buffersTrue并不等于同步 BatchNorm 的统计量。它的实际行为是在每次 forward 之前把 rank 0 上的 buffer包括 BN 的running_mean和running_var广播给所有其他 rank。也就是说所有卡用的都是 rank 0 统计出来的那份均值和方差。而running_mean的更新发生在 forward 过程中各卡是基于自己那片数据更新的所以如果不同步各卡的 buffer 会逐渐漂移最终导致各卡的输出不一致、进而参数不一致。那什么时候该上SyncBatchNorm当每卡的 batch size 太小的时候。假设你全局 batch size 是 328 张卡分下来每卡只有 4 个样本BN 用 4 个样本估出来的均值方差噪声极大训练会非常不稳定。这时候用model torch.nn.SyncBatchNorm.convert_sync_batchnorm(model) model DDP(model, device_ids[local_rank])SyncBatchNorm会在 forward 时做一次额外的 AllReduce把各卡的 batch 统计量合并起来算全局均值和方差。代价是通信量增加、速度变慢而且在梯度累积场景下它会和no_sync打架因为 SyncBN 的通信不受no_sync控制需要额外注意。如果你的网络允许我更推荐直接换掉 BN用GroupNorm或者LayerNorm。这两个都不依赖 batch 维统计天然免疫这个问题在小 batch 场景下比 SyncBN 稳定得多还省了通信。视觉检测和分割任务里现在用 GroupNorm 的越来越多了这是个趋势。4.3 NCCL 卡死、掉卡、超时的排查链路这类问题最烦人因为现象往往是程序不动了也不报错。我一般的排查顺序是这样的。第一步确认是不是卡在通信上。用py-spy dump --pid PID看一眼每个进程的调用栈如果全部停在nccl相关的函数上那基本确定是通信问题如果有的进程停在数据加载上那可能是 DataLoader 死锁。第二步打开详细的调试信息重跑。设这几个环境变量export NCCL_DEBUGINFO export NCCL_DEBUG_SUBSYSINIT,COLL export TORCH_DISTRIBUTED_DEBUGDETAILNCCL_DEBUGINFO会打印每张卡用了什么拓扑NVLink 还是 PCIe、走不走 RDMA这一步能暴露出很多硬件层面的问题。我用这个方式定位过一次机房里的显卡 P2P 通路异常日志里明确写着 fallback 到了 shared memory。第三步确认各 rank 的配置一致。最常见的坑是nproc_per_node设成了 8但实际CUDA_VISIBLE_DEVICES只给了 4 张卡于是后面的进程拿不到 GPU 就卡在初始化。另一个高频坑是不同机器上的 PyTorch 版本、NCCL 版本不一致导致握手成功但通信协议对不上。第四步如果报的是超时把超时时间调长再看dist.init_process_group(backendnccl, timeouttimedelta(minutes60))如果调长之后能跑完说明是某一步特别慢比如某个 rank 在做大量数据处理这时候问题不在通信而在数据管道。顺便说一句从 PyTorch 1.10 开始TORCH_NCCL_ASYNC_ERROR_HANDLING1会让 NCCL 的错误及早暴露成异常而不是静默挂起调试阶段建议打开。4.4 学习率、batch size 与 warmup 的换算换到多卡之后全局 batch size 变大了学习率要不要跟着调要但要算清楚。先明确有效 batch size 的算法effective_batch per_gpu_batch × world_size × accum_steps比如单卡 batch 168 卡梯度累积 4 步有效 batch 就是 512。如果单卡时代你的 batch 是 64、学习率是 1e-4换到 DDP 之后有效 batch 变成 512按线性缩放规则linear scaling rule学习率应该近似放大到 8e-4。但这个规则只在 batch 不是特别大的时候成立超过某个阈值通常几千后线性关系就失效了这时候要用平方根缩放或者干脆保持学习率不变、只延长训练步数。我的经验是有效 batch 在 512 以内线性缩放基本安全1024 到 4096 之间先按平方根试再往上就别动学习率了改改 warmup。warmup 一定要跟着调。batch 变大之后每一步的梯度方差变小前期可以用更大的步长但如果一上来就用大学习率很容易直接发散。我通常设置 warmup 步数为总步数的 3% 到 5%或者固定 2000 步取小值。用torch.optim.lr_scheduler.OneCycleLR或者get_cosine_schedule_with_warmup都可以但要注意总步数必须按 DDP 之后的步数来算也就是len(loader) × epochs而len(loader)已经变成单卡样本数除以 batch size 了别用单卡时代的数字。有个特别容易踩的坑如果你用了no_sync做梯度累积那么一个 optimizer step 对应accum_steps个loader迭代scheduler 的步数基准要用 optimizer step 数而不是迭代数。写错了会导致学习率曲线走完得太快或太慢。5. 调优旋钮通信、显存与精度能跑通之后接下来就是抠性能。这部分内容比较细但对吞吐的影响很直接。5.1 bucket、gradient_as_bucket_view 与 static_graphDDP 有三个值得调的参数。第一个是gradient_as_bucket_viewTrue。默认情况下DDP 会把 AllReduce 的结果从通信 bucket 拷贝回各参数自己的.grad上多了一次显存拷贝。开启这个选项后.grad直接指向 bucket 里的那块内存省下的显存大致等于全部梯度的量。对一个 1 亿参数的模型来说FP32 梯度就是 400 MB这 400 MB 有时就是 OOM 和 OOM 不出来的区别。model DDP(model, device_ids[local_rank], gradient_as_bucket_viewTrue)第二个是bucket_cap_mb默认 25 MB。这个值控制每次通信的粒度bucket 太小通信次数多、启动开销大bucket 太大通信的启动被推迟overlap 效果变差而且峰值显存更高。我一般在大模型上试 50 到 100小模型上试 10 到 25实测哪个快用哪个没有通用最优值。第三个是static_graphTruePyTorch 1.10。如果你的模型每次迭代的计算图完全一样没有动态分支、没有条件执行的层、没有依赖迭代次数的控制流那开启它能省掉 DDP 每轮的正向图分析还能配合find_unused_parametersTrue一起用而不产生额外开销。反过来说如果你的模型有/if判断或者动态深度千万别开会产生一个很难查的 error。注意static_graphTrue和find_unused_parametersTrue同时用时DDP 要求未使用的参数集合在整个训练过程中保持不变。如果你有时候用某个头、有时候不用这个条件就不满足会报Static graph mismatch之类的错调起来很痛苦。5.2 AMP 与 DDP 组合时的梯度缩放陷阱混合精度训练在 DDP 下用起来和单卡差不多scaler torch.cuda.amp.GradScaler() with torch.autocast(cuda, dtypetorch.bfloat16): loss model(**batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()梯度缩放本身和 AllReduce 是兼容的因为缩放是线性操作先把每张卡的梯度乘上 scale 再求和等于先求和再乘 scale结果一样。所以 AllReduce 拿到的是正确的缩放后全局梯度scaler.step()内部 unscale 之后再更新参数逻辑是通的。问题出在溢出检测上。GradScaler检测 inf/nan 是在本地做的如果 rank 3 的某一步梯度溢出了而其他 rank 没有rank 3 会跳过这次 step、并把 scale 减半其他 rank 照常 step 并保持 scale 不变——参数就不一致了。轻则后续 step 结果错乱重则直接报 NCCL 相关的古怪错误。PyTorch 较新的版本对这个情况做了一些处理但如果你用的是老版本、或者自己写了自定义的训练循环稳妥做法是手动同步一次found_inf torch.tensor( 1.0 if scaler._get_inf_count() else 0.0, devicedevice ) dist.all_reduce(found_inf, opdist.ReduceOp.MAX) if found_inf.item() 0: scaler._scale scaler._scale / 2 # 跳过这次 step并让所有 rank 一起跳过如果你觉得上面这段太脏最省事的方案是直接用 bfloat16Ampere 及以上架构它不需要损失缩放压根没这个问题。BF16 的动态范围和 FP32 一样宽精度略低但现代模型基本都能扛住我现在新项目基本默认用它。5.3 从 DDP 迈向 FSDP/ZeRO 的判定条件DDP 有一个无法绕过的限制每张卡都要存一份完整的模型参数、梯度和优化器状态。用 Adam 的话显存占用大概是参数量的 16 倍FP16 权重 2 字节 FP32 主权重 4 字节 两个动量各 4 字节 梯度 2 字节。10 亿参数的模型光这一套就是 16 GB还没算激活值。所以当你遇到下面这几种情况时就该考虑 FSDPFully Sharded Data Parallel或者 DeepSpeed 的 ZeRO 系列了单卡装不下模型本身无论 batch size 调到多小都 OOM优化器状态占用太大导致 batch size 被压到个位数训练效率极低需要把参数或优化器状态 offload 到 CPU 内存甚至 NVMe。FSDP 的思路是把参数、梯度、优化器状态都按卡数切分每张卡只存 1/N计算到某一层时临时把所有分片 AllGather 起来组成完整的层算完立刻释放。它的通信量比 DDP 大大约 1.5 倍但显存占用能降一个数量级。代价是代码改动比 DDP 大需要处理auto_wrap_policy、mixed_precision策略、分片 checkpoint 的保存加载调试难度也高不少。我的建议是能用 DDP 解决的绝不上 FSDP。只有当显存真的不够时再上而且一旦上了留出足够的调试时间FSDP 的报错信息普遍比 DDP 难懂。6. checkpoint、续训与推理导出的收尾工作训练能跑、速度快了最后收尾的部分反而是最容易被糊弄的。断点续训没做好一次意外中断就白跑三天。6.1 断点续训必须落盘的东西只存模型权重是绝对不够的一个完整的 checkpoint 至少包括这些内容为什么必须存model.state_dict()模型权重本身optimizer.state_dict()Adam 的动量和二阶矩丢了等于从头训lr_scheduler.state_dict()学习率进度丢了会回到初始学习率scaler.state_dict()AMP 的缩放因子丢了前期会重新试探epoch和global_step恢复训练进度和日志对齐sampler的 epoch保证数据 shuffle 顺序能接上RNG 状态Python / NumPy / CUDA 三套随机数状态RNG 状态这一项经常被忽略但如果你的数据增强里有随机裁剪、随机遮挡不恢复 RNG 会导致每次续训后前几十步的样本分布和之前不一样loss 曲线会出现一个明显的小凸起。恢复写法state torch.load(path, map_locationcpu) model.load_state_dict(state[model]) optimizer.load_state_dict(state[optimizer]) scheduler.load_state_dict(state[scheduler]) scaler.load_state_dict(state[scaler]) start_epoch state[epoch] 1 torch.set_rng_state(state[rng]) torch.cuda.set_rng_state_all(state[cuda_rng])还有一个容易忽略的点加载之后要保证所有 rank 的参数一致。虽然各 rank 读的是同一个文件理论上没问题但如果文件是 rank 0 在训练过程中写的、而其他 rank 读到的是写了一半的版本就会出问题。解决办法是加一个dist.barrier()rank 0 写完之后过 barrier其他 rank 收到信号才去读。6.2 保存策略rank0 汇总 vs 分片保存DDP 下所有 rank 的模型完全一致所以最简单也最推荐的策略是只让 rank 0 保存完整权重。这份权重可以直接加载到单卡模型上做推理不需要任何额外处理。但如果你用的是 FSDP每个 rank 只有一部分参数分片就必须用FSDP.state_dict_type()配合FullStateDictConfig汇聚成完整权重或者保存分片权重ShardedStateDictConfig以实现快速的分布式加载。前者适合导出给推理引擎用后者适合大规模续训避免 rank 0 内存被打爆。推理导出的时候记得去掉module.前缀或者干脆在保存时就存model.module.state_dict()。如果忘了推理端加载时会看到满屏的Unexpected key(s) in state_dict: module.xxx虽然可以用strictFalse硬加载但不如一开始就存干净。还有个小技巧分享如果你的模型用了 EMA指数移动平均记得把 EMA 的权重也单独存一份并且在验证时用 EMA 权重而不是训练权重。我见过不少项目训完之后直接拿训练权重推理效果比验证时看到的差一截查半天才发现是 EMA 忘了切回来。我个人在多卡训练上折腾了几年最深的体会是DDP 的代码改动量其实很小真正的成本在于把分布式环境的各种边界情况搞清楚——端口冲突、进程数不匹配、sampler 忘了 set_epoch、梯度累积的除法写重了。这些东西在文档里都是一句话带过但落到具体项目上每一个都能耗掉你半天。我的做法是给自己维护一个dist_utils.py把init_process_group、日志分流、指标 AllReduce、checkpoint 存取这些封装成函数新项目直接 import前面踩过的坑就不用再踩第二遍。另外第一次上多机的时候先拿一个两层的小 MLP 跑通全流程确认通信、checkpoint、续训都没问题了再换成真正的模型——用大模型排查分布式问题纯粹是给自己找罪受。
返回列表