ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

混合精度与分布式训练:大模型显存优化与多卡并行实战指南

混合精度与分布式训练:大模型显存优化与多卡并行实战指南 1. 训练大模型显存先从“抠”开始做了几年大模型训练我越来越觉得所谓“底层学习”其实就是一门关于“资源守恒”的学问。当你第一次把目光从模型代码挪到训练框架的底层实现时最先撞到你脸上的问题一定是显存怎么就那么不够用呢我印象很深有一回在单张A10080GB上尝试训练一个7B参数的模型最开始用的全是FP32单精度。模型权重占28GBAdam优化器状态占56GB梯度占28GB光这三项加起来就已经112GB远超单卡上限。还没算激活值、中间变量、通信缓冲区这些东西。显然这条路走不通。于是混合精度训练几乎成了必选项——它的核心思想用一个词就能概括在不明显损失精度的前提下把内存占用砍下去一大截。那为什么前几年大家没把混合精度当回事因为模型小显存还够用训练时间也短单卡跑个几十小时也能出结果。但大模型时代参数量从几亿涨到几百亿显存和算力的矛盾彻底暴露混合精度才从“锦上添花”变成“雪中送炭”。再加上多卡并行也就是所谓的分布式训练两者一拍即合成了现代大模型训练的两根柱子。这篇文章我根据自己的实际经验把这两块掰开揉碎讲清楚混合精度的底层原理、损失缩放是怎么运作的、分布式有哪几种并行策略、实际部署时怎么配参数、还有我踩过的那些坑。内容偏工程实践尽量讲人话不会上来就甩公式。提示这篇文章适合有一定深度学习基础、准备把项目从单卡扩展到多卡的读者。如果你还在跑几百M的小模型也可以先收藏等显存告急的时候再回头来看。2. 混合精度训练原理与核心细节2.1 三种精度到底差在哪FP32、FP16、BF16先说精度。FP32是单精度浮点数1位符号位、8位指数位、23位尾数位动态范围大概是 ( \pm 3.4 \times 10^{38} )一般训练任务里的数值都在这个范围内所以早年大家直接用它算没什么大毛病。FP16是半精度1位符号位、5位指数位、10位尾数位。优点很明显内存占用是FP32的一半计算速度在支持FP16加速的GPU上能翻倍甚至更多但缺点也致命尾数只有10位有效精度大约只有3位十进制数字指数位只有5位最大表示到65504。如果梯度或者激活值超过了这个范围直接变成无穷大然后整个loss变成NaN训练直接崩掉。BF16bfloat16是Google为深度学习专门设计的格式1位符号位、8位指数位、7位尾数位。它的思路很巧妙保留和FP32一样大的指数范围砍掉的是尾数精度。什么意思就是说BF16能表示的数的“量级范围”和FP32几乎一样不会轻易溢出但每个数值的“精细程度”降低了。对训练来说权重更新的方向主要看数值的相对大小关系对绝对精度没那么敏感所以BF16在大多数深度学习场景下都非常稳。我把常见的三种精度整理成一张表方便对照精度类型符号位指数位尾数位显存占比主要问题FP3218234字节显存开销大训练慢FP1615102字节动态范围窄需损失缩放BF161872字节精度较低个别任务会掉点用生活化的话说FP16像一个量程只有65504的量杯水稍微大点就溢出来了你得时刻盯着、时不时把水倒掉一点BF16是量程很大但刻度粗糙的量杯基本不会溢出来但你读数时要心里有数它没那么精细。这也是为什么现在大模型预训练基本都在用BF16。2.2 混合精度不是“无脑换成半精度”把模型从FP32换成FP16或BF16看起来很简单但实际上没那么简单。直接全部换成FP16大概率会遇到两个问题第一个问题是梯度下溢。深度学习中很多梯度的数值非常小小到接近 ( 10^{-7} ) 甚至更小。在FP16里最小的正规正数大约是 ( 6.1 \times 10^{-5} )比这个数还小的梯度直接变成0。梯度变0权重就不更新了模型就像被按了暂停键早期某些层永远学不动。第二个问题是动态范围溢出。某些层的梯度偶尔会出现很大的瞬时值一旦超过65504直接变成Inf然后loss变NaN整个训练流程白跑。那怎么办混合精度训练不是所有环节都用低精度它保留三份关键数据权重主副本始终保存一份FP32的权重用于精确的参数更新累积。梯度在反向传播时用FP16或BF16计算然后转回FP32再用于更新。优化器状态像Adam这类优化器会保存动量项一阶矩和方差项二阶矩这两个状态也需要FP32精度否则更新会漂移。一句话总结计算用低精度存储和更新用高精度。这套组合既获得了速度收益又保住了稳定性。2.3 损失缩放是怎么回事针对FP16的下溢问题最巧妙的手段就是“损失缩放”loss scaling。思路特别直白既然小梯度在FP16里会变成0那我先把loss乘以一个比较大的系数比如 ( 1024 ) 或 ( 4096 )这样反向传播算出来的梯度整体就放大到了FP16能安全表示的范围。等梯度算完再除以这个系数恢复到原始大小去更新参数。实际操作中缩放系数的选择有两种策略固定缩放选一个固定值从头用到尾适合比较稳定的训练任务但如果模型结构特殊容易踩坑。动态缩放定期尝试加大缩放系数如果连续一段时间没出现溢出就说明还有余量再往上调一旦发现Inf或NaN就回退到之前的系数同时跳过这一轮更新。我一般建议用动态缩放。PyTorch的torch.cuda.amp和深水区的DeepSpeed/Megatron-LM都内置了这套机制默认配置能覆盖绝大多数场景。这里有个非常重要的点梯度裁剪gradient clipping和损失缩放要放在一起考虑。如果你先做loss缩放再对参与更新的FP32梯度做裁剪没问题但如果对低精度梯度先裁剪再转成FP32裁剪阈值没有跟缩放联动很容易造成裁剪失效或者梯度爆炸。我之前帮朋友排查过一个训练发散问题最后发现就是裁剪顺序搞反了。2.4 混合精度在PyTorch里的实现细节PyTorch的AMPAutomatic Mixed Precision实现起来非常简单核心代码就那么几行import torch from torch.cuda.amp import autocast, GradScaler model MyModel().cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) scaler GradScaler() for batch in dataloader: with autocast(): loss model(batch).loss optimizer.zero_grad() scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update()这段代码里有几个容易被忽略的细节autocast()不是把整个模型强行转成FP16而是自动选择合适的op精度。卷积、矩阵乘这些计算密集型的操作走FP16而像归一化、softmax这些对数值范围敏感的op仍然是FP32这是混合精度性能好的关键。scaler.scale(loss)执行两步放大loss然后反向传播这一步记录放大的梯度。scaler.unscale_(optimizer)把梯度恢复成原始值同时检测是否有Inf/NaN。如果发现异常scaler.step(optimizer)会识别出来并跳过这一次参数更新不至于直接崩掉。scaler.update()根据最近几轮是否出现溢出动态调整缩放系数。如果你是新手建议先原样跑通这套模板再去动细节。很多人觉得AMP就是“加两行代码”的事但实际里面每一步的顺序、时机都是有讲究的我后面在问题排查部分还会补充。3. 分布式训练从三板斧到组合拳3.1 单卡不够用自然想到多卡当你把混合精度打开显存压力小了一些但模型继续变大你会发现单卡怎么都不够。要么模型权重放不下要么batch size被迫开得很小导致训练效率低下。这时候就只能上多卡了。分布式训练的本质就是把一个大任务拆成多个子任务在多张GPU上同时执行再想办法把结果合到一起。听起来简单但“怎么拆”和“怎么合”决定了系统的效率上限。大模型领域的典型拆分方式有四种数据并行、张量并行、流水线并行、以及以ZeRO为代表的分片策略。3.2 数据并行最简单的并行方式数据并行Data ParallelismDP的逻辑是每张卡都放一份完整的模型副本然后把一份大数据集切成多份每张卡处理不同的小批次。所有卡算完梯度之后用AllReduce通信把每张卡的梯度做一次全局求平均然后每张卡拿着平均后的梯度更新自己的模型副本。用生活化的类比来说这就像一个年级考试后把答案发给大家互批每个人批一批卷子最后大家聚在一起把每个题的扣分情况合计一遍得到统一的评分标准再各自核对自己的卷子。数据并行在中小模型上非常有效代码也最成熟PyTorch的DistributedDataParallel就是典型实现。但它的缺陷也很明显每张卡都要放一份完整模型显存不足时根本无法启动。每轮训练都要做一次全局梯度同步通信量跟模型参数量成正比模型越大通信压力越大。对大模型来说单独用数据并行解决不了根本问题它必须跟其他策略组合使用。3.3 张量并行把模型“横着切”张量并行Tensor ParallelismTP的思路是把某一层的权重矩阵按行或按列切成多块分别放在不同GPU上计算时通过AllReduce通信把拆分的结果拼回来。比如Transformer里的自注意力层可以把多头拆到不同卡上计算输出再合并。这种方式的优点是规避了单卡显存不足的问题缺点是拆分的粒度很细每做一次前向和反向都会引入额外的通信。如果跨节点比如跨服务器的NVLink网络做张量并行通信延迟会显著拖慢训练。所以实践上张量并行基本限制在单节点内部卡间通过NVLink或PCIe高速互联通信。Megatron-LM里张量并行默认是1D切分只按一个维度切更高维度的切分2D、3D会进一步增加通信开销一般情况不建议一上来就上。3.4 流水线并行把网络“竖着切”流水线并行Pipeline ParallelismPP的思路和张量并行正相反不是把一层切开而是把整个网络按层切成几段每段放在一张卡上。数据流像流水线一样从第一段模型流到第二段再到第三段每张卡只负责自己那几层的计算。我见过一个很好的类比流水线并行就像工厂里的流水线每个工人只负责自己工位的组装大批零件从一个工位流向另一个工位。虽然单个零件在流水线上从头到尾的时间变长了延迟增加但整条产线单位时间的产量吞吐量提升了。这就是为什么系统吞吐量上去了但单个batch的端到端延迟不一定变快。流水线并行主要问题是“气泡”bubble初始阶段后面的卡在等前面卡算完收尾阶段前面的卡在等后面卡算完导致部分GPU空转。业界通过更细粒度的微批次调度比如1F1B调度策略来减小气泡占比但没法完全消除。3.5 ZeRO把“没必要重复”的冗余删掉数据并行里的每张卡保存一份完整模型副本其实有大量冗余。ZeRO零冗余优化器就是用来做“去重”的它不复制完整的模型状态权重、梯度、优化器状态而是把不同部分分片到不同卡上每张卡只存自己负责的那一瓣。先说一个关键点Adam优化器状态占据的空间甚至比模型权重本身还大。对于一个7B参数的模型权重用BF16存储是14GB梯度又是14GBAdam状态一阶矩、二阶矩是56GB总计84GB。如果数据并行有16张卡这84GB每张卡都要存一份总共占84×161344GB只是重复存储实际每卡84GB。而ZeRO-3的分片策略能把这份开销平摊到16张卡上每张卡只需要存约5.25GB省下来的显存是巨大的。ZeRO分成几个级别ZeRO-1只分片优化器状态。ZeRO-2分片优化器状态梯度。ZeRO-3分片优化器状态梯度模型权重显存压到最低但通信量也最大。DeepSpeed是ZeRO的主要实现框架配置极其方便{ zero_optimization: { stage: 3, offload_optimizer: { device: cpu }, overlap_comm: true, contiguous_gradients: true } }注意offload_optimizer开启后优化器状态会放到CPU内存里这会显著降低单卡显存需求但CPU和GPU之间的传输会成为新的瓶颈。我一般建议只有当显存真的不够时才考虑CPU offload如果显存能硬扛别开。3.6 四种策略怎么配合实际大模型训练很少只用一种并行策略通常是组合拳。常见组合如下并行策略核心思想适合场景主要瓶颈数据并行全模型重复切数据小模型、大吞吐显存冗余、通信量大张量并行层内切分权重和计算大模型单节点多卡通信延迟敏感流水线并行层间切分网络超深网络、多节点存在气泡ZeRO分片状态分片到各卡大模型显存不够时动态通信开销常见的部署模式是节点间用流水线并行GPipe/PP节点内用张量并行TP配合数据并行DP再用DeepSpeed的ZeRO-3吃显存红利。具体怎么配跟你的模型大小、GPU数量、节点拓扑都有关系没有放之四海而皆准的配置得做实验对比。4. 实操配置与参数调节经验4.1 我的实际部署从单卡到8卡说一个我最近亲测过的案例。模型是13B参数的LLM预训练数据大概500GB文本单卡A100 80GB直接训练很快就OOM。我最后花了点时间配置成了8卡A100单机8卡NVLink全互联加ZeRO-3加BF16混合精度。大概步骤如下先用BF16跑一遍纯单卡确认模型本身没问题记录baseline loss曲线。在单卡上打开混合精度确认loss曲线跟FP32基本一致记录显存和吞吐。切到单机8卡用torchrun启动数据并行核对梯度同步是否正确。启动DeepSpeed ZeRO-3配合zero_allow_untested_optimizer等配置项逐步调整。做一次小规模的batch size和learning rate联动测试确认稳定性。用torchrun启动8卡的命令大概长这样torchrun --nproc_per_node8 train.py \ --model_name_or_path your_13b_model \ --bf16 True \ --deepspeed ds_config.json \ --per_device_train_batch_size 8 \ --gradient_accumulation_steps 4 \ --learning_rate 1e-5 \ --num_train_epochs 3注意这里的gradient_accumulation_steps它的作用是在不增大单卡batch size的前提下累积多次小batch的梯度再统一更新。当时我单卡只能放batch size 8但用4步累积后等效batch size是8×8×4256既保证了显存不爆又保证了训练的稳定性。4.2 学习率和batch size的联动法则很多人忽略了一个问题多卡训练时总batch size变了learning rate也得跟着调。最简单的经验法则是“线性缩放”——batch size变成原来的K倍learning rate也应该变成原来的K倍左右。比如单卡batch size32、lr1e-4现在用8卡、batch size256那么初始lr可以调整到 ( 1e-4 \times 8 8e-4 ) 左右。但线性缩放不是万能的batch size大到一定程度比如超过训练集大小的1/10线性缩放就失效了这时候需要改用“平方根缩放”或者加warmup阶段。我实操中的习惯是先跑几百步小规模的实验观察loss曲线和梯度范数再决定最终lr。另外一个细节是梯度累积与学习率的配合如果通过gradient_accumulation_steps把等效batch size调大了learning rate同样需要相应拉高但很多人只改了数据并行数量忘了累积步数也算进去结果模型训得又慢又稳过头学不动。4.3 通信库和网络拓扑的影响分布式训练的性能很大程度取决于通信效率。NVIDIA的GPU之间有NVLink单机内高带宽低延迟和InfiniBand跨机两种主流互联方式。通信库选择上NCCL是事实标准PyTorch和DeepSpeed都默认走它。我在排查性能问题时发现NCCL的通信时间和计算时间经常是重叠的。比如AllReduce可以在反向传播还没完全结束时就启动一部分通信这就是所谓的“通信计算重叠”。想让重叠生效除了框架支持还需要你开启overlap_comm这样的配置项尤其是在ZeRO-3下通信量更大重叠优化带来的收益非常明显。如果你有多台机器节点间的网络带宽非常关键。以太网千兆和InfiniBand 200Gbps之间的训练速度差距可能有十倍这不是量变而是质变。跑千亿级模型时网络基本就是第一瓶颈这个钱省不得。4.4 如何找到对应你显卡的合理配置很多新手喜欢copy网上的配置文件但你的显卡、显存、模型大小跟人家不一定一样。我建议用以下步骤判断先用单卡跑一次看能不能塞下模型显存占用多少。估算“模型状态”的显存需求。以13B模型为例BF16权重26GB、梯度26GB、Adam优化器状态104GB总计156GB。单卡80GB自然放不下得靠分片。想看ZeRO-3能帮你省多少可以先用DeepSpeed自带的显存计算工具或者直接配置好之后跑一个很小的step看实际显存占用。如果显存还是不够再考虑梯度/激活值检查点activation checkpointing用时间换空间。激活值检查点这个技术容易被忽略但对大模型训练帮助很大。它的做法是不存储所有前向传播的中间激活值而是在反向传播时重新计算一遍。用一次额外的前向计算换取大量显存释放。PyTorch里开一行就行model.gradient_checkpointing_enable()通常能帮你把激活值显存砍掉60%-70%但训练时间会增加15%-25%左右。5. 常见问题与排错实录5.1 loss变成NaN或Inf这是混合精度训练里最经典的问题。我遇到过的情况大致分三类梯度溢出。fp16下数值超过65504导致Inf。解决办法确认开启了动态损失缩放如果要手动设建议先设GradScaler(init_scale1024)观察一段时间。学习率过高。多卡并行后总batch size变大lr没跟着调导致更新步长过大。排查方法是把lr降到原来的1/10看是否恢复稳定。数据本身包含异常值。我没太想到的是某次预训练语料里有几篇异常文档token里混入了大段乱码导致loss直接爆炸。后来加了clip_grad_norm_和异常样本过滤问题才彻底解决。5.2 通了分布式但整体速度没有提升比如8卡训练跟你单卡训练一对比吞吐量竟然只提升了1.5倍最常见的几个原因batch size太小通信开销占比过高。每轮通信都要同步梯度小batch下计算时间短通信时间被拉长相对占比。解决办法适度增大batch size或者调大gradient_accumulation_steps。跨节点网络带宽不够。如果你用的是千兆以太网AllReduce通信会让网络拥塞到怀疑人生。最好改成InfiniBand或者至少用RDMA/GPUDirect。CPU数据加载成瓶颈。GPU算得快但数据来不急只能空等。排查方法看GPU利用率如果频繁掉到0多半是数据加载问题。用DataLoader的num_workers调高或者用tf.data? 不对用PyTorch的话开num_workers8以上并配合prefetch_factor。5.3 手把手排查训练速度慢我自己习惯用这样的排查顺序先开nvidia-smi看GPU利用率和显存占用。GPU利用率持续低于90%大概率是CPU数据加载或者通信热点。用nsys profile或者PyTorch Profiler看时间分布。如果反向传播时间占比巨大可能需要关注通信与计算重叠是否生效。查NCCL日志。设NCCL_DEBUGINFO看AllReduce是否频繁报错或重试。重点看跨节点通信时间。如果一次AllReduce几百毫秒可能得考虑减少TP并行度、增加PP并行度。5.4 分布式下模型精度掉点经常有人问同一份代码单卡跑出来的效果很好8卡跑完效果打折到底差在哪别的先不提最可疑的是loss scaling导致的梯度差异、以及通信同步浮点精度问题。分布式训练里每张卡算的梯度汇总做平均时浮点数相加的先后顺序会影响微小精度这种影响通常可以忽略不计但如果你开了FP16浮点误差会被放大。另一个容易忽略的是随机种子多卡训练时如果没有给每张卡设置不同的随机种子数据加载顺序会出问题相当于某些样本被反复看到。正确的处理是主进程统一设种子然后通过DistributedSampler保证每张卡拿到的数据互斥且可shuffle。5.5 一张踩坑清单我把这些年踩过的坑整理成一份速查表每次出问题先对一遍问题现象推荐排查方向OOM训练一开始就内存溢出检查是否开启混合精度/ZeRO/激活检查点吞吐量低多卡利用率不足检查通信库配置、数据加载速度训练发散loss升高或NaN检查lr、损失缩放、梯度裁剪顺序通信超时NCCL报错检查网络带宽、NCCL版本、环境变量模型不收敛loss下降但效果差检查随机种子、数据并行切分是否正确6. 最后的实操建议我一直觉得混合精度和分布式训练这两个话题光看文档是学不会的。真正有价值的是你亲手在显卡上踩过坑然后把那些报错信息一个接一个解决掉的过程。根据我自己的经验可以给几个可以直接抄的结论第一30B以下的模型优先尝试BF16混合精度加ZeRO-3这是性价比最高、配置也最简单的方式。如果遇到loss不稳定再加warmup和动态损失缩放。第二如果单节点8卡以内能解决的问题尽量别跨节点部署通信延迟会带来很多额外的麻烦。第三所有参数调节都别一步到位先在“百步级”小实验中摸清楚loss曲线和梯度范数的走势再放大到完整训练。第四日志和监控一定要从一开始就建立起来我用的是WandB加nvidia-smi监控确保任何一个环节出问题都能在十分钟内定位。我踩过的最深的一个坑是某次64卡训练时NCCL因为一台机器的网卡驱动版本不同反复导致AllReduce超时排查了整整两天才发现是环境不一致。那之后我就规定所有机器必须用同一份容器镜像和驱动版本环境问题从源头杜绝。大模型训练这条路上难度不在于某个单一技术有多复杂而在于所有环节是环环相扣的。混合精度解决显存与速度的矛盾分布式解决单机算力不足的问题二者结合才有了今天几十B、上百B模型训练的可能。希望这篇文章能帮你少走一些弯路。真到了显存吃紧、多卡部署的时候回头看看这些细节应该能省不少事。
返回列表