拓冰建站拓冰建站
首页 / 资讯中心 / 正文

大模型训练显存估算与混合精度实战:从OOM到BF16选型

开头先从一次真实翻车现场说起。去年我把一个 13B 模型放到单卡上做微调盯着nvidia-smi看显存从 12GB 往上涨然后眼睁睁看它撞上 80GB 的墙——OOM 报错弹出来那一刻我才意识到自己对大模型训练显存估计的理解有多肤浅。后来换了混合精度训练又踩了 BF16 和 FP16 的坑才把这一整套逻辑理顺。这篇就把大模型训练里显存估算的方法和混合精度训练的底层机制一次讲透包括怎么算、怎么配、踩过的坑怎么排。1. 训练显存的五个去向先搞清楚钱花在了哪1.1 参数、梯度、优化器状态三种最直接的“大头”问一个实际问题训练一个模型显存到底被谁吃掉了很多人的第一反应是“模型参数”。这话对了一半。真正吃显存的其实是三份数据参数本身weight、反向传播算出来的梯度gradient、以及优化器内部维护的状态optimizer states。参数所占空间最直观——模型有多少个参数每个参数几个字节一乘就出来。梯度呢模型反向传播时需要把 loss 对每个参数的导数暂存下来形状和参数一模一样所以它占的空间也和参数等量。优化器状态就容易被忽略了但它往往才是最大的开销尤其当你用 AdamW 这类自适应优化器时。为什么优化器状态这么大以 AdamW 为例它给每个参数额外保存两样东西一阶动量 m 和二阶动量 v都是 FP32 格式。对于混合精度训练还得再保存一份 FP32 的 master weight主权重。这三项加起来每参数 12 字节而参数本身用 BF16 存才 2 字节梯度 2 字节。一对比你就知道优化器状态有多大分量了。1.2 激活值随 batch size 和序列长度膨胀的“隐形开销”第四个大头是激活值activation values。前向传播每一层的输出需要保存在内存里供反向传播计算梯度时使用。这一部分和模型参数量的关系不大而是和你的 batch size、序列长度、隐藏层维度直接挂钩。它的增长模式很吓人batch size 翻一倍激活内存也差不多翻一倍序列长度翻一倍激活内存同样可能翻数倍。你可以把它理解成一次性的“过路货”——算完就扔但在算完之前必须完整存着。对于长序列场景激活内存完全可能追平甚至超过参数内存。1.3 通信缓冲区与显存碎片容易被低估的杂项第五类是杂项开销。多卡训练时梯度同步需要临时存储通信缓冲区像 NCCL 的 all-reduce 操作每张卡都要预留一块发送和接收数据的空间。此外还有显存碎片动态分配和释放过程中产生的空洞不会因为你操作完就自动消失。显存碎片在长训练任务中尤其烦人明明nvidia-smi显示还有 10GB 空闲程序却告诉你 OOM。原因就是大块连续内存被切碎了无法满足模型运行时对连续显存的请求。后面我会专门说怎么处理碎片问题。2. 显存估算公式从参数规模直接推算出“得用几块卡”2.1 不同优化器下的单位参数量成本搞清楚了显存去向估算就有章可循。先把“每多少个字节每参数量”这个基础数字记熟这套体系建立后任何模型都能快速估算。训练配置参数梯度优化器状态合计字节/参数FP32 SGD440纯 SGD 无动量8FP32 SGD带动量44412FP32 AdamW448m v16BF16/FP16 AdamW2212master m v16BF16 AdamW ZeRO-12212/NN 卡分片随时 N 缩小BF16 AdamW ZeRO-32/N2/N12/N随时 N 缩小注意一个关键细节混合精度训练省显存重点并不在参数和梯度那几字节而是省了激活值可用 FP16 半精度存储同时保证优化器状态依然用 FP32 维持训练稳定性。如果你用 BF16 AdamW 但不开 ZeRO光参数梯度优化器状态每参数还是要 16 字节和纯 FP32 AdamW 几乎一样。这是很多人对混合精度的第一个误解。2.2 实际计算公式与示例13B 模型到底需要多少显存用 13B 模型做例子算一笔账。模型 130 亿参数混合精度 AdamW每参数 16 字节那么运行权重、梯度、优化器状态的静态开销就是13 × 10^9 × 16 字节 ≈ 208GB是的光这三样就要 208GB还没算激活值、通信缓冲和中间碎片。这就是为什么 13B 单卡微调即便用 BF16 也要把 batch 调很小否则直接爆显存。加上激活值我给一个工程上的粗略公式总显存需求 ≈ 参数量 × 16字节静态部分 batch_size × seq_len × hidden_size × num_layers × 激活系数激活系数通常在 2 到 20 之间取决于是否开梯度检查点、是否存 FP16、实现细节等。开梯度检查点能把这个系数压到接近 1 到 2。保守估算时我一般把激活部分按静态部分的 20% 到 40% 算然后再加上 2GB 的通信冗余和碎片余量。算出来的结果再比对照你手头 GPU 的显存用整除确定需要几张卡、要不要上 ZeRO。例如 13B 模型静态 208GB四张 80GB 卡总共 320GB 显存那大概率够用但如果是两张 80GB 卡只有 160GB就必须开 ZeRO-1 或 ZeRO-2 来分片优化器状态了。2.3 脚本实测用 PyTorch 快速清点模型显存纸上算完还要实测验证。你在本地用一段小代码就能量出模型权重到底占多少字节。import torch def count_params_and_bytes(model): total_params 0 total_bytes 0 for name, param in model.named_parameters(): if param.requires_grad: n param.numel() total_params n total_bytes n * param.element_size() if total_params 5_000_000: print(f{name}: {n} params, {param.element_size()} bytes/elem) return total_params, total_bytes model get_your_model() total_params, total_bytes count_params_and_bytes(model) print(fTotal params: {total_params:,}) print(fModel weights memory: {total_bytes / 1024**3:.2f} GB)跑完后打印出来的权重块配合上面那张每参数成本表能反推更精细的需求。还有两个 API 在训练时值得盯torch.cuda.memory_allocated()显示当前实际分配的显存torch.cuda.max_memory_allocated()显示到目前为峰值。把这两行代码放在一个训练 step 的首尾每次打印会对显存随 batch 变化的趋势有很直观的感受。3. 混合精度训练的原理与选型为什么 BF16 是训练首选3.1 FP16、BF16、FP32表示范围和精度的区别混合精度训练的核心是让计算和存储用低精度格式同时又避免精度过低导致训练崩溃。从底层原理看FP16 和 BF16 都是 2 字节数据类型但分配方式完全不同。FP16 有 5 位指数位和 10 位尾数位表示范围大约在 6e-8 到 65504 之间。范围小是它的硬伤一旦数值超出 65504 就会溢出为无穷大inf。BF16 则是 8 位指数位加 7 位尾数位指数范围和 FP32 几乎一致因为 FP32 也是 8 位指数但尾数只有 7 位——精度低得多可范围要安全得多。训练场景中梯度的大小往往起伏很大FP16 的窄范围几乎是天然劣势必须靠 loss scaling 硬撑。而 BF16 因为范围和 FP32 一致不需要太多额外保命机制就能稳定训练。这就是为什么当 A100、H100、以及新出的消费级显卡支持 BF16 之后训练社区迅速转向 BF16 的原因。3.2 master weight 与 loss scaling混合精度的两个核心机制混合精度训练不是“所有数值都用半精度”这么简单。整个机制里最容易被忽视也最重要的两件事master weight以及针对 FP16 的 loss scaling。master weight 是说模型需要保留一份 FP32 格式的权重副本训练过程中用它来更新参数再把它转成半精度用于前向和反向计算。为什么不能直接在半精度权重上更新因为半精度的尾数太短一次学习率的微小增量可能比尾数能表示的最小步长还要小更新了几百上千个 step 后误差累积起来loss 就完全不收敛了。所以 master weight 相当于一个高精度“账本”每轮计算完再以低精度副本参与训练。loss scaling 则是 FP16 训练专属手段。反向传播算出的梯度普遍数值很小FP16 能表达的最小正数有限极小梯度会直接变成 0。处理方法是把 loss 乘一个大系数比如 1024梯度整体放大到 FP16 可表示的范围再反向传播等优化器取到梯度后除以同样的系数恢复真实大小。PyTorch 的GradScaler会在训练过程中动态调整这个缩放系数——当发现某个 step 的梯度溢出为 inf就把系数调小持续若干 step 没溢出再尝试调大。3.3 什么场景用 FP16、什么用 BF16做选择之前先看你的 GPU 支持情况。RTX 3090、V100 这类老一点的卡原生支持 FP16但对 BF16 不友好或根本加速不了。A100 及以后的服务器卡基本都原生支持 BF16。消费级 RTX 4090 也支持 BF16但也要确认是原生计算还是模拟。如果显卡不支持 BF16那就用 FP16 loss scaling靠 GradScaler 动态维护数值范围。支持 BF16 时我强烈优先选 BF16。理由很简单省心。你不用整天为 loss 爆炸、loss 变 NaN 发愁可以把精力放到模型本身的问题上。BF16 尾数少这个缺点在大多数 LLM 训练任务中并不致命因为训练的收敛主要仰仗优化器的累加和 master weight 的纠偏。一句话总结选型标准能用 BF16 就用 BF16原生不支持再退到 FP16千万别用半精度把所有值都直接降了——那是性能灾难也是炼丹事故的源头。4. 实战PyTorch/DeepSpeed 混合精度训练配置与显存优化组合拳4.1 最小可用的 PyTorch AMP 训练示例在实际工程中PyTorch 的torch.autocast搭配GradScaler是上手最快的混合精度方案。我通常这样写训练循环import torch model model.cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-5) scaler torch.cuda.amp.GradScaler() use_bf16 torch.cuda.is_bf16_supported() for epoch in range(epochs): for input_ids, labels in dataloader: input_ids, labels input_ids.cuda(), labels.cuda() optimizer.zero_grad() dtype torch.bfloat16 if use_bf16 else torch.float16 with torch.autocast(device_typecuda, dtypedtype): loss model(input_ids, labelslabels).loss if not use_bf16: scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() else: loss.backward() optimizer.step()关键点在于BF16 模式下不需要 GradScaler直接普通 backward 和 optimizer.step() 即可。FP16 模式下则必须使用 scaler否则极小的梯度会被下溢掉训练直接停在原地。这个差异刚接触混合精度时很容易写错。4.2 DeepSpeed 与混合精度、ZeRO 的配合方案单卡能跑动 7B、13B 吗能但把 batch 压到 1 之后显存可能依然不够这时候就要上 ZeRO。ZeRO 的核心思想一句话说就是把模型训练过程的数据从“每张卡都存一份副本”变成“分片存储、集体通信”ZeRO-1/2 主要切优化器状态ZeRO-3 把参数、梯度、优化器状态全部切分。DeepSpeed 下的配置我直接给一份能跑的 JSON{ train_batch_size: 32, gradient_accumulation_steps: 4, gradient_clipping: 1.0, bf16: { enabled: true }, zero_optimization: { stage: 1, allgather_partitions: true, reduce_scatter: true }, optimizer: { type: AdamW, params: { lr: 1e-5, betas: [0.9, 0.999], eps: 1e-8 } }, scheduler: { type: WarmupDecayLR, params: { warmup_min_lr: 0, warmup_max_lr: 1e-5, warmup_num_steps: 100, total_num_steps: 10000 } } }这条配置的含义是bf16.enabled true开启混合精度zero_optimization.stage 1只切分优化器状态适合显存差点意思但不多的情况gradient_accumulation_steps 4把 4 个小 batch 的梯度累加后再更新一次相当于扩大了有效 batch size还不需要额外增加显存。注意一点DeepSpeed 里fp16和bf16两个配置只能选一个不能同时开启。用 FP16 时fp16.initial_scale_power一般设成 32loss_scale_window设成 1000 左右让动态 loss scaling 在一个合理范围内波动。4.3 梯度检查点与激活重计算牺牲速度换显存开完混合精度 ZeRO-1模型可能还是差一口气而你想省显存最简单的手段是开梯度检查点gradient checkpointing又叫 activation checkpointing 或 activation recomputation。原理一句话前向传播时不保存每一层的激活值只保存一小组“检查点”反向传播需要某一层激活时临时把之前的前向路径重新算一遍。代价是大约 10% 到 30% 的训练速度损失换来的是激活内存的直接“减半再减半”。实际跑大模型时我的经验是优先开它尤其在序列长度长、batch size 又压缩不了的场景里比降 batch 更划算。PyTorch 里的打开方式也很直接model.gradient_checkpointing_enable()HuggingFace 的 Transformer 模型基本都内置该方法。DeepSpeed 配置里也可以通过设置activation_checkpointing片段启用。实测一个 7B 模型显存峰值可能从 60GB 掉到 40GB 左右7B 就拥有了在单张 48GB 卡上训练的空间。这个“速度换显存”的 trade-off在大模型时代的收益非常可观。4.4 显存实测案例7B 模型三种方案对比分享一组我实测过的显存占用数据模型是 7Bbatch size 2序列长度 2048单卡训练A100 80GB 环境。不同方案组合下峰值显存差异明显。方案权重梯度优化器状态激活内存实测峰值显存每 step 耗时全 FP32不开 ZeRO约 112GBOOM无法评估OOM无数据BF16 ZeRO-1约 28GB约 30GB约 65GB约 21sBF16 ZeRO-1 梯度检查点约 28GB约 8GB约 42GB约 26sBF16 ZeRO-3 梯度检查点约 12GB约 8GB约 26GB约 33s从这组数据能看出单纯开混合精度对静态数据权重梯度优化器状态的削减有限真正的“显存杀手”往往在激活值和优化器状态上。想要极致的显存控制就得混合精度 ZeRO 梯度检查点三管齐下。5. 训练中的显存与精度问题排查实录5.1 Loss 变 NaN/Inf八成跟混合精度有关训练到一半 loss 突然变成 NaN这大概是混合精度训练里出现频率最高的事故。很多人第一反应是调学习率但我建议先按下面顺序排查一圈。先看是不是 FP16 的溢出问题。检查scaler.get_scale()输出的缩放系数如果它一直在自动下降说明梯度频繁溢出。对治手段是把loss_scale_window调大或者手动把initial_scale_power从 32 降到 24也能换取更保守的范围。第二个常见根源是学习率过大大模型训练里混合精度会导致每个 step 的有效更新量变大原来 FP32 下能跑的 1e-4 在混合精度下可能就崩了先降一半再继续观察。第三个可能是数据里有异常值比如 label 出现 inf或者 embedding 层输入没归一化。如果你已经切到 BF16那 NaN 的概率本来就低很多真出现了优先查模型结构或数据别再怀疑是混合精度框架的锅。5.2 OOM 了怎么办三步排查法OOM 报错人人都遇到过排查也有套路别一上来就降 batch。第一步看torch.cuda.memory_allocated()和torch.cuda.memory_reserved()。前者是真正在用的显存后者是缓存池保留的大小。如果reserved远超allocated说明主要问题是碎片或者缓存不释放。试试设置环境变量export PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True这个选项让 PyTorch 用可扩展内存段来分配能显著缓解碎片问题。第二步把梯度检查点打开把激活重计算的收益吃下来。第三步才轮到调 batch size配合gradient_accumulation_steps把有效 batch 补回来。这三步走完绝大多数单卡和单节点 OOM 都能解决。还有一个细节多卡训练时通信缓冲区也会占显存。减少通信峰值的方法是把梯度同步操作拆小或者换allreduce的通信后端但最简单的办法其实是再开一个zero_optimization.stage。通常 stage 1 升到 stage 2 能多省一点一块stage 2 升到 stage 3 能省更多但通信量会明显增加训练会变慢得权衡。5.3 常见问题速查表把我在实际训练中遇到最多的几个问题整理成一张速查表熟读这张表能帮你省下大把 debug 时间。症状可能原因解决手段Loss 在某个 step 突然变 NaN/InfFP16 loss scale 过高导致梯度溢出调大loss_scale_window或改用 BF16或降低学习率训练速度慢但显存充足梯度检查点重计算开销过大只对部分层开启 checkpoint或减少 checkpoint 密度step 一开始就直接 OOM激活值计算量巨大开梯度检查点或者把序列长度先缩短到一半验证reserved多但allocated少显存碎片化缓存设置expandable_segments或定期重启训练进程多卡训练时速度和显存都不理想梯度同步通信量大优先开 ZeRO-2/3或用梯度累积延长同步周期FP16 训练比 BF16 loss 抖动严重FP16 动态范围窄换成 BF16或调大initial_scale_power并加强 clipping5.4 一个容易忽略的细节控制变量比调参更重要做这些排查时最重要的一条是每次只改变一个变量。我在调显存优化时吃过亏同时开了梯度检查点、换了 ZeRO stage、又调了 batch size结果训练速度暴跌根本分不清是哪一项的影响。后来学乖了每次只动一个开关记录max_memory_allocated()和每 step 耗时测试四五轮后再根据数据决定保留哪些优化项。这个习惯看起来笨但在大模型训练这种纯试错成本极高的场景里反而最节约时间。回归到混合精度和显存本身我个人在实际操作中的体会是显存估计永远要有 20% 的安全余量。理论上算出来 62GB你别真拿 64GB 的卡去跑因为 PyTorch 缓存池、CUDA context 和 NCCL 都会额外吃掉几个 GB。预先留好余量、把估算公式做成习惯再配合混合精度和 ZeRO 这把组合拳大模型训练就算搬到一张消费级显卡上也不是什么不可能完成的任务。接下来我大概率会写一写多卡分布式训练里的通信开销怎么优化那又是一个全新的坑。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门