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

资讯详情

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

大模型训练显存估计与混合精度训练实战指南

大模型训练显存估计与混合精度训练实战指南

1. 大模型训练显存估计与混合精度训练详解

显存不够用,几乎是每个做大模型训练的人都会撞上的第一堵墙。你可能也经历过:模型代码写完了,数据管道跑通了,满心欢喜地按下训练启动脚本,结果几秒钟后终端弹出一行红字——CUDA out of memory。然后开始反复调 batch size,从 32 降到 16,再降到 8,最后发现 batch size 等于 1 都跑不起来。这时候你才意识到,问题不是 batch size 太大,而是你根本不知道显存到底花在了哪里。

这篇文章就是来解决这个问题的。我会把大模型训练中的显存估计方法拆开讲清楚,让你在启动训练之前就能算出一张卡到底能不能装下你的模型;同时把混合精度训练(FP16 和 BF16)的原理、实操配置和踩坑经验一并说透。适合正在做或准备做大模型训练的同学,无论你是刚入门还是已经跑过几轮实验,都能从里面找到可以直接用的东西。

2. 显存到底被谁吃掉了

2.1 模型参数、梯度、优化器状态:显存的三座大山

很多人第一次算显存的时候,只算了模型参数的大小。比如一个 7B 参数的模型,用 FP16 存储,那就是 7 × 10^9 × 2 字节,大约 14 GB。然后一看自己手里是 24 GB 显存的卡,觉得绰绰有余。结果一跑就炸。为什么?因为你只算了三分之一。

在标准的 Adam 优化器训练中,显存消耗主要来自四个部分:

  • 模型参数(Parameters):模型本身的权重。
  • 梯度(Gradients):反向传播时每个参数对应的梯度,大小和参数一样。
  • 优化器状态(Optimizer States):Adam 会为每个参数维护一阶矩估计(动量)和二阶矩估计(方差),所以是参数量的两倍。
  • 激活值(Activations):前向传播过程中每一层的输出,需要保留到反向传播时计算梯度。

这四部分加起来,才是你真正需要的显存。而且激活值这一块,往往是大头,尤其是序列长度比较长的时候。

2.2 混合精度下每部分显存怎么算

混合精度训练的核心思路是:前向和反向计算用 FP16 或 BF16,但优化器状态和模型的主权重用 FP32 保存。所以显存的计算会稍微复杂一点。

以一个参数量为 P 的模型为例,在混合精度 + Adam 的训练配置下:

组成部分数据类型显存占用
FP16 模型参数FP162P 字节
FP32 主权重副本FP324P 字节
FP16 梯度FP162P 字节
FP32 优化器一阶矩FP324P 字节
FP32 优化器二阶矩FP324P 字节
激活值FP16与 batch size、序列长度、层数相关

把前五项加起来,每个参数大约需要 2 + 4 + 2 + 4 + 4 = 16 字节。也就是说,一个 7B 模型,光是参数、梯度、优化器状态这三块,就需要 7 × 10^9 × 16 = 112 GB。这已经远远超过单张 80 GB 卡的容量了。所以实际训练中,7B 模型通常需要配合 ZeRO 或张量并行等分布式策略才能跑起来。

如果不用混合精度,全部用 FP32 训练,那每个参数的显存开销是 4(参数)+ 4(梯度)+ 4(一阶矩)+ 4(二阶矩)= 16 字节,和混合精度下的总开销一样。但混合精度下计算用的是 FP16,速度更快,而且激活值占用的显存也减半。所以混合精度几乎是必选项。

2.3 激活值显存:最容易被低估的部分

激活值的显存估算是最复杂的,因为它和模型结构、序列长度、batch size 都相关。一个粗略的估算公式是:

激活值显存 ≈ batch_size × seq_len × hidden_size × num_layers × 系数

这个系数取决于具体的模型结构和实现方式,通常在 10 到 20 之间。以 LLaMA-7B 为例,hidden_size 是 4096,num_layers 是 32。如果 batch_size 是 1,seq_len 是 2048,那么激活值大约需要 1 × 2048 × 4096 × 32 × 15 ≈ 4 GB。如果 seq_len 翻倍到 4096,激活值也会翻倍到 8 GB。如果 batch_size 再翻倍,那就是 16 GB。这就是为什么长序列训练特别吃显存。

实际训练中,激活值可以通过梯度检查点(Gradient Checkpointing)来大幅降低。梯度检查点的思路是:前向传播时不保存所有中间激活值,只保存部分检查点的激活值,反向传播时重新计算被丢弃的激活值。这样可以把激活值显存降低到原来的 1/√N 到 1/N(N 是层数),代价是增加大约 30% 的计算时间。在显存紧张的时候,这是一个非常划算的 trade-off。

3. 混合精度训练:FP16 和 BF16 到底怎么选

3.1 FP16 和 BF16 的本质区别

FP16 和 BF16 都是 16 位浮点数,但它们的位分配不同:

  • FP16:1 位符号位,5 位指数位,10 位尾数位。
  • BF16:1 位符号位,8 位指数位,7 位尾数位。

关键区别在指数位。FP16 的指数位只有 5 位,能表示的数值范围是大约 6 × 10^-5 到 65504。BF16 的指数位有 8 位,和 FP32 一样,所以数值范围和 FP32 相同,大约是 10^-38 到 10^38。

这意味着什么?FP16 在训练中很容易溢出。当梯度值小于 6 × 10^-5 时,FP16 会把它变成 0,这就是下溢(underflow)。当梯度值大于 65504 时,FP16 会变成 inf,这就是上溢(overflow)。而 BF16 因为指数位和 FP32 一样,基本不会出现溢出问题。

但 BF16 的尾数位只有 7 位,比 FP16 的 10 位少,所以精度更低。不过在大模型训练中,精度损失可以通过其他方式补偿,而溢出问题一旦出现就很难处理。所以现在主流的大模型训练,比如 GPT-3、LLaMA、PaLM,都优先使用 BF16。

3.2 什么时候必须用 FP16

虽然 BF16 是更好的选择,但有一个现实问题:不是所有硬件都支持 BF16。BF16 需要 NVIDIA Ampere 架构及以上的 GPU(比如 A100、A30、RTX 30 系列及以上)。如果你用的是 V100 或更早的卡,那就只能用 FP16。

用 FP16 训练时,必须配合损失缩放(Loss Scaling)。损失缩放的原理是:在计算损失时乘以一个很大的缩放因子(比如 2^16),这样反向传播得到的梯度也会被放大同样的倍数,从而避免下溢。在更新参数之前,再把梯度除以这个缩放因子,恢复原来的尺度。

实际操作中,通常使用动态损失缩放:先从一个较大的缩放因子开始,如果连续多个 step 没有出现 inf 或 NaN,就增大缩放因子;如果出现了,就减小缩放因子并跳过这个 step。PyTorch 的torch.cuda.amp模块已经内置了动态损失缩放,用起来很方便。

3.3 混合精度训练的实操配置

在 PyTorch 中,混合精度训练的标准写法是:

from torch.cuda.amp import autocast, GradScaler model = MyModel().cuda() optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) scaler = GradScaler() for input_ids, labels in dataloader: optimizer.zero_grad() with autocast(dtype=torch.bfloat16): # 或 torch.float16 outputs = model(input_ids) loss = loss_fn(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

如果你用的是 BF16,其实可以不用 GradScaler,因为 BF16 基本不会溢出。但为了代码统一,保留 scaler 也没问题,它不会对 BF16 造成负面影响。

注意:使用 autocast 时,模型的前向传播会自动把部分算子转换成 FP16/BF16,但有些算子(比如 softmax、layer norm)仍然会用 FP32 计算,以保证数值稳定性。这是框架自动处理的,不需要你手动干预。

4. 显存估计的实操方法

4.1 用公式快速估算

在启动训练之前,你可以用下面的公式快速估算显存需求:

总显存 ≈ 参数量 × 16 字节 + 激活值显存 + 临时缓冲区

其中参数量 × 16 字节是参数、梯度、优化器状态的总和。激活值显存可以用前面提到的公式估算。临时缓冲区通常不大,但也要留出 1-2 GB 的余量。

举个例子:一个 13B 参数的模型,用 BF16 + Adam 训练,batch_size 为 1,seq_len 为 2048。

  • 参数、梯度、优化器状态:13 × 10^9 × 16 = 208 GB
  • 激活值:假设 hidden_size 为 5120,num_layers 为 40,系数取 15,则 1 × 2048 × 5120 × 40 × 15 ≈ 6.3 GB
  • 临时缓冲区:2 GB

总计约 216 GB。这显然单卡放不下,需要至少 4 张 80 GB 的卡配合 ZeRO-3 才能跑起来。

4.2 用工具精确测量

公式估算只能给你一个大概的范围,实际显存占用还会受到框架实现、算子优化等因素的影响。如果你想精确测量,可以用 PyTorch 提供的工具:

import torch # 训练前记录初始显存 torch.cuda.reset_peak_memory_stats() initial_mem = torch.cuda.memory_allocated() # 训练几步 for step, batch in enumerate(dataloader): train_step(batch) if step == 5: break # 查看峰值显存 peak_mem = torch.cuda.max_memory_allocated() print(f"峰值显存: {peak_mem / 1024**3:.2f} GB")

这个方法可以让你在跑了几步之后就知道实际峰值显存是多少,比公式估算准确得多。建议在正式训练之前,先用小规模数据跑几步,测一下实际显存,然后再决定 batch size 和并行策略。

4.3 显存不够时的应对策略

如果测出来显存不够,有几个方向可以调整:

  • 减小 batch size:最直接的方法,但可能会影响训练稳定性。
  • 启用梯度检查点:用计算换显存,激活值显存可以降低到原来的 1/3 到 1/5。
  • 使用 ZeRO 优化:把优化器状态、梯度、参数分片到多张卡上,单卡显存大幅降低。
  • 使用张量并行:把单个 Transformer 层的计算拆分到多张卡上,适合超大模型。
  • 使用 CPU Offload:把优化器状态放到 CPU 内存,进一步降低显存,但会拖慢训练速度。

这些策略可以组合使用。比如 7B 模型单卡 80 GB 跑不动,可以用 ZeRO-2 把优化器状态分片到 2 张卡上,就能跑起来了。13B 模型可能需要 ZeRO-3 加梯度检查点。70B 模型就需要 ZeRO-3 + 张量并行 + 梯度检查点 + CPU Offload 全套上了。

5. 常见问题与排查技巧实录

5.1 训练中突然 OOM 怎么办

有时候训练刚开始没问题,跑了几百步之后突然 OOM。这种情况通常是因为显存碎片化。PyTorch 的显存分配器在反复分配和释放显存后,会产生碎片,导致没有足够的连续显存来分配新的张量。

解决方法有两个:一是设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,让分配器使用可扩展的显存段,减少碎片;二是定期调用torch.cuda.empty_cache(),但这会拖慢训练速度,不建议频繁使用。

另一个可能的原因是动态损失缩放导致某些 step 的梯度异常大,触发了额外的显存分配。可以检查一下 scaler 的缩放因子是否在合理范围内。

5.2 FP16 训练出现 NaN 怎么排查

FP16 训练出现 NaN 是常见问题,排查思路如下:

  1. 检查损失缩放是否开启:如果没开损失缩放,FP16 训练几乎必然出现 NaN。
  2. 检查缩放因子是否过大:缩放因子太大会导致梯度上溢,变成 inf,然后变成 NaN。可以手动调小初始缩放因子。
  3. 检查模型中有没有不稳定的算子:比如 exp、log、softmax 等,在 FP16 下容易溢出。可以强制这些算子用 FP32 计算。
  4. 检查数据中是否有异常值:比如标签越界、输入包含 NaN 等。

如果排查了一圈还是找不到原因,可以先用 BF16 跑一遍,确认模型本身没问题,再切回 FP16 排查。

5.3 BF16 训练 loss 不下降怎么处理

BF16 的精度比 FP16 低,有时候会出现 loss 不下降或者下降很慢的情况。这时候可以尝试:

  • 提高学习率:BF16 的梯度精度较低,适当提高学习率可以加快收敛。
  • 使用 FP32 主权重:确保优化器更新的是 FP32 的主权重,而不是直接更新 BF16 参数。
  • 检查梯度裁剪:BF16 下梯度裁剪的阈值可能需要调整。
  • 混合使用 FP16 和 BF16:有些框架支持在前向用 BF16、反向用 FP16,兼顾速度和精度。

5.4 常见问题速查表

问题现象可能原因解决方法
启动即 OOMbatch size 太大或模型太大减小 batch size,启用梯度检查点
训练中途 OOM显存碎片化设置 expandable_segments
FP16 出现 NaN损失缩放未开启或缩放因子过大开启动态损失缩放,调小初始因子
BF16 loss 不下降学习率过低或精度不足提高学习率,确保 FP32 主权重
显存占用忽高忽低动态损失缩放导致检查 scaler 状态,固定缩放因子

6. 一些实操心得和避坑建议

显存估计这件事,公式只能给你一个起点,真正的数字一定要实测。我自己的习惯是:在正式训练之前,先用 1/10 的数据量跑 50 步,用torch.cuda.max_memory_allocated()记录峰值显存,然后按比例放大到全量数据。这样估算出来的数字比任何公式都准。

混合精度方面,如果你的卡支持 BF16,那就无脑用 BF16,省心省力。如果只能用 FP16,那损失缩放一定要开,而且初始缩放因子不要设太大,从 2^12 开始比较稳妥。另外,FP16 训练时建议把 layer norm 和 softmax 强制用 FP32 计算,这两个算子对数值精度很敏感。

还有一个容易被忽略的点:数据加载器也会占显存。如果你用了pin_memory=True和多个 worker,数据会在 CPU 和 GPU 之间频繁传输,有时候会占用不少显存。如果显存实在紧张,可以把pin_memory关掉,或者减少 worker 数量。

最后说一个我踩过的坑:有一次用 FP16 训练一个对话模型,loss 一直很正常,但生成出来的回复全是乱码。排查了很久才发现是保存模型的时候没有把 FP16 参数转回 FP32,导致推理时精度损失严重。所以训练完保存模型时,一定要确认保存的是 FP32 权重,或者在推理时做相应的精度转换。这个坑不常遇到,但一旦遇到就很难排查,希望大家引以为戒。

返回列表