十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

混合精度与分布式训练:大模型显存优化的核心实战指南

混合精度与分布式训练:大模型显存优化的核心实战指南 大模型训练跑到一定阶段很多人都会遇到同一个坎模型规模一上去单卡就装不下了。就算勉强塞进去训练速度也慢得让人怀疑人生。这个时候混合精度训练和分布式训练就成了绕不开的两个话题。这两个东西听起来像是两个独立的技术方向但实际上在大模型训练里是牢牢绑在一起的——没有混合精度分布式训练要花的卡数和成本会翻好几倍没有分布式混合精度帮你省下来的显存也撑不起真正的大模型。这篇是这个系列的第四篇我把这两个主题放在一起讲因为分开学容易学成两张皮合在一起才能理解大模型训练到底是怎么在硬件上跑起来的。这篇文章并不是一个纯理论科普也不是纯工程手册而是从“显存从哪里来、又去了哪里”这条主线出发把混合精度的原理、分布式训练的几种并行方式以及它们如何组合才能真正落地的完整链路理清楚。适合已经跑通简单模型训练、准备迈向真正大模型训练的开发者阅读。如果你对训练脚本里那些--fp16、--zero-stage参数知其然不知其所以然这篇文章就是帮你补上底层那块拼图的。1. 为什么说混合精度和分布式是“绑在一起”的先不急着写代码先用一道算术题把问题逼出来。假设我要训练一个7B参数量的模型也就是70亿参数。如果全程用FP32单精度浮点数每个数占4字节训练光是模型本身的参数就有70亿乘以4字节大约28GB。放到一张80GB显存的A100或H100上参数就占了三分之一强的显存。但这只是参数本身。训练过程中系统还要保存梯度梯度在FP32下又是28GB。更狠的是优化器状态——用Adam优化器的话每个参数要额外保存一阶动量m和二阶动量v这两个同样是FP32精度加起来又是56GB。把这三块加一加参数28GB加梯度28GB加优化器状态56GB总共112GB。一张80GB的卡已经装不下了。这就是大模型训练遇上的第一堵墙模型参数、梯度、优化器状态三座大山把显存吃干抹净还不够。而混合精度训练和分布式训练是从两个完全不同的维度来解决这个问题。混合精度解决的是“单卡内每一份数据能不能更省”的问题。核心思路很直白在计算过程中用FP16半精度浮点数占2字节或BF16另一种半精度格式也占2字节来存储参数和梯度把显存需求直接砍半。按上面的算法参数加梯度从56GB降到28GB优化器状态还是FP32的56GB总量84GB。虽然一张80GB的卡仍然紧张但已经从一个完全不可能的情况变成了一个勉强有希望的情况。分布式训练则解决的是“一张卡装不下就多张卡一起来”的问题。它把模型、数据或者计算切到多张GPU上让每张卡只负责一部分最后通过通信把结果合并。还是上面那个场景如果我有4张80GB的卡总显存就有320GB哪怕不引入混合精度也能跑起來。但注意这里有个关键的交织点如果你的分布式方案采用的是数据并行也就是每张卡上都放一份完整的模型副本那么每张卡依然要承受“参数加梯度加优化器状态”的全部压力。这时如果不用混合精度每张卡至少要112GB4张80GB的卡照样全炸。这就是我为什么说混合精度和分布式必须合在一起学分布式决定了你要切几刀混合精度决定了每一刀下去之后那份肉有多大。两个技术一起上才能用合理的成本把一个大模型真的训起来。所以这篇文章的讲述顺序也按这个逻辑来先讲混合精度搞清楚单卡上的显存怎么省、精度怎么保再讲分布式搞清楚多卡之间怎么分工、怎么协作最后把两者组合成实际训练方案并聊聊我在真实训练中踩过的坑。2. 混合精度训练的底层原理FP16计算、FP32主权重与Loss Scaling混合精度训练这个概念听上去高深实际拆开就三件事参数和梯度用半精度来存、计算也发生在半精度下、但优化器更新权重时用一份额外的FP32副本来保证精度。很多人第一次写混合精度代码以为只要把模型参数改成.half()就完事了结果训练直接发散或者精度崩掉。这是因为没有真正理解混合精度的核心设计主权重Master Weights。2.1 为什么不能全程用半精度FP16这种格式用1位符号位、5位指数位、10位尾数位来表示一个浮点数。它的数值范围大约在6×10^-8到65504之间。听着范围不小但对训练来说有两个致命问题。第一个问题是数值溢出。FP16能表示的最大值是65504而模型训练过程中某些中间结果、损失值或者梯度一旦超过这个数就会变成无穷大inf。尤其在模型比较大的时候早期训练阶段loss值可能非常高或者某些层的梯度异常大一个inf产生之后就会在整个反向传播过程中被不断放大最后所有参数都变成NaN。我见过好几次训练跑着跑着loss突然变成NaN查到最后都是FP16下某个梯度爆了。第二个问题更隐蔽是精度下溢。在反向传播过程中很多梯度是非常小的数动辄在1e-5甚至1e-7级别。FP16能表示的最小正规格化数大约是6×10^-8一旦梯度小于这个数就直接被舍入成0。等于说这一层的参数在更新时根本没有获得任何学习信号。小梯度直接归零深层网络参数学不动这就是为什么纯FP16训练几乎不可能收敛。所以混合精度的核心设计来了计算图用FP16跑但每个参数同时保留一个FP32的副本每次更新权重时是在FP32副本上完成的然后再转成FP16用于前向传播。这个FP32副本就是主权重。有了它哪怕FP16在反向传播时损失了一点细节更新本身依然发生在高精度的主权重上不会因为累计误差而跑偏。2.2 Loss Scaling是怎么工作的解决了主权重的问题还剩梯度下溢的问题。既然反向传播算出的小梯度在FP16下会归零能不能先把梯度放大让它落进FP16能表示的范围内等用完之后再缩小回去这就是Loss Scaling的核心思路。做法其实很朴素在反向传播之前把loss乘以一个缩放因子比如1024或4096loss变大了反向传播算出来的梯度也等比例变大从而落在FP16的安全范围内。等梯度计算完毕在更新权重之前把梯度除以这个缩放因子恢复其真实数值再交给优化器使用。实际工程中缩放因子有两种配置方式。静态缩放就是固定一个值比如128或者1024全程不变。好处是简单坏处是对不同模型、不同训练阶段不够灵活——loss已经很小了梯度都缩到正常范围了你还在死命放大反而又搞出inf来。动态缩放则是自适应的训练开始时设一个较大的scale如果连续N步没有出现inf或NaN就尝试把scale调大一些如果一旦出现inf/NaN就先跳过这一步的参数更新再把scale调小。这样能自动找到一个既不会下溢又不会频繁上溢的平衡点目前绝大多数框架默认用的都是动态Loss Scaling。我们在实际训练中通常的做法是初始化scale设为2的16次方65536每2000步检查一次如果2000步内没出现inf或NaN就把scale翻倍上限一般设到2的24次方左右一旦出现overflow就把当前batch跳过并回退scale。听起来复杂但PyTorch的torch.cuda.amp.GradScaler这些都帮你封装好了你只需要理解它背后在做什么才好在它工作异常时知道去哪里排查。2.3 混合精度训练过程中的数据流把上面这些东西串起来混合精度训练的一个标准迭代长这样把模型参数以FP32的主权重形式存放在显存中同时保留一份FP16版本用于前向和反向。将输入转为FP16执行前向传播得到FP16或FP32的loss。调用scaler.scale(loss)把loss放大后再执行反向传播得到FP16的梯度。这里的核心是缩放后的梯度。调用scaler.unscale_()把梯度缩小回真实值只对本地梯度进行操作不涉及跨卡通信时更安全。执行梯度裁剪时也必须在unscale之后再做。optimizer.step()优化器读取FP32主权重和unscale后的FP32梯度完成更新。把更新后的FP32主权重再转成FP16供下一轮前向使用。这套流程中最容易出错的组合是梯度裁剪和Loss Scaling的顺序。很多人只记得放大loss结果梯度裁剪的时候用的还是缩放过的梯度clip阈值完全失效等于白裁。正确顺序是先unscale再clip再step。类似这种因为“顺序错了导致训练看起来在跑、但实际已经不收敛”的情况在混合精度实践里非常常见。我见过同事排查了大半天loss为什么不降最后发现就是梯度裁剪位置放错了。3. FP16的接班人BF16凭什么成了大模型训练的新宠如果你去GitHub上看现在的大模型训练仓库会发现越来越多项目默认用BF16而不是FP16。这个变化背后有个很实际的原因理解了它你就能明白混合精度这个领域其实也一直在演进。BF16的全称是Brain Floating Point最早来自Google的TPU设计后来被英伟达、AMD等厂商广泛支持。它同样是16位但位分配完全不同1位符号位8位指数位7位尾数位。这和FP32的位分配几乎一致只是尾数从23位砍到了7位。这意味着什么意味着BF16的数值范围和FP32完全一样都是大约1.18×10^-38到3.4×10^38。也就是说BF16根本没有FP16那个高溢出、低下溢的问题。这一点对训练来说极其关键。前面提到FP16下梯度小于6×10^-8就会归零但在BF16下最小正规格化数大约是9.2×10^-41正常训练中几乎不可能出现比这还小的梯度。换句话说BF16天然规避了梯度下溢问题连Loss Scaling在很多情况下都可以不用了。很多大模型训练框架在BF16模式下直接把动态缩放功能关闭整个流程都简单了。那BF16没有代价吗当然有。它的尾数只有7位比FP16的10位还少单个数能表达的精度更差舍入误差更大。做个类比FP16像是用两页纸记笔记BF16也是两页纸但很多页都用来记“范围”了留给“细节”的反而更少。7位尾数意味着一个像3.14159这样的数存到BF16里可能变成3.14更细的小数部分直接被丢弃。这个精度损失在训练中会不会造成问题实际经验表明大模型训练过程中这个误差通常可以接受原因在于每个参数在训练中会被海量梯度更新随机方向的舍入误差在统计上会相互抵消一部分。而且我们还有FP32主权重兜底梯度更新发生在FP32精度上BF16只负责前向和反向的计算过程。所以BF16的工程实践往往比FP16更省心不用调Loss Scaling没有那么多溢出问题训练更加稳定。但有一个场景必须小心小模型或者对精度极其敏感的任务BF16的尾数不足可能导致问题。经验上如果模型只有几千万参数训练一个需要高精度的回归模型BF16可能会跑出FP16完全没有的精度问题。这种情况下要么继续用FP16加动态缩放要么干脆用FP32训练不要盲目追新。我自己在训练一个亿级参数的模型时遇到过一种现象用BF16训练loss曲线看起来比FP16平滑很多但最终验证集精度反而比FP16低了一点点。原因大概率就是尾数精度损失的累积效应。所以现在很多仓库的做法是提供开关FP16和BF16都支持让用户根据自己模型的实际表现来选择。选哪个不是看paper吹什么而是看你的模型在验证集上的最终表现。4. 分布式训练的三条路线数据并行、张量并行、流水线并行混合精度解决的是单卡上的显存和精度问题但单卡再省物理上限就摆在那里。真要跑百亿千亿参数的模型就必须上分布式训练。分布式训练的口号是“多卡一起干”但“一起干”的拆分方式完全不同最常见的三条路线是数据并行、张量并行、流水线并行。我一个个说清楚它们各自解决什么问题、代价是什么。4.1 数据并行每张卡都装一份完整模型分工处理不同数据数据并行最容易理解假设有8张卡把训练数据切成8份每张卡持有一份完整的模型副本各自跑各自的batch算各自的梯度然后通过通信把所有卡的梯度加到一起求平均再用平均后的梯度更新每张卡上的模型。之后所有卡的模型参数保持一致继续下一轮。数据并行的通信模式核心是AllReduce。简单说就是所有卡把梯度拿出来做一个全局求和然后广播回去最终每张卡拿到的梯度都一样。实际工程中用的多是Ring AllReduce把N张卡连成一个环通信量不随卡数线性增长所以8张、32张、甚至128张卡都能扩展。数据并行的好处是简单、直观所有框架都支持改造代价极小。但它有一个显眼的问题每张卡都要完整保存一份模型副本前面说的参数加梯度加优化器状态的显存压力一张卡上也跑不掉。所以数据并行通常会和混合精度一起用混合精度先把单卡压力减半数据并行再把数据量平摊到多卡。在我实际使用中如果模型单卡能塞下比如7B模型配合混合精度和梯度检查点单张80GB勉强能跑数据并行是最舒服的方案因为它不需要改动模型结构只是在训练循环外层套一个DDPDistributed Data Parallel包装。4.2 张量并行把单个层切成几块多张卡联合算一个层模型大到单卡塞不下的时候数据并行就失效了。这时候必须从模型结构本身下手张量并行就是这条路。张量并行解决的是“单个Transformer层太大一张卡装不下”的问题。它的做法是把一个层的权重矩阵按行或按列切分到多张卡上每张卡只保存这个层的一部分做矩阵乘法时各算各的最后通过通信把部分结果拼接起来。比如一个4096×4096的权重矩阵用8卡切分每张卡只需要保存512列。张量并行的通信开销特别大因为每算完一个矩阵乘法就要做一次全量的AllReduce来同步结果。在大模型里一个Transformer层里可能有好几个Linear层和Attention层每一层都带通信所以张量并行通常只在节点内部使用走NVLINK这类超高速卡间互联跨节点用它的通信延迟会让人崩溃。实际使用中张量并行度一般不超过单节点内的卡数常见的是2、4、8。4.3 流水线并行把模型切成一段段各张卡接力算流水线并行的思路完全不同不切层内部分而是按层切。假设模型有40层Transformer8张卡那就每张卡分到5层。数据流是从GPU0算完1到5层把中间结果传给GPU1算6到10层一路接力到最后一张卡。每个GPU只负责一部分层显存压力自然下来了。流水线并行的问题是“流水线气泡”。如果数据只能串行走完整个模型那同一时刻只有一张卡在计算其他卡全在等硬件利用率惨不忍睹。为了解决这个问题工程上采用“微批次”技术把一个大的batch切成多个micro-batch让它们像流水线一样依次进入模型。GPU0在算第1个micro-batch的1到5层时算完就传给GPU1马上接着算第2个micro-batch的1到5层这样多张卡就能同时忙碌起来。不过即便如此流水线必然存在气泡也就是某张卡没有活干的空闲时间所以它更适合超大模型因为显存放不下的问题比算力浪费更致命。4.4 三条路线怎么选实际训练中这三种并行通常是组合使用的而不是三选一。我的经验是这样判断的单卡能塞下直接数据并行加混合精度最早跑通最简单。单卡塞不下但单节点多卡能塞下张量并行配合数据并行每张卡持有模型的一个分片多节点之间用数据并行扩展吞吐。模型大到单节点都放不下张量并行加流水线并行再加数据并行参考Megatron-LM的“3D并行”设计这一般是千亿参数以上模型的标配。这个选择逻辑可以套一个简单判断先看模型能不能单卡装下决定要不要上模型并行再看训练速度够不够决定要不要加数据并行。顺序不要搞反否则很容易给一个本该简单的训练任务硬套一个复杂并行方案通信开销比省下的显存还贵得不偿失。5. 显存都去哪儿了混合精度下的显存账单与优化器秘密很多初学者对分布式训练有个错误理解以为把模型放到多张卡上显存压力就自动均摊了。这其实要看你用哪种并行。如果用的只是数据并行每张卡上的显存占用量和单卡训练几乎一样因为你保存的是一份完整的模型副本。所以搞清楚显存到底花在了哪里是理解分布式训练后续所有优化的前提。5.1 模型训练期的显存四大开销训练过程中显存主要被四类东西占用模型参数、梯度、优化器状态、激活值。激活值指前向传播中每一层的中间输出它们需要保留到反向传播时用来计算梯度。这四块中参数、梯度、优化器状态的大小和参数量成正比一目了然。激活值则和序列长度、batch size、模型层数都有关系在大模型训练里往往占大头所以才会衍生出梯度检查点技术用“反向传播时重新计算一部分激活值”来换显存。以7B模型为例如果只算参数、梯度、优化器状态三块混合精度下的显存账单大概是这样的参数用FP16存储7GB梯度和参数同尺寸也是FP16的7GB但优化器状态是按FP32算的Adam的m和v各占28GB再加上主权重FP32一份又是28GB优化器相关总共84GB。三块相加约98GB已经超过了单张80GB的物理上限这还没有算激活值、中间缓冲和通信缓冲区。所以光靠混合精度7B模型依然不能舒服地跑在单卡上。这时候分布式训练的价值就体现出来了。但具体怎么切很有讲究。如果把这98GB分摊到8张卡上每张卡只要大约12.25GB这就很舒服。关键问题在于98GB里占比最大的是优化器状态总共84GB如果不针对优化器状态做文章只切参数和梯度节省效果非常有限。5.2 Adam优化器的“隐藏账单”Adam是当前大模型训练的事实标准优化器但它也是最占显存的优化器之一。它除了保存一份FP32的模型参数主权重以外还要为每个参数保存两个状态一阶动量m和二阶动量v全都是FP32。也就是说Adam的显存开销是参数量乘以12字节光优化器部分就是模型参数量的6倍对比SGD只保存一个动量Adam昂贵得多。正因为Adam如此吃显存混合精度加分布式训练才必须仔细设计优化器状态的存放位置。如果把每个参数在FP32下的主权重和m、v都保存在本地显存那么数据并行下每张卡依然要承担完整的84GB优化器状态这显然不可接受。于是就有了下一节要讲的ZeRO——它的核心洞察正是“我看到优化器状态占了训练显存的大头那为什么不把优化器状态分片到多张卡上呢”。5.3 显存账本与分布式训练选型的联动我来给你一个可以直接抄的判断流程。当你拿到一个模型准备训练时先用公式粗算一下参数量为P单位是B也就是十亿混合精度FP16参数加FP16梯度加FP32的m、v和主权重下基础显存约等于P乘以14字节单位是GB。7B模型就是7乘以14等于98GB。再加上激活值你大概能知道自己需要几张卡。然后用这个数字决定要不要上ZeRO、上几阶段。比如7B模型单卡80GB放不下基础显存和激活值那就得上模型并行或者ZeRO。如果模型只有1B到3B混合精度后约14到42GB单卡勉强能跑这时候数据并行最常见也不需要用ZeRO去分片优化器状态。这个判断看起来简单但我见过太多人一上来就抄别人的超大模型训练配置明明3B模型还用ZeRO-3加32卡结果通信开销比实际节省的显存还大训练速度反而更慢。6. ZeRO让显存账单从“一人扛”变成“大家摊”的关键设计ZeRO是DeepSpeed框架提出的一套显存优化方案全称是Zero Redundancy Optimizer。它解决的核心问题就是数据并行下每张卡都重复保存了一整套参数、梯度和优化器状态这里有大量的冗余。ZeRO的思想是打破这种冗余把模型状态分片到所有卡上让每张卡只保存一份全局状态的1/N。6.1 ZeRO的三个阶段从砍优化器状态到砍参数ZeRO分三个阶段对应三种不同的取舍。ZeRO-1只分片优化器状态。参数和梯度仍然每张卡各保存一份但Adam中的主权重、m和v只保存1/N份。按照前面7B模型的账本优化器状态84GB在8卡下变成84除以8约等于10.5GB每卡加上参数和梯度各7GB总共约25GB。这个阶段的好处是通信开销基本没有增加因为梯度同步在数据并行里本来就要做AllReduce你只是顺便把梯度分片之后只更新自己那块优化器状态而已。实际使用中ZeRO-1是性价比最高的起点。ZeRO-2在ZeRO-1的基础上把梯度也分片了。每张卡只保存自己需要负责更新的那部分梯度。这带来一个额外的需求梯度在反向传播计算出完整值之后需要通过Reduce-Scatter操作把各卡梯度按分片方式求平均每张卡只保留自己要的那部分。这个操作会引入一些额外的通信但通常会比一遍完整全量梯度同步省一部分开销。显存上梯度7GB每卡在8卡下降到不到1GB总显存又降了一个台阶。ZeRO-3更进一步参数也分片。前向传播时每用到一个层需要先通过AllGather把所有卡上的该层参数收集到本地反向传播时同样要收集参数来计算梯度。这导致通信开销大幅增加因为每一层都要做一次全量的参数收集。ZeRO-3的显存节省最明显基础三件套降到每卡大约98除以8等于12.25GB的水平但通信开销也最高。6.2 通信开销与显存的博弈选ZeRO几阶段本质上是在“显存不够”和“通信太慢”之间找平衡。如果显存还没见底优先用ZeRO-1因为它几乎不增加通信负担。如果显存真的吃紧ZeRO-2能有效缓解代价是每次梯度同步多一步Reduce-Scatter。ZeRO-3则尽量少用它虽然省显存最猛但对节点内卡间通信带宽要求极高跨节点时网络延迟会严重影响训练效率。我在公司实际训练一个13B模型时用的方案是ZeRO-2加混合精度配合8张A100。显存占用大约稳定在65到72GB之间训练速度基本能接受。后来为了在同样的卡上塞更大的batch尝试切到ZeRO-3结果显存确实降到了40GB左右但每个step的时间翻了将近一倍最后权衡下来还是换回了ZeRO-2。这就是典型的“省了显存、亏了速度”不要只看显存数字好看。6.3 混合精度和ZeRO的配合细节混合精度和ZeRO配合时有一个容易忽略的细节值得专门拎出来说。ZeRO分片优化器状态时分出去的重头就是你那把FP32主权重、m和v。如果没有混合精度参数、梯度和优化器状态全部是FP32分片后通信数据量更大引入FP16参数和梯度之后虽然通信数据量变小了但分片后的FP32主权重和优化器状态依然要给每一层更新使用。另外ZeRO对梯度通信也做了特殊处理。以ZeRO-2为例反向传播过程不是等所有层的梯度全部算完再一次同步而是每算完一部分就启动Reduce-Scatter边算边通信让通信和计算重叠起来。这个设计非常关键否则梯度同步的等待时间会让训练慢得难以忍受。理解这一点你就能明白为什么建议用DeepSpeed等成熟框架来实现ZeRO而不是自己手写一套——看起来不难但通信与计算的流水线重叠才是真正的工程门槛。7. 我在实际训练中踩过的几个坑混合精度与分布式组合实战最后这部分我梳理几个真实项目中踩过的坑每一个都是我或身边同事花了很长时间排查才解决的。这些问题在官方文档里很难找到现成答案但一旦碰到非常容易把人心态搞崩。7.1 梯度未反缩放就通信和裁剪混合精度训练时loss被缩放反向传播出来的梯度也是缩放过的。在DDP或者ZeRO中梯度最终要在卡之间做AllReduce这时有一个必须注意的顺序问题如果先做了AllReduce梯度的通信再进行unscale等于把缩放后的梯度在卡间做了平均这个缩放因子被混进了通信结果里。表面上看梯度还是同步的但更新时的数值就会被scale污染导致学习率实际变成“学习率除以scale”。而不同的scale会增加排查难度训练loss可能一直不稳定。正确做法是先对本地梯度unscale再做梯度裁剪再进行梯度通信。DeepSpeed的ZeRO会在内部自动处理好这个顺序但使用手写DDP或一些轻量框架时这个顺序很容易被忽略。所以我的习惯是只要发现训练曲线在混合精度下异常抖动先检查unscale和通信的顺序成本最低。7.2 没设好随机种子各卡模型初始化不一致分布式训练下所有卡的模型初始参数必须严格一致。如果初始化不同梯度同步就没有意义训练会直接乱掉。常见的错误是在写代码时没有为每张卡设置不同的随机种子或只设置了主进程的种子导致每张卡的数据加载顺序和模型初始化参数都不一样。实际训练中每个进程的分工是模型初始化用固定的全局种子保证一致性数据加载时在固定种子基础上加上进程rank偏移保证每张卡拿到不同的数据。这两套种子必须分开设置。我在一次多机训练里吃过亏两台机器的随机种子一致但数据增强部分用了全局时间戳作为种子结果两个节点上数据分布完全不一致训练指标奇差无比。排查了很久最后才发现是数据加载种子的设置问题。7.3 混合精度下的loss spike先查scale再查数据大模型训练过程中loss突然变成NaN或者炸到极大值是每个训练者都躲不掉的噩梦。在混合精度下你第一个要查的东西永远是scale的状态动态loss scaling有没有在某个step突然触发了回退是不是某个batch的数据里混进了异常值导致前向或反向出现inf一个我印象深刻的case训练一个多模态模型每过几百步loss就spike一次。检查数据发现某几个样本在预处理时出现了空白张量前向计算时产生了异常值导致梯度过大直接爆掉FP16的表示范围。换到BF16之后这个问题明显缓解了因为BF16的数值范围大得多。所以如果你用FP16频繁遇到spike不妨先排查数据中是否有极端异常样本同时考虑切换到BF16试试。7.4 通信后端选不对多机训练慢得离谱分布式训练的通信后端也很关键。单机多卡时PyTorch DDP的NCCL后端能利用卡间高速互联效率极高。但有些场景下比如容器环境里NCCL初始化失败、或者某些云计算平台上NCCL被限制有人会退回GLOO后端结果训练速度断崖式下降。因为GLOO是为CPU通信设计的通用后端在GPU场景下效率远不如NCCL。我的排查经验是遇到多卡训练非常慢先用nvidia-smi确认GPU利用率。如果GPU利用率低大概率是通信在等着如果GPU利用率高但速度慢可能是单卡计算本身或数据加载瓶颈这时换个后端不会有任何改善。集群环境下还需要额外确认NCCL的网卡绑定和网络拓扑避免跨节点通信走了低带宽通道。7.5 显存看似够用但一跑就OOM还有一种很常见的坑计算下来显存明明够一跑就OOM。这种情况通常不是模型参数超了而是隐藏的显存开销没算进去。混合精度训练中PyTorch会为通信引擎分配额外的梯度缓冲区数据并行下梯度AllReduce之前每张卡要把梯度搬到连续的通信缓冲区中这会额外占用显存。ZeRO中也要为分区通信保留缓冲区。解决思路是要么减小batch size要么开梯度检查点gradient checkpointing来舍去激活值要么调小通信缓冲区要么从ZeRO-1换到ZeRO-2减少一部分冗余。不要一上来就认为是卡数不够先看看显存具体被谁吃掉。可以用torch.cuda.max_memory_allocated()和torch.cuda.memory_reserved()打点统计通常能发现被忽略的显存大户。8. 混合精度加分布式的最终配置建议纸上谈兵这么多最后给一套可以直接参考的选型组合。模型规模不同最优配置差异很大。1B到3B规模的模型混合精度加单机多卡数据并行DDP就够了。单张卡在用混合精度后基本能装下数据并行用来加速提升吞吐。不要上ZeRO更不要上张量并行通信开销大于收益。7B到13B规模的模型推荐混合精度加ZeRO-2加数据并行。显存基本能稳在80GB以内训练速度也可控。如果激活值太大就开梯度检查点不要贸然上ZeRO-3。如果有条件用BF16优先用BF16省心很多。30B到70B规模的模型单机多卡加ZeRO-3是底线最好配合张量并行。此时通信开销高需要InfiniBand或高速RDMA网络否则训练效率会被通信托垮。这个规模下混合精度和分布式已经不是选不选的问题而是必须共同设计的系统工程。上面这些配置不是绝对标准但可以作为起步参考。训练是一场实验科学最优配置不是算出来的而是在具体硬件和模型下调试出来的。理解底层原理知道每个配置在调什么才能在训练出问题时找到方向。混合精度训练和分布式训练本质上都是在和硬件资源博弈。混合精度博弈的是单卡的显存和精度分布式博弈的是多卡之间的通信和分工。两者结合才撑起了当前大模型训练的基本盘。学完这篇如果以后再有人问你为什么大模型训练要用BF16而不用FP16、为什么数据并行还会显存不够、为什么ZeRO-2比ZeRO-3快你应该都能给出一个清晰的技术回答了。
返回列表