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

资讯详情

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

大模型训练显存估算与混合精度调优实战

大模型训练显存估算与混合精度调优实战

干过大模型训练的人都清楚,“显存不够”往往是压垮训练任务的第一根稻草。不管是7B还是13B模型,开工前如果不先算清楚显存预算,等跑到一半OOM再回头调batch size、重开任务,损失的不只是时间,还有心态。这篇是这个系列的第三篇,专门讲两件事:怎么手算模型训练要吃多少显存,以及混合精度训练里那些经常被忽略的技术细节。适合正要准备全参数微调、在选卡或者规划序列长度的朋友;读完你至少能回答两个问题:这条训练管线需要几张80G卡?混合精度到底帮我省了哪部分显存?

显存估计这件事,很多人习惯直接抄别人给的“7B大概要多少G”,但别人的batch size、序列长度、是否开重计算都和你不一样,抄来必翻车。不如花五分钟把手里的模型配置代进公式,得到一个误差可接受的量级。混合精度则更微妙:它确实能把训练跑得更快、中间张量更小,但如果你以为用了AMP就万事大吉,结果显存没降多少、loss还时不时NaN,那大概率是没理解FP16/BF16的底牌。这篇就把账一五一十算给你看,再附上实操排坑清单。

1. 显存都去哪儿了:训练时的四类开销

1.1 模型参数本身:最直白的一笔

模型参数有多少个,就有多少权重数值占显存。但同样一组数字,用FP32存和用FP16存体积差一半。FP32每个参数占4字节,FP16/BF16每个参数占2字节。7B参数模型如果全用FP32存储,光权重就是 7e9 × 4B ≈ 28GB;FP16则约14GB。所以“这个模型多少个B”只是起点,真正的显存还要看存储精度和训练方式。

1.2 梯度:反向传播的“路程记录”

训练必须做反向传播,每个参数都会对应一个梯度,用于指导参数往哪个方向更新。梯度张量的大小和参数一样大,因此参数如果是FP32,梯度FP32就是每参数4字节;如果参数用FP16存储,梯度通常也会被算成FP16,每参数2字节。这里多说一句:在PyTorch标准的torch.cuda.amp里,如果你的模型参数保持FP32,反向传出来的梯度也会是FP32,显存占用并不会因为这个模式而下降。很多新手在这里有误解,以为用了AMP参数内存就减半,其实不是。

1.3 优化器状态:真正的大头,不是参数

以最常用的Adam/AdamW为例,它要给每个参数额外保存两个状态:一阶动量(momentum)和二阶动量(variance),这两个都按FP32存储,每个状态4字节。另外,如果我们把模型参数转成FP16训练,为了不让更新过程因为精度太低而发散,还需要单独保留一份FP32的“主权重”副本,又是4字节。算下来,混合精度训练中每个参数的固定开销可拆成:

  • FP16参数副本:2字节
  • FP16梯度:2字节
  • FP32主权重:4字节
  • Adam一阶动量:4字节
  • Adam二阶动量:4字节
  • 合计:16字节/参数

看到没?即使模型FP16存储,更新时还得背一份FP32主权重和两个FP32状态,每个参数依然要16字节。这一点常被营销话术掩盖:混合精度训练并没有让Adam优化器下的“固定显存”减半,它真正省下的是中间激活张量,以及通过降低计算精度换取的算力优势。用纯FP32训练,参数4字节+梯度4字节+两个状态8字节,也是16字节。两者打平。

不同优化器和精度组合的每参数固定字节数,可以直接查下面这张表:

训练配置参数梯度优化器状态每参数总字节
FP32 + Adam4B4B8B16B
FP16/FP32主权重 + Adam2B2B12B16B
FP16无主权重 + Adam2B2B8B12B
BF16 + AdamW + ZeRO offload2B2B状态可被切分/卸载不定

最后一行意味着通过ZeRO等策略,可以把优化器状态切到多卡甚至卸载到CPU,这才能突破单卡固定开销的瓶颈。

1.4 激活值与中间张量:藏得最深的“刺客”

前向传播时,每一层都会产生中间结果,比如Transformer里QKV投影输出、注意力分数、MLP中间层输出等。这些结果在反向传播计算梯度时还要再用一次,所以不能算完就扔,必须暂存在显存里。它们的大小和batch size、序列长度、注意力头数强相关,而和参数量关系不大。很多模型看着参数不多,却因为序列长、batch大,显存直接爆掉,罪魁祸首往往是激活值。这一块也是我们可以手动估算的核心部分,下一节专门展开。

除了上述四类,CUDA context、cuDNN workspace、PyTorch缓存分配器的预留区也会占用一小部分显存,通常几十MB到几GB不等,但不会成为决定性因素,估算时可以忽略,实测时再通过监控去看。

2. 显存估计:用一张纸算出你的卡够不够

2.1 固定开销:先把参数规模装进去

固定开销只依赖参数量P和训练配置,和batch size、序列长度无关。计算公式非常简单:

[ 固定显存(GB) = P \times bytes_per_param / 1024^3 ]

P是参数个数,bytes_per_param按上一节表格取。例如一个7B参数模型,用FP16+FP32主权重+Adam训练,每参数16字节,固定开销就是:

[ 7 \times 10^9 \times 16 / 1024^3 \approx 112GB ]

也就是说,只看参数、梯度、优化器状态,就已经需要1.5张80G卡了。还没算激活值和batch。这也是为什么很多人说全参数微调7B模型,单卡80G根本放不下,除非用LoRA、DeepSpeed ZeRO、CPU offload等手段把它拆走。

2.2 激活值粗略估算:一个能上手算的经验式

严格精确计算激活显存需要逐层追踪张量生命周期,工程上太累。这里给一个足够用于“拍脑袋”的近似式。假设我们要训练一个标准Transformer,模型配置为:

  • 层数L
  • 隐藏维H
  • 注意力头数A
  • batch size B
  • 序列长度S
  • 使用FP16/BF16存储中间张量

每个Transformer层需要暂存的激活量大约为:

[ 2 \times (19 \times B \times S \times H + 3 \times B \times A \times S^2) ]

单位是字节。解释一下两部分来源:前半部分包括层的输入、QKV投影结果、注意力上下文输出、MLP两层中间结果、残差连接需要保留的输入等,大约折合19份B×S×H的二维矩阵;后半部分是注意力分数矩阵的三种形态(QK^T的结果、softmax之后的概率、dropout掩码或概率版),每个都是B×A×S×S的大方块。因为按2字节精度的FP16存储,所以前面乘2。这个式子不追求精确,但能反映出两个关键趋势:序列长度S影响极大,因为注意力矩阵是S²增长;batch size和隐藏维是线性增长。

如果开启了激活检查点(gradient checkpointing),代价是反向传播时重算一部分前向,激活显存可以大幅降低,估算时大约只需要保留“每层输入激活”和一层重算所需的中间量,整体量级大约是一个layer的激活量,可以按上式“乘以1”而不是“乘以L”来近似。

2.3 一个7B模型的计算实例

我们用常见配置:7B模型,L=32,H=4096,A=32,序列长度S=2048。先算batch size=1时每一层激活:

[ 19 \times 1 \times 2048 \times 4096 = 159,383,552 ] [ 3 \times 1 \times 32 \times 2048^2 = 402,653,184 ]

两部分相加约5.62亿个元素,乘2字节得到约1.12GB。这是单层激活,32层都不开检查点则约36GB。固定开销112GB + 激活36GB ≈ 148GB,两张80G卡能放下但比较紧张。如果不开激活检查点,再把batch size提到4,激活变成144GB,总显存约256GB,两张80G就不够了,至少需要4张,而且这只是压着内存线跑,实际情况还要看PyTorch缓存碎片,建议再多留余量。

配置固定开销激活开销(估算)总显存
7B,B=1,S=2048,无检查点112GB36GB~148GB
7B,B=1,S=2048,开检查点112GB约4GB~116GB
7B,B=4,S=2048,无检查点112GB144GB~256GB
13B,B=1,S=2048,开检查点约208GB约5GB~213GB

13B的固定开销:13e9×16B≈208GB,所以要跑13B全参训练,4卡80G是起步,5卡更稳。这也是为什么现在大家更青睐LoRA:LoRA只训练少量新增参数,主模型冻结后可省掉优化器状态和梯度的巨额开销,大大降低门槛。

2.4 从估算到代码:用memory stats验证

纸面估算完成后,一定要让训练代码“自报家门”。PyTorch的CUDA caching allocator提供的数字比nvidia-smi更接近真实逻辑使用量,因为显卡驱动显示的显存可能包含预留而未用的部分。在训练循环里打印:

import torch def print_gpu_memory(): allocated = torch.cuda.memory_allocated() / 1024**3 reserved = torch.cuda.memory_reserved() / 1024**3 print(f"allocated: {allocated:.2f}GB | reserved: {reserved:.2f}GB")

每次step结束时调用一次,观察allocated的峰值和reserved的差值。如果allocated离80G很远就开始OOM,多半是缓存碎片问题。这时可以设置环境变量:

export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True

让它利用可扩展内存段,通常能显著缓解碎片导致的假OOM。另一种方案是max_split_size_mb调小,但实际体验上expandable_segments更有效。

3. 混合精度训练:它做什么,不做什么

3.1 FP16和BF16:一个怕下溢,一个怕精度细

FP16用1位符号、5位指数、10位尾数,能表示的数值范围很小:最大约65504,最小的normal值约6e-5。它的问题不是“存不下大数”,而是“存不了特别小的数”,一旦某个中间结果小于约1e-5,再往下就变成0,这叫下溢。梯度在深层网络中经常会出现非常小的数值,正是FP16的噩梦。

BF16则用了和FP32相同的8位指数,最大值也是FP32同量级,范围宽阔得多,几乎不会出现下溢;但它只有7位尾数,精度比FP16更粗。打个比方:FP16像一把刻度更细但量程有限的尺子,BF16像一把量程很广但刻度很粗的尺子。大模型训练里,梯度数值范围的最主要矛盾往往是小数下溢,所以新一代GPU(比如A100、H100、RTX 30系以后)上BF16往往比FP16更稳,不需要额外做loss scaling也能跑。反过来,早期的V100、T4这些没有BF16硬件加速的卡,就只能用FP16 + loss scaling。

3.2 Loss Scaling:为什么损失放大一倍不是玄学

既然FP16存不了太小的数,那操作办法很朴素:既然梯度在进入FP16存储前变得太小,那就先把loss乘一个大因子,例如1024。因为链式法则,反向传播中所有梯度都会跟着“放大同样倍数”,原本可能小于FP16下限的值被搬进可表示范围。等真正更新参数前,再把梯度除以这个因子,恢复真实梯度。

这个“先乘后除”的过程就是loss scaling。手动调一个固定scale也行,但不保险,因为不同训练阶段的梯度量级差很多。于是就有了动态策略:每隔一段时间,如果梯度没出现inf/NaN,就把scale往大调;一旦发现梯度中有inf/NaN,说明放大过头了,就回退半分,并跳过这次优化器更新。PyTorch的GradScaler封装的就是这套逻辑。BF16下通常不需要它,因为指数范围足够大,但如果训练早期loss突然发散,也可以临时罩一层来兜底,代价不大。

3.3 两种混合精度“流派”:torch AMP与完整半精度

这是最容易被混在一起说的地方,必须讲清楚。

第一种是PyTorch原生AMP,典型写法是torch.autocast+GradScaler。它做的事是在算子层面动态选择精度:比如矩阵乘法、卷积这类计算密集且适合TensorCore的操作,自动把输入转成FP16去算;LayerNorm、Softmax这类对精度敏感的操作,仍保持FP32。模型参数本身不会整体转成FP16,优化器更新依然在FP32参数上做,因此不需要额外维护FP32主权重。它的价值是:代码改动极小、速度提升明显、中间张量也是FP16从而省内存。但注意,它并没有省下“参数+梯度+优化器状态”那部分固定内存。

第二种是“完整半精度训练”,典型代表是DeepSpeed和部分手动实现。先调用model.half()或者model.to(dtype=torch.bfloat16)把模型参数整体转成低精度存储,再额外保留一份FP32主权重给优化器用。这时训练管线里的权重和中间激活大部分都是低精度,整体显存占用会更接近我们第一节表里的“FP16参数2B”那一行。它的固定开销依然是每参数16字节,但因为参数本体是2字节而不是4字节,通信和某些矩阵乘法的内存带宽压力会小一些。

所以一句话总结:真正省显存的是“削减优化器状态/激活”,而不是把参数从FP32换成FP16这件事本身。如果你的目标是塞进一张卡,最有效的手段是LoRA、ZeRO、激活检查点,而不是单纯打开AMP。

3.4 什么时候用FP16,什么时候用BF16

我自己的选择标准很简单:

  • 如果卡是A100/H100/4090/3090这类支持BF16的,优先BF16,因为省事,不需要跟loss scale斗智斗勇,对学习率也不那么敏感。
  • 如果只有V100/T4/2080Ti这类不支持BF16加速或支持的生态不完善的卡,就用FP16 + GradScaler。
  • 如果模型本身有非常多的Softmax、LayerNorm、注意力打分这类对精度敏感的结构,BF16偶尔会在小模型上出现精度损失,可以先跑10个step对比FP16和BF16的loss曲线,再决定用哪个。
  • 如果发现GradScaler的scale一路降到16甚至更低,而梯度还在出现inf,那大概率不是scale的问题,而是梯度爆炸。先看学习率,再看梯度裁剪,最后看数据里是不是混入了异常样本。

3.5 混合精度训练的关键流程

无论用哪种流派,优化器更新前都有一个必须完成的顺序:

  1. 前向时启用autocast,loss得到FP16/BF16计算下的输出。
  2. 将loss乘上当前scale,再backward,梯度被放大。
  3. 反向后先unscale梯度,即除以scale,并趁机检查是否有inf/NaN。
  4. 如果梯度正常,做梯度裁剪,然后调用optimizer.step()更新参数。
  5. 最后通过scaler.update()调整下次的scale。

这里特别容易踩的坑是梯度裁剪的时机。很多人习惯直接调用torch.nn.utils.clip_grad_norm_,但如果你用GradScaler,梯度此时还是放大后的状态,必须先scaler.unscale_(optimizer)再裁剪。顺序错了,裁剪阈值等于形同虚设。

4. 实战配置、问题速查与显存兜底方案

4.1 PyTorch混合精度训练循环最小可跑代码

给一段简洁但完整的最小示例。以FP16 + GradScaler为例:

import torch import torch.nn as nn model = MyTransformer() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5) scaler = torch.amp.GradScaler("cuda") # PyTorch 2.x 新写法 # 旧写法:torch.cuda.amp.GradScaler() for batch in dataloader: optimizer.zero_grad() with torch.autocast(device_type="cuda", dtype=torch.float16): loss = model(batch) scaler.scale(loss).backward() # 关键:先unscale再clip scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update()

如果换成BF16,可以用:

with torch.autocast(device_type="cuda", dtype=torch.bfloat16): loss = model(batch)

BF16通常不需要GradScaler,但unscale_和step的代码结构仍然建议保留,因为万一训练中途出现NaN想启用scale时,改起来只差一行。另外注意:scaler.unscale_在同一个optimizer上只能调用一次,重复调用会报错;如果你没有调用scaler.unscale_,clip_grad_norm_之前也要保证梯度已经是真实值。把上面的流程当成标准模板存下来,能少踩很多坑。

4.2 显存不够时的“救命”优先级

如果真的OOM了,按下面这个顺序调整,性价比从高到低:

  1. 缩小batch size到1。这是最直接的激活显存压缩器。batch缩小后,用梯度累积来补偿batch size下降带来的不稳定,但梯度累积不省显存,它只是等价放大batch。
  2. 打开激活检查点。PyTorch自带接口非常方便:
model.gradient_checkpointing_enable()

代价是大约20%-30%的训练速度损失,但激活内存能从几十GB压到几GB,划算得不能再划算。 3. 使用ZeRO或CPU offload。ZeRO阶段1把优化器状态切分到多卡,阶段2把梯度也切分,阶段3把参数也切分。如果只有单卡,可以用DeepSpeed的offload_optimizer把优化器状态放到CPU内存。 4. 改用LoRA/QLoRA。这是“显存实在不够”后的终极方案,只训练极小一部分低秩适配器,主模型冻结或量化,7B模型能压到一张消费级显卡。

记住,梯度累积不省显存,能不写就别为了“显存”去写它;真正需要它时,它的意义是大batch,不是省内存。

4.3 常见问题速查表

现象可能原因处理建议
训练到一半OOM激活缓存积累/峰值过高减小batch、开激活检查点、扩容缓存配置
OOM但memory_allocated并不高CUDA缓存碎片设置expandable_segments:True或max_split_size_mb
loss变成NaN/Inf学习率过大、梯度爆炸降低lr、增加梯度裁剪、检查数据中NaN
GradScaler的scale持续下降梯度确实有inf先查lr,再查模型结构是否数值不稳定,不要只盯loss scaling
用了AMP后显存没降参数仍是FP32,AMP只省中间张量改用完整半精度方案,或用LoRA
训练速度反而更慢频繁cast/小算子太多/CPU瓶颈增大batch,减少张量转精度次数,检查GPU利用率
BF16训练收敛明显变差尾数精度不足换FP16+scaler,或提高模型宽度而不是深度

4.4 容易被忽略的“收尾”细节

训练结束保存模型时,如果用的是FP16 + FP32主权重方案,不要只保存model.state_dict(),否则下次加载的可能是FP16权重,直接拿去推理精度可能打折扣。正确做法是保存训练用的主权重,或者在保存前用主权重覆盖模型并转回所需精度。PyTorch AMP路线下模型参数本来就在FP32,没这个问题,但完整半精度方案必须注意。

另外,我在实际训练中习惯每100步打印一次当前的scale、loss和最大梯度范数。如果scale一直在涨,说明训练稳定;如果scale掉到16以下且频繁触发跳过step,那基本可以断定是学习率太大或者数据出了问题。这比看loss曲线更早暴露危险。

最后分享一个经验

显存估算这件事,我踩过最大的坑是过分相信网上的显存占用截图。别人说7B全参数微调只要70G,我照着配,结果batch size、seq length一不一样都不知道,最后浪费了半天调试。后来我学乖了:先把固定开销用每参数16字节秒算出来,再根据激活公式估算一个数量级,然后直接跑一个batch=1、seq短的小实验,用memory_allocated验证,误差控制在10GB以内后才去正式调参。这套流程虽然不精致,但从来不会让我在设备规划上翻车。

混合精度也一样,它救不了显存规划错误,它的价值是让TensorCore干活,让中间张量减半,让训练速度变快。真正想要把70B模型塞进单卡,还得靠LoRA、ZeRO和重计算这些“开源节流”的组合拳。希望这篇能让你下次新开训练任务时,第一件事不再是焦虑,而是冷静算一遍这张卡到底装不装得下。

返回列表