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 StatesAdam 会为每个参数维护一阶矩估计动量和二阶矩估计方差所以是参数量的两倍。激活值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 是 4096num_layers 是 32。如果 batch_size 是 1seq_len 是 2048那么激活值大约需要 1 × 2048 × 4096 × 32 × 15 ≈ 4 GB。如果 seq_len 翻倍到 4096激活值也会翻倍到 8 GB。如果 batch_size 再翻倍那就是 16 GB。这就是为什么长序列训练特别吃显存。实际训练中激活值可以通过梯度检查点Gradient Checkpointing来大幅降低。梯度检查点的思路是前向传播时不保存所有中间激活值只保存部分检查点的激活值反向传播时重新计算被丢弃的激活值。这样可以把激活值显存降低到原来的 1/√N 到 1/NN 是层数代价是增加大约 30% 的计算时间。在显存紧张的时候这是一个非常划算的 trade-off。3. 混合精度训练FP16 和 BF16 到底怎么选3.1 FP16 和 BF16 的本质区别FP16 和 BF16 都是 16 位浮点数但它们的位分配不同FP161 位符号位5 位指数位10 位尾数位。BF161 位符号位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(), lr1e-4) scaler GradScaler() for input_ids, labels in dataloader: optimizer.zero_grad() with autocast(dtypetorch.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 为 1seq_len 为 2048。参数、梯度、优化器状态13 × 10^9 × 16 208 GB激活值假设 hidden_size 为 5120num_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_CONFexpandable_segments:True让分配器使用可扩展的显存段减少碎片二是定期调用torch.cuda.empty_cache()但这会拖慢训练速度不建议频繁使用。另一个可能的原因是动态损失缩放导致某些 step 的梯度异常大触发了额外的显存分配。可以检查一下 scaler 的缩放因子是否在合理范围内。5.2 FP16 训练出现 NaN 怎么排查FP16 训练出现 NaN 是常见问题排查思路如下检查损失缩放是否开启如果没开损失缩放FP16 训练几乎必然出现 NaN。检查缩放因子是否过大缩放因子太大会导致梯度上溢变成 inf然后变成 NaN。可以手动调小初始缩放因子。检查模型中有没有不稳定的算子比如 exp、log、softmax 等在 FP16 下容易溢出。可以强制这些算子用 FP32 计算。检查数据中是否有异常值比如标签越界、输入包含 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_segmentsFP16 出现 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_memoryTrue和多个 worker数据会在 CPU 和 GPU 之间频繁传输有时候会占用不少显存。如果显存实在紧张可以把pin_memory关掉或者减少 worker 数量。最后说一个我踩过的坑有一次用 FP16 训练一个对话模型loss 一直很正常但生成出来的回复全是乱码。排查了很久才发现是保存模型的时候没有把 FP16 参数转回 FP32导致推理时精度损失严重。所以训练完保存模型时一定要确认保存的是 FP32 权重或者在推理时做相应的精度转换。这个坑不常遇到但一旦遇到就很难排查希望大家引以为戒。
阅读完成 · 觉得有帮助?