大模型训练绕不开两个硬骨头:显存不够和精度怎么选。我见过太多团队在单卡上跑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 × (2+4) = 42GB。梯度也存在FP16和FP32两份,又是42GB。这一下就84GB了,单卡80GB直接爆掉。这也是为什么混合精度训练通常要配合ZeRO或者模型并行。
1.2 优化器状态的显存黑洞
优化器状态是大头,尤其Adam系列。Adam为每个参数维护两个状态:一阶矩估计(动量)和二阶矩估计(方差)。如果优化器状态用FP32存储,每个参数需要4+4=8字节。加上参数本身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_size=4096,num_layers=32,seq_len=2048,batch_size=1。粗略算: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两份,所以是2+4=6字节每参数,参数加梯度共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,加激活值8GB,268GB。至少4张80GB卡。
70B模型:70B × 20 = 1400GB,加激活值40GB,1440GB。至少18张80GB卡,实际要20张以上留余量。
| 模型规模 | 参数量 | 静态显存(20字节/参数) | 激活值(bs=1) | 总显存 | 80GB卡数量 |
|---|---|---|---|---|---|
| 7B | 7B | 140GB | 4GB | 144GB | 2 |
| 13B | 13B | 260GB | 8GB | 268GB | 4 |
| 70B | 70B | 1400GB | 40GB | 1440GB | 18 |
这个表是保守估算,实际用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等老卡上不得不用。
| 对比项 | FP16 | BF16 |
|---|---|---|
| 指数位 | 5位 | 8位 |
| 尾数位 | 10位 | 7位 |
| 动态范围 | 6e-5 ~ 65504 | 1e-38 ~ 3e38 |
| 损失缩放 | 必须 | 不需要 |
| 硬件要求 | 多数GPU | Ampere以上 |
| 训练稳定性 | 一般,可能溢出 | 好 |
| 速度 | 快 | 快 |
2.4 混合精度训练的代码实操
PyTorch的torch.cuda.amp是标准做法。下面是一个训练循环的骨架:
import torch from torch.cuda.amp import autocast, GradScaler model = MyModel().cuda() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) scaler = GradScaler() # FP16需要,BF16不需要 for batch in dataloader: inputs, labels = batch optimizer.zero_grad() with autocast(dtype=torch.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(dtype=torch.bfloat16),去掉scaler相关调用。注意autocast只影响前向,反向的梯度计算自动用对应精度。优化器更新时,PyTorch会自动把梯度转成FP32再更新master weight。
注意:
autocast区域内的操作要检查是否支持低精度。有些自定义算子可能不支持,会回退到FP32,影响速度。用torch.autocast的enabled参数可以临时关闭。
3. 显存优化实战:从OOM到跑通7B模型的完整过程
理论算完,上手跑7B模型还是OOM。我记录了一次完整的排查过程,从单卡80GB开始,一步步调到能跑。
3.1 第一次尝试:单卡全量训练直接爆
7B模型,FP16混合精度,AdamW,batch_size=1,seq_len=2048。按公式算静态显存140GB,单卡80GB肯定不够。但我想试试PyTorch的torch.cuda.amp能不能省点。结果加载模型就占了14GB(FP16参数),优化器初始化又占了56GB(FP32参数+优化器状态),还没开始训练就70GB了。前向传播一跑,激活值加上去直接OOM。
这里有个细节:PyTorch加载模型时,如果直接.cuda(),参数是FP32的。用.half()转FP16会省一半,但优化器状态还是FP32。所以加载完模型先转FP16,再初始化优化器,能省14GB。
3.2 第二次尝试:梯度累积加小batch
batch_size降到1已经最小了,只能减seq_len。从2048降到1024,激活值减半。但静态显存没变,还是140GB。单卡无解,必须上多卡。
3.3 第三次尝试:ZeRO-2分片优化器状态和梯度
用DeepSpeed的ZeRO-2,把优化器状态和梯度分片到2张卡上。每张卡的静态显存变成:参数14GB(FP16)+ 参数28GB(FP32 master)+ 梯度14GB(FP16)+ 梯度28GB(FP32)/ 2 + 优化器状态56GB / 2。算下来每张卡约14+28+14+14+28=98GB,还是超80GB。
等等,ZeRO-2的分片逻辑是:优化器状态和梯度分片,参数不分片。所以每张卡都有完整的FP32参数28GB和FP16参数14GB,梯度只存自己那份。重新算:FP16参数14GB + FP32参数28GB + FP16梯度14GB/2 + FP32梯度28GB/2 + 优化器状态56GB/2 = 14+28+7+14+28 = 91GB。还是超。
3.4 第四次尝试:ZeRO-3加CPU offload
ZeRO-3把参数也分片,每张卡只存1/2的参数。FP16参数7GB + FP32参数14GB + 梯度分片 + 优化器状态分片。算下来每张卡约7+14+7+14+28=70GB,加上激活值4GB,74GB,勉强能跑。但ZeRO-3通信开销大,训练速度降了约30%。
最后用ZeRO-3加CPU offload优化器状态,每张卡降到50GB左右,速度降了40%但能跑通。如果换成4张卡,ZeRO-2就够了,速度损失小很多。
| 方案 | 卡数 | 每卡显存 | 速度损失 | 可行性 |
|---|---|---|---|---|
| 单卡全量 | 1 | >140GB | 0 | 不可行 |
| ZeRO-2 | 2 | 91GB | 10% | 不可行 |
| ZeRO-3 | 2 | 74GB | 30% | 勉强 |
| ZeRO-3+offload | 2 | 50GB | 40% | 可行 |
| ZeRO-2 | 4 | 45GB | 10% | 推荐 |
这个排查过程说明:显存估算要准,优化方案要按卡数选。2张卡优先ZeRO-3,4张卡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适合对精度要求不高的场景,比如分类、抽取。
| 对比项 | BF16 | INT8 |
|---|---|---|
| 位数 | 16 | 8 |
| 显存 | 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_size=1,累积4步,等效batch_size=4。但要注意: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_last=True。
5.4 混合精度下的数值稳定性
BF16虽然稳定,但某些操作还是要注意。比如softmax、layer_norm、loss计算,最好在FP32下做。PyTorch的autocast会自动处理这些,但自定义算子要手动加@torch.cuda.amp.custom_fwd(cast_inputs=torch.float32)装饰器。
还有,优化器的eps参数在混合精度下要调大。AdamW默认eps=1e-8,FP16下可能下溢,改成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_CONF=expandable_segments:True能减少碎片。
这些细节看起来琐碎,但每一个都可能导致训练失败。我踩过最坑的一次是优化器eps没调,FP16训练到一半loss变NaN,排查了两天才发现是下溢。
6. 从估算到落地:一套可复用的决策流程
最后把整个流程串起来,形成一套可复用的决策逻辑。拿到一个新模型,先算参数量,再按20字节每参数估静态显存,加上激活值,看单卡够不够。不够就上ZeRO,2张卡用ZeRO-3,4张卡用ZeRO-2,8张卡以上用ZeRO-1加数据并行。精度优先BF16,老卡用FP16加损失缩放。推理阶段显存够用BF16,不够用INT8但接受精度损失。
这套流程我在7B、13B、70B上都验证过,估算误差在10%以内。关键是别偷懒,每一步都算清楚,比盲目试错省时间。显存估算不是玄学,是算术。混合精度不是魔法,是权衡。把这两件事搞明白,大模型训练的门就推开了一半。