ARTICLE DETAIL

资讯详情

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

大模型训练显存估算与混合精度实战:从OOM到跑通7B模型

大模型训练显存估算与混合精度实战:从OOM到跑通7B模型 大模型训练绕不开两个硬骨头显存不够和精度怎么选。我见过太多团队在单卡上跑7B模型刚把batch size调到8就OOM然后开始盲目换卡、换框架最后发现是优化器状态没算对。这篇就把显存估计和混合精度训练这两件事拆开揉碎讲清楚从公式推导到代码实操从FP16的坑到BF16的甜再到INT8量化推理的边界全部基于实际项目经验。不管你是刚接触大模型训练的新手还是已经调过几轮参数的老手这里面的估算方法和精度选择逻辑都能直接拿去用。1. 显存到底被谁吃掉了逐项拆解与估算公式很多人估显存就是“参数量乘以4再乘个系数”这种粗估在7B以下还能凑合到了13B以上误差能到几十GB。要算准必须把显存消耗拆成四块模型参数、梯度、优化器状态、激活值。前三个是静态开销跟batch size无关激活值是动态开销随batch size和序列长度线性增长。1.1 模型参数与梯度的显存占用模型参数就是权重矩阵每个参数占多少字节取决于精度。FP32下每个参数4字节FP16和BF16下每个参数2字节。梯度跟参数一一对应所以梯度占用的字节数跟参数精度一致。举个例子一个7B模型用FP16训练参数占用7B × 2 14GB梯度同样14GB加起来28GB。如果用FP32训练参数28GB梯度28GB直接56GB单张80GB的卡只剩24GB给优化器状态和激活值基本跑不动。这里有个容易忽略的点混合精度训练时模型会保留一份FP32的master weight。也就是说即使你用FP16做前向和反向优化器更新的还是FP32的那份参数。所以实际参数显存是FP16副本加上FP32主副本总共7B × (24) 42GB。梯度也存在FP16和FP32两份又是42GB。这一下就84GB了单卡80GB直接爆掉。这也是为什么混合精度训练通常要配合ZeRO或者模型并行。1.2 优化器状态的显存黑洞优化器状态是大头尤其Adam系列。Adam为每个参数维护两个状态一阶矩估计动量和二阶矩估计方差。如果优化器状态用FP32存储每个参数需要448字节。加上参数本身FP32的4字节和梯度的4字节Adam的总开销是每个参数16字节。7B模型就是7B × 16 112GB这还没算激活值。AdamW跟Adam的区别在于权重衰减的实现方式显存开销一样。SGD就省多了没有状态每个参数只要参数4字节加梯度4字节共8字节。但SGD收敛慢大模型训练基本都用AdamW。所以实际估算时优化器状态按8字节每参数算加上参数和梯度的FP32副本总共16字节每参数。如果用ZeRO-1把优化器状态分片到N张卡上每张卡的优化器状态显存变成原来的1/N。ZeRO-2再分片梯度ZeRO-3连参数都分片。这是后话先记住单卡全量训练的公式。1.3 激活值随batch size线性增长的变量激活值是前向传播过程中每一层的输出反向传播时需要用来计算梯度。激活值的大小跟batch size、序列长度、隐藏层维度、层数都相关。粗略估算公式是激活值显存 ≈ batch_size × seq_len × hidden_size × num_layers × 系数。系数取决于具体实现PyTorch的checkpoint机制能大幅降低激活值但会增加计算量。以7B模型为例hidden_size4096num_layers32seq_len2048batch_size1。粗略算1 × 2048 × 4096 × 32 × 2字节 ≈ 0.5GB。但实际因为注意力矩阵、中间激活等因素会到2-4GB。batch_size翻倍激活值翻倍。所以batch size从1调到8激活值从3GB变成24GB这就是OOM的直接原因。注意激活值估算没有精确公式不同框架、不同注意力实现如FlashAttention差异很大。建议用torch.cuda.memory_allocated()实测或者用框架自带的显存估算工具。1.4 完整估算公式与实战案例把上面四块加起来单卡全量训练AdamW的显存估算公式总显存 参数量 × (2 4) 梯度 × (2 4) 优化器状态 × 8 激活值 参数量 × 16 激活值等等这里参数和梯度各算了FP16和FP32两份所以是246字节每参数参数加梯度共12字节优化器8字节合计20字节每参数。但很多框架实现不同有的只保留FP32 master weight梯度只有FP16一份。保守估算按20字节每参数。7B模型7B × 20 140GB加上激活值至少4GB总共144GB。单卡80GB肯定不够需要至少2张卡做ZeRO-2或者3张卡做ZeRO-3。13B模型13B × 20 260GB加激活值8GB268GB。至少4张80GB卡。70B模型70B × 20 1400GB加激活值40GB1440GB。至少18张80GB卡实际要20张以上留余量。模型规模参数量静态显存(20字节/参数)激活值(bs1)总显存80GB卡数量7B7B140GB4GB144GB213B13B260GB8GB268GB470B70B1400GB40GB1440GB18这个表是保守估算实际用ZeRO-3加CPU offload能进一步降低。但估算逻辑要清楚先算静态再算动态最后留20%余量。2. FP16与BF16的精度博弈为什么BF16成了大模型标配混合精度训练的核心是用低精度做前向和反向用高精度做参数更新。FP16和BF16都是16位但动态范围天差地别。FP16有10位尾数、5位指数动态范围约6e-5到65504。BF16有7位尾数、8位指数动态范围跟FP32一样约1e-38到3e38。尾数决定精度指数决定范围。2.1 FP16的溢出与下溢问题FP16的指数位只有5位能表示的最大值65504。大模型训练中梯度值很容易超过这个数尤其是深层网络。一旦溢出梯度变成Inf参数更新后变成NaN训练直接崩。下溢也一样梯度小于6e-5就变成0参数不更新模型学不动。解决FP16溢出的标准做法是损失缩放Loss Scaling。原理很简单反向传播前把loss乘以一个大的缩放因子比如1024梯度也跟着放大避免下溢。更新参数前再除以这个因子恢复原值。动态损失缩放会根据梯度是否溢出自动调整因子溢出就减小正常就增大。但损失缩放不是万能的。如果梯度本身动态范围很大缩放因子很难兼顾。而且每次溢出都要跳过这一步更新训练效率受影响。我实测过7B模型用FP16加动态损失缩放训练初期每几百步就溢出一次虽然能恢复但浪费了不少计算。2.2 BF16的天然优势与硬件门槛BF16的指数位跟FP32一样动态范围完全覆盖不需要损失缩放。梯度再大也不会溢出再小也不会下溢。代价是尾数只有7位精度比FP16低。但大模型训练对精度没那么敏感参数更新时那点误差被优化器的动量平滑掉了。BF16的硬件门槛是需要Ampere架构以上的GPU比如A100、A30、RTX 30系以上。V100不支持BF16只能用FP16。所以选BF16之前先确认卡的支持情况。用torch.cuda.is_bf16_supported()可以查。提示BF16训练时优化器状态和master weight仍然用FP32保证参数更新的精度。前向和反向用BF16速度跟FP16差不多但稳定性好太多。2.3 实测对比FP16 vs BF16在7B模型上的表现我在A100上跑过7B模型的对比实验同样的数据、同样的超参只换精度。FP16加动态损失缩放训练到1000步时loss曲线有几次明显抖动对应损失缩放因子调整。BF16全程平滑没有溢出。最终收敛后的loss值BF16比FP16低0.02左右差异不大但稳定。速度方面两者几乎一样因为A100对FP16和BF16的Tensor Core吞吐相同。显存占用也相同都是2字节每参数。所以只要卡支持无脑选BF16。FP16只在V100等老卡上不得不用。对比项FP16BF16指数位5位8位尾数位10位7位动态范围6e-5 ~ 655041e-38 ~ 3e38损失缩放必须不需要硬件要求多数GPUAmpere以上训练稳定性一般可能溢出好速度快快2.4 混合精度训练的代码实操PyTorch的torch.cuda.amp是标准做法。下面是一个训练循环的骨架import torch from torch.cuda.amp import autocast, GradScaler model MyModel().cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) scaler GradScaler() # FP16需要BF16不需要 for batch in dataloader: inputs, labels batch optimizer.zero_grad() with autocast(dtypetorch.bfloat16): # 或torch.float16 outputs model(inputs) loss loss_fn(outputs, labels) scaler.scale(loss).backward() # FP16 # loss.backward() # BF16直接用这个 scaler.step(optimizer) scaler.update()BF16时把autocast(dtypetorch.bfloat16)去掉scaler相关调用。注意autocast只影响前向反向的梯度计算自动用对应精度。优化器更新时PyTorch会自动把梯度转成FP32再更新master weight。注意autocast区域内的操作要检查是否支持低精度。有些自定义算子可能不支持会回退到FP32影响速度。用torch.autocast的enabled参数可以临时关闭。3. 显存优化实战从OOM到跑通7B模型的完整过程理论算完上手跑7B模型还是OOM。我记录了一次完整的排查过程从单卡80GB开始一步步调到能跑。3.1 第一次尝试单卡全量训练直接爆7B模型FP16混合精度AdamWbatch_size1seq_len2048。按公式算静态显存140GB单卡80GB肯定不够。但我想试试PyTorch的torch.cuda.amp能不能省点。结果加载模型就占了14GBFP16参数优化器初始化又占了56GBFP32参数优化器状态还没开始训练就70GB了。前向传播一跑激活值加上去直接OOM。这里有个细节PyTorch加载模型时如果直接.cuda()参数是FP32的。用.half()转FP16会省一半但优化器状态还是FP32。所以加载完模型先转FP16再初始化优化器能省14GB。3.2 第二次尝试梯度累积加小batchbatch_size降到1已经最小了只能减seq_len。从2048降到1024激活值减半。但静态显存没变还是140GB。单卡无解必须上多卡。3.3 第三次尝试ZeRO-2分片优化器状态和梯度用DeepSpeed的ZeRO-2把优化器状态和梯度分片到2张卡上。每张卡的静态显存变成参数14GBFP16 参数28GBFP32 master 梯度14GBFP16 梯度28GBFP32/ 2 优化器状态56GB / 2。算下来每张卡约142814142898GB还是超80GB。等等ZeRO-2的分片逻辑是优化器状态和梯度分片参数不分片。所以每张卡都有完整的FP32参数28GB和FP16参数14GB梯度只存自己那份。重新算FP16参数14GB FP32参数28GB FP16梯度14GB/2 FP32梯度28GB/2 优化器状态56GB/2 142871428 91GB。还是超。3.4 第四次尝试ZeRO-3加CPU offloadZeRO-3把参数也分片每张卡只存1/2的参数。FP16参数7GB FP32参数14GB 梯度分片 优化器状态分片。算下来每张卡约7147142870GB加上激活值4GB74GB勉强能跑。但ZeRO-3通信开销大训练速度降了约30%。最后用ZeRO-3加CPU offload优化器状态每张卡降到50GB左右速度降了40%但能跑通。如果换成4张卡ZeRO-2就够了速度损失小很多。方案卡数每卡显存速度损失可行性单卡全量1140GB0不可行ZeRO-2291GB10%不可行ZeRO-3274GB30%勉强ZeRO-3offload250GB40%可行ZeRO-2445GB10%推荐这个排查过程说明显存估算要准优化方案要按卡数选。2张卡优先ZeRO-34张卡ZeRO-2更划算。4. INT8量化推理加速与训练精度的边界INT8在推理场景很常见训练场景用得少。原因很简单INT8只有8位精度损失太大训练时梯度更新会不稳定。但推理时不需要反向传播INT8能把显存和计算量都降一半。4.1 INT8量化的基本原理INT8把FP16的权重和激活值映射到8位整数。映射公式int8_value round(fp_value / scale) zero_point。scale是缩放因子zero_point是零点偏移。反量化时逆运算。关键是找合适的scale让浮点值的动态范围刚好覆盖INT8的-128到127。训练后量化PTQ直接用校准数据算scale简单但精度损失大。量化感知训练QAT在训练时模拟量化误差精度好但需要重新训练。大模型通常用PTQ加GPTQ或AWQ等高级算法能在4位甚至3位下保持精度。4.2 INT8与BF16的模型区别BF16模型是训练时的精度参数和激活都是16位。INT8模型是推理时的精度参数和激活都是8位。BF16模型可以直接训练INT8模型只能推理。BF16的显存是FP32的一半INT8是FP32的四分之一。速度上INT8的矩阵乘法在支持INT8 Tensor Core的GPU上比BF16快一倍左右。但INT8的精度损失在生成任务上很明显。我实测过同一个7B模型BF16推理和INT8推理在长文本生成上INT8会出现重复、逻辑断裂。短文本问答差异不大。所以INT8适合对精度要求不高的场景比如分类、抽取。对比项BF16INT8位数168显存2字节/参数1字节/参数训练支持是否需QAT推理速度快更快精度损失无明显适用场景训练推理推理4.3 实际部署中的选择逻辑训练阶段用BF16推理阶段看场景。如果显存够BF16推理最稳。如果显存紧张INT8能省一半显存但要做好精度下降的准备。折中方案是FP16推理显存跟BF16一样但老卡不支持BF16时用FP16。提示INT8量化后的模型推理时要注意校准数据的分布。校准数据跟实际输入分布差太多量化误差会放大。建议用真实业务数据做校准。5. 那些文档不会告诉你的实操细节显存估算和混合精度训练的理论不难难的是实操中的各种意外。我整理了几个踩过的坑和对应的解法。5.1 激活值检查点的取舍PyTorch的torch.utils.checkpoint能把激活值显存降一个数量级代价是反向传播时重新计算前向。7B模型用checkpoint激活值从4GB降到0.5GB但训练速度慢20%。如果显存够别用checkpoint如果OOM这是最直接的救命稻草。用的时候注意checkpoint只对nn.Sequential或自定义的forward有效注意力层要单独处理。而且checkpoint后的层参数梯度计算会受影响要确保requires_grad设置正确。5.2 梯度累积与batch size的等效关系显存不够时用梯度累积模拟大batch。比如batch_size1累积4步等效batch_size4。但要注意BatchNorm层在累积时统计量会偏大模型通常用LayerNorm没这个问题。另外学习率要按等效batch size调整线性缩放规则batch size翻倍学习率翻倍。梯度累积的代码很简单accum_steps 4 for i, batch in enumerate(dataloader): loss model(batch) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()注意loss要除以累积步数否则梯度会累积成4倍。5.3 多卡训练时的显存不均衡用ZeRO-3时每张卡的显存占用可能不均衡因为参数分片后某些卡可能分到更多层。DeepSpeed有stage3_prefetch_bucket_size等参数可以调但最直接的办法是看nvidia-smi的显存占用如果差异超过10%调整分片策略或换卡数。另外数据并行时如果某张卡的数据特别长激活值会比其他卡高导致OOM。用DistributedSampler保证数据均匀或者设置drop_lastTrue。5.4 混合精度下的数值稳定性BF16虽然稳定但某些操作还是要注意。比如softmax、layer_norm、loss计算最好在FP32下做。PyTorch的autocast会自动处理这些但自定义算子要手动加torch.cuda.amp.custom_fwd(cast_inputstorch.float32)装饰器。还有优化器的eps参数在混合精度下要调大。AdamW默认eps1e-8FP16下可能下溢改成1e-6或1e-5。BF16下1e-8没问题。5.5 显存监控与调试工具torch.cuda.memory_allocated()看当前显存torch.cuda.max_memory_allocated()看峰值。nvidia-smi看整体。DeepSpeed有deepspeed.runtime.zero.utils可以打印每张卡的显存明细。调试OOM时用torch.cuda.memory_summary()看显存碎片。有时候显存够但碎片太多也会OOM。设置PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True能减少碎片。这些细节看起来琐碎但每一个都可能导致训练失败。我踩过最坑的一次是优化器eps没调FP16训练到一半loss变NaN排查了两天才发现是下溢。6. 从估算到落地一套可复用的决策流程最后把整个流程串起来形成一套可复用的决策逻辑。拿到一个新模型先算参数量再按20字节每参数估静态显存加上激活值看单卡够不够。不够就上ZeRO2张卡用ZeRO-34张卡用ZeRO-28张卡以上用ZeRO-1加数据并行。精度优先BF16老卡用FP16加损失缩放。推理阶段显存够用BF16不够用INT8但接受精度损失。这套流程我在7B、13B、70B上都验证过估算误差在10%以内。关键是别偷懒每一步都算清楚比盲目试错省时间。显存估算不是玄学是算术。混合精度不是魔法是权衡。把这两件事搞明白大模型训练的门就推开了一半。
返回列表