新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/29 20:53:10来源:尧图网络
大模型训练显存估算与混合精度实战:从OOM到BF16调优
很多人第一次真正上手大模型训练第一个感觉往往是“这玩意怎么这么吃显存”我当年在公司搭预训练环境时拿到一张A100 80G心想怎么也能跑个10B模型吧结果连一个7B模型的前向都过不去。后来才知道显存不是只给模型参数用的参数、梯度、优化器状态、激活值、通信缓冲、CUDA context每一项都在跟你要空间。而混合精度训练就是解决这个问题的第一把钥匙。这篇内容我围绕“显存估计”和“混合精度训练”两个主题展开梳理清楚大模型训练时的显存到底花在哪、怎么在开训前就算出大概用量、FP16/BF16混精到底怎么工作、以及遇到OOM和Loss NaN时怎么排查。适合正准备训自己的模型、或者在公司搭训练脚本遇到瓶颈的人不涉及太多分布式并行细节先把单机多卡下的显存账算明白后面做DeepSpeed、FSDP也不会发怵。1. 显存花在哪先搞懂训练时显存的四个去向1.1 训练比推理吃的显存多在哪很多人一开始是用推理的思维理解显存的。推理时模型参数是只读的你只需要把权重加载进显存跑前向计算就行所以7B模型用BF16加载14GB显存就跑起来了。但训练完全不同训练时你需要反向传播计算梯度还要用优化器去更新权重这意味着同样的参数训练阶段要在显存里放好几份。这里有个生活化类比推理就像去餐厅点菜菜前向计算结果端上来你吃掉就行后厨只需要一份食材训练则是你自己开餐厅不仅要存食材参数还要准备切菜用的砧板梯度以及放调料配方的地方优化器状态。所以同一份7B模型推理14GB够用训练呢光静态开销就轻松上百GB。这就是为什么大模型训练几乎必须上多卡分布式、或者做各种Offload。还有一个很多人忽略的点训练时CUDA context本身也会占显存通常在几百MB到1GB左右多卡环境下每张卡都有一份。虽然不大但在你差最后几百MB显存的时候它就是你OOM的元凶之一。1.2 静态开销参数、梯度、优化器状态的精确账本静态开销是指不管batch size为1还是32都会固定占用的那部分显存。以最常用的AdamW优化器 混合精度训练为例每个参数实际占用的显存是这样分的模型参数本身以FP16或BF16存储2字节/参数梯度常见实现里也是FP16或BF162字节/参数有些框架用FP32下面会提优化器状态AdamW需要保存一份FP32主权重副本、一阶动量m、二阶动量v这三项各4字节/参数合计12字节/参数所以混合精度训练下每个参数大约要占2 2 12 16字节。参数量N的模型静态显存就是16 × N字节。7B模型7 × 1e9 × 16字节 ≈ 112GB。13B模型约208GB。70B模型约1120GB也就是超过1TB。如果你用的是纯FP32训练参数4字节、梯度4字节、优化器12字节总计也是16字节/参数显存一点没少只是计算精度高。这就是为什么很多人误以为“混合精度训练就是为了省显存”其实它对显存的贡献主要在于激活值和通信量而不是静态参数占用。它真正明显省下的是前向/反向使用半精度后激活值减半以及通信时梯度体积减半。静态开销这块混合精度相比纯FP32并没有数量级优势。如果一个框架告诉你它能“节省8倍显存”那一定不是单纯靠混合精度而是靠ZeRO分片、Offload这些分布式策略。这些策略我会在第3.3节再讲。这里先给一个常见模型的静态显存速查表方便你心里有底模型规模参数量纯FP32静态开销BF16AdamW静态开销单卡80G需要卡数7B7e9112GB112GB2张不含激活13B13e9208GB208GB3张不含激活70B70e91120GB1120GB14张不含激活注意这张表还不含激活值真实训练需要的显存只会更多。你细品一下为什么单卡不Offload训7B会OOM就明白这个账了。1.3 动态大头激活值才是最不可控的变量如果说静态开销是固定房租那激活值就是浮动水电费。激活值是前向传播过程中每一层临时算出来的中间结果比如Attention的Q/K/V、MLP层的中间输出、LayerNorm的均值方差等。反向传播时要用这些值去算梯度所以在反传结束前它们都得留在显存里。激活值大小和很多因素有关batch_size、序列长度、模型隐藏层维度、层数、注意力头数甚至Attention的实现方式。一个粗略的经验公式是激活值显存 ≈ batch_size × seq_len × hidden_size × num_layers × C其中C是一个经验系数通常可以取12到24取决于是否使用FlashAttention、是否开启激活重算、是否用半精度存储。比如一个隐藏层4096、32层的7B模型batch size为4、序列长度2048用BF16C取12来粗算4 × 2048 × 4096 × 32 × 12 × 2字节 ≈ 96.4GB你没看错这已经是一个很大的数了。如果不开激活重算batch size稍微调大一点显存就爆给你看。所以激活值才是训练显存里最不可控的部分也是你在调batch大小时最需要盯住的变量。这里就出现了两个最常用的手段激活重算Gradient Checkpointing和更省内存的Attention实现。激活重算的思路很简单前向时不去存每一层激活只存少量必要信息反传时重新算一遍激活。计算量大概多出30%左右但能把激活显存从层数 × 每层激活降到单层激活级别。另一个是FlashAttention它通过分块计算减少显存中Attention矩阵的存储。这两个手段几乎是现代大模型训练脚本的默认配置我建议只要序列长度超过1024就直接都打开。2. 混合精度训练FP16/BF16为什么能省显存还提速2.1 FP16、BF16、FP32别只看“半精度”三个字混合精度训练的基础是FP16和BF16但我发现身边经常有人把这两个搞混。它们在显存上都占2字节但数值表示方式完全不同。FP32有8位指数、23位尾数FP16只有5位指数、10位尾数。这意味着FP16能表示的数值范围很窄最大只有65504一旦梯度或者中间计算结果超过这个范围直接变Inf尾数只有10位大数加小数时容易把小数“吃掉”造成精度损失。BF16则不同它用8位指数、7位尾数指数范围和FP32完全一样范围不受限但尾数更少精度更低。用尺子类比FP32是游标卡尺精细但占地方FP16是毫米刻度尺常用但量程短稍微大点的东西就测不了BF16是那种量程很大但刻度很粗的工程卷尺它测量范围足够广但读数的精细程度不如毫米尺。对大模型训练来说BF16比FP16其实更重要因为训练过程中梯度的数值范围波动很大FP16需要靠Loss Scaling来避免溢出BF16则基本不用操心范围问题尾数精度不够的问题可以靠优化器里维护FP32主权重来弥补。这就引出了下一个话题。2.2 混合精度的三件套半精度计算、主权重、Loss Scaling混合精度训练不是一个简单的“用FP16代替FP32”它包含三个核心机制第一前向和反向计算用FP16或BF16执行。这能减少显存中激活值的占用同时利用GPU上Tensor Core对半精度运算的加速能力让计算速度大幅提升。第二优化器里维护一份FP32主权重副本。模型权重在前向时被转成FP16/BF16使用但优化器更新时是在FP32主权重上做动量更新和权重衰减再转回FP16/BF16。这样既能享受半精度计算的速度又不会让参数精度损失累积到无法收敛。第三Loss Scaling。FP16的数值范围太窄反传时梯度容易下溢成0所以要把Loss乘一个放大系数比如65536让梯度在FP16表示范围内保持合理大小反传完成后再把梯度缩小回去。PyTorch里的GradScaler会自动做这件事动态调整缩放因子如果梯度溢出就减小Loss Scale如果一段时间没溢出就增大。如果你用BF16理论上Loss Scaling的需求会小很多因为BF16的范围和FP32一致。实际操作里BF16训练几乎不需要做额外的梯度缩放但也要注意学习率过大导致的数值发散。所以现在很多框架默认就直接上BF16省心。2.3 混合精度到底帮你省了多少显存混合精度对显存的贡献主要在两方面一是激活值减半二是分布式训练时通信量减半。以1.3节那个7B模型为例如果你用FP32存激活96GB可能会变成192GB直接多出一张卡而用BF16激活就保持在96GB级别。但这里必须说清楚混合精度不能让你把一个7B模型塞进单卡80G。静态开销112GB摆在那光参数、梯度、优化器状态就放不下了。所以混合精度搭配的其实是“让同样的显存装下更大batch、更长序列”这件事而不是“让模型体积缩小”。真正解决模型放不下的是ZeRO分片和Offload。另外还有一个容易被忽略的点混合精度训练时如果某个算子不在autocast覆盖范围内比如LayerNorm或者Softmax里的某些操作框架仍然会用FP32计算这些地方在显存里依然是FP32尺寸。所以你实际看到的显存节省不会是理论上的精确半幅而是大致接近。不要拿公式去硬套要以nvidia-smi和PyTorch的显存统计为准。3. 实操训练前怎么把显存算得八九不离十3.1 一个可复用的手工估算脚本训练前做显存估算可以让你提前知道单卡能不能跑、需要几张卡、以及设置多大的batch。我自己习惯写一个很小的Python脚本在模型还没加载前就先算一遍静态开销代码大致长这样def estimate_static_memory(params_billion: float, use_bf16: bool True, optimizer: str adamw): params params_billion * 1e9 if use_bf16: param_bytes 2 else: param_bytes 4 model_mem params * param_bytes grad_mem params * 2 # 常见实现梯度也是bf16/fp16 if optimizer adamw: optimizer_mem params * 12 # fp32主权重 m v elif optimizer sgd: optimizer_mem params * 4 # 只有fp32主权重 total (model_mem grad_mem optimizer_mem) / 1024**3 print(f模型参数: {params_billion}B) print(f模型权重: {model_mem / 1024**3:.1f}GB) print(f梯度: {grad_mem / 1024**3:.1f}GB) print(f优化器: {optimizer_mem / 1024**3:.1f}GB) print(f静态总计: {total:.1f}GB) return total estimate_static_memory(7)跑出来的结果就是我上面说的112GB级别。这还没算激活值和CUDA context所以你在多卡训练时就得把总显存减去这个静态值再来看剩下能分给激活值的空间。通常我还会额外预留20%的显存缓冲区因为PyTorch的显存管理器不会把所有显存都用尽碎片化也会吃掉一部分空间。更进一步我建议直接在训练循环里挂一个峰值显存记录器import torch # 在某次迭代之后 peak torch.cuda.max_memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 print(f分配峰值: {peak:.2f}GB, 预留峰值: {reserved:.2f}GB)memory_reserved是PyTorch向CUDA申请的总量memory_allocated是实际使用的量。两者差值就是PyTorch显存缓存池里的空余这个差值太大通常说明显存碎片化严重或者你分配了太多不连续的小张量。3.2 训练过程中的实测与校准理论估算再准不如直接跑一遍看实测。我的做法是先估算一个“绝对安全的batch size”比如batch为1、seq_len为256跑通一个step记录峰值显存然后逐步把batch size翻倍观察显存增长曲线。这里有个实用经验显存随batch size的增长接近线性随序列长度的增长也接近线性但如果你把序列长度从1024加到2048Attention矩阵的显存是按照seq_len^2增长的所以显存增长会明显加速。当观察到显存增长速度远超线性时基本可以判断是Attention矩阵在作怪这时候优先开FlashAttention或激活重算比单纯缩减batch更划算。多卡训练时nvidia-smi显示的是每张卡的显存占用但这个数字包含CUDA context和通信缓冲而且有一定延迟。我更推荐在代码里用torch.cuda.memory_summary()查看指定设备的详细分配情况或者直接打印torch.cuda.max_memory_reserved(device)来看某张卡在训练中的真实峰值。实测下来nvidia-smi的数值和PyTorch内部统计之间经常会有1到2GB的差异不要因此惊慌。另一个值得提的细节Dataloader的num_workers开太多也会在数据预处理时占用额外CPU内存但一般不影响显存。不过如果你用了pin_memoryTrueCPU上的固定内存缓冲区也会变大多卡时会增加主机和设备间的拷贝压力这个不是显存问题但会导致内存吃紧。3.3 显存不够时的优先级清单当你发现显存确实放不下时不要第一反应就上多卡。我整理了一个处置优先级按性价比从高到低排列检查是否已开启激活重算Gradient Checkpointing。这个选项往往能把激活值从十几GB降到几GB代价是训练时间变长。尤其长序列场景收益非常明显。检查是否用了FlashAttention或类似的Memory-Efficient Attention。如果代码支持基本是零成本收益。把batch size调小用梯度累积凑有效batch size。注意梯度累积不会减少显存峰值它只是减小单次前向和反传的batch从而降低激活值。如果你原来batch为4爆显存改成batch为1并累积4步有效batch还是4但激活值只有原来的四分之一。优化器从AdamW换成Adafactor或者狮身人面像优化器比如狮身人面像其实就是Sophia能省掉一部分优化器状态显存但对模型收敛的影响要实验验证。使用ZeRO Stage 1或Stage 2把优化器状态和梯度分片到多张卡上。这是7B模型单机多卡训练的标准配置。CPU Offload把优化器状态放到内存里显存能省一大截但代价是训练速度明显下降。只能作为最后的兜底。这个顺序背后的逻辑是先压缩激活值再调整batch最后才动优化器和参数存储。因为前三项几乎不影响模型收敛行为后面两项会改变训练动态需要重新验证超参。4. 常见问题与排查OOM和Loss NaN的实战处理4.1 OOM爆显存最常见的几个原因OOM是训练大模型时最常遇到的报错但原因往往各不相同。我遇到过的典型情况有这么几类第一种是激活值峰值爆掉。现象是训练到某个step时报OOM但前面几个step都好好的。此时优先怀疑序列长度变化或者Attention实现导致的峰值内存。排查方式是缩小batch size到1如果还能跑基本就是激活值问题如果batch为1都OOM那就是静态开销问题。第二种是显存碎片化。PyTorch的缓存分配器是按块申请显存的如果训练过程中频繁创建和销毁不同大小的张量显存里会出现很多碎片总剩余空间够但连续空间不够最终OOM。这种情况可以设置环境变量export PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:64把分配拆成更小的块碎片化问题通常会缓解。我自己的经验是这个设置在训练早期效果明显训练后期因为显存总量紧张帮助有限。第三种是被忽视了CUDA context和其他库的占用。有些库加载时会申请额外显存比如flash-attn的不同版本、triton的kernel缓存都会吃几百MB到1GB的显存。如果你离OOM就差几百MB试着把不用的库延迟导入能省一点是一点。第四种是分布式通信缓冲。使用torch.distributed时NCCL会分配通信缓冲区多卡训练时这部分显存不可忽略。如果你在单卡测试时不OOM但多卡训练时OOM先检查NCCL缓冲区和torch.cuda.set_per_process_memory_fraction这样的设置。我给一个通用的处置速查表方便你快速定位现象可能原因处置方法起步即OOMCUDA context或参数静态开销超限开ZeRO/Offload调小模型精度跑一会儿才OOM激活值峰值、长序列Attention开激活重算、FlashAttention、减batch偶尔OOM重启后变好显存碎片化设置max_split_size_mb换更小的分配块多卡OOM单卡正常NCCL通信缓冲调NCCL缓冲大小减batch训练后期OOM缓存池增长、梯度累积后的峰值累积定期torch.cuda.empty_cache()但要谨慎使用4.2 Loss NaN/Inf与Loss Scale的相爱相杀混合精度训练下Loss变成NaN或Inf是很常见的事情。我一开始遇到这个情况第一反应是调低学习率但其实原因可能完全不同。如果你用的是FP16要首先看GradScaler的Loss Scale状态。PyTorch的GradScaler会在梯度溢出时自动减小Loss Scale如果连续溢出Loss Scale会一路降到很低此时模型参数更新幅度可能已经失衡Loss直接发散。排查方法是打印scaler.get_scale()如果它在一路下降那说明模型里有数值溢出的操作而不是学习率的问题。常见解决办法包括把初始Loss Scale调大、把inf_scale调小、或者干脆换用BF16。如果你用的是BF16Loss NaN的原因多半不是下溢反而是上溢的对称问题BF16虽然范围大但尾数精度低当你做Softmax或者LayerNorm时如果某些数值被削减到极少位数梯度里的噪声会被放大累积到一定程度就会出现训练不稳定。这种情况我比较建议直接开梯度裁剪把梯度的范数clip到1.0左右能挡掉很大一部分发散问题。还有一个排查顺序可以分享先跑一个FP32的小规模基线用同样数据和超参如果FP32也NaN那就是模型结构或数据问题跟混合精度无关如果FP32正常、混合精度异常那才是精度相关的问题。我每次遇到NaN第一件事就是开一个纯FP32的小实验做对照这一步能帮你省下大量无头绪的调试时间。4.3 显存估算偏差大的幕后黑手有时候你按1.3节的公式估算显存结果实际训练时发现差出好几GB这是正常的。主要偏差来源是这些地方第一PyTorch的缓存池。torch.cuda.memory_reserved()往往大于实际使用量因为PyTorch会多预留一些显存避免反复向CUDA申请释放。所以你用nvidia-smi看到的“已用显存”其实是预留值不是你模型实际占用值。第二CUDA context和cuDNN算法选择。cuDNN会根据输入尺寸选择不同卷积算法这个算法搜索过程会临时申请显存可能爆掉你的预算。可以在代码里设置torch.backends.cudnn.benchmark False减少这种搜索开销代价是部分算子速度略降。第三框架自动混合精度策略的细节。比如torch.autocast默认对每个算子单独决定是否用低精度有些算子仍然用FP32计算这些“漏网”的FP32激活会额外占用显存。你粗略算出来的“半精度显存”和实际值有出入很正常。所以我一直强调估算只用来做“方向判断”上线训练前一定要用真实模型、真实batch、真实seq_len跑一两个step拿实测峰值来校准。训练脚本里挂一个torch.cuda.max_memory_allocated()的日志是成本最低又能看清显存真相的做法。5. 我个人在实际训练中的一点体会最后说个实打实的经验。我最早自己训一个1.3B模型时用BF16混合精度满心以为显存能比FP32省下一半结果发现省下来的空间远没有想象中多速度提升倒是实实在在的。后来排查才发现我的batch size设得特别小显存大头全在静态开销和频繁的kernel启动带来的缓存碎片上。后来把batch调到能塞满显存的级别再配合梯度累积和激活重算训练吞吐直接翻倍。如果你现在正准备训一个新模型我的建议是把显存估算脚本跑一遍确定静态开销再预估激活值区间用保守batch先跑通一个step记录实测峰值再逐步放大。这套流程看着简单却能帮你少走我当年走的弯路。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

DeepMind Lab Python 模块:从 Bazel 构建到 pip 安装的完整实战指南 2026/9/29 21:33:30

DeepMind Lab Python 模块:从 Bazel 构建到 pip 安装的完整实战指南

人工智能强化学习机器学习 【免费下载链接】lab A customisable 3D platform for agent-based AI research 项目地址: https://gitcode.com/gh_mirrors/la/lab 点击查看 免费下载 导读 DeepMind Lab 是一个面向智能体(agent)研究的可定制 3…

阅读更多 →
树莓派4B玩转RPLIDAR C1:从串口配置到ROS2 SLAM建图完整指南 2026/9/29 21:33:29

树莓派4B玩转RPLIDAR C1:从串口配置到ROS2 SLAM建图完整指南

说实话,这块RPLIDAR C1激光雷达和树莓派4B的组合,折腾得我比预想中久不少。去年年底拿到C1的时候,我第一次操作还是在X86笔记本上,插上USB转串口,make完官方SDK,三分钟就看见了角度和距离数据刷屏。可等我把…

阅读更多 →
VoxelFM:学习稳健的 CT 视觉特征,实现面向临床任务的高效迁移 2026/9/29 21:33:23

VoxelFM:学习稳健的 CT 视觉特征,实现面向临床任务的高效迁移

论文:Learning Robust Visual Features in Computed Tomography Enables Efficient Transfer Learning for Clinical Tasks 作者:Rubn Moreno-Aguado、Alba Magalln、Victor Moreno、Yingying Fang、Guang Yang arXiv:2604.04133v1&#xff0…

阅读更多 →
35岁后端转Agent一年,说几句得罪人的话 2026/9/29 21:33:23

35岁后端转Agent一年,说几句得罪人的话

标题写了"得罪人",就得真得罪。下面这几句,可能会让一些人不舒服,但都是我转型这一年多亲眼看到的。 不是为了骂谁,是因为我自己也踩过这些坑,被这些内容误导过时间、走过弯路。 第一句:天天发&…

阅读更多 →
关于画眉App开发的好处及相关功能介绍 2026/9/29 21:33:23

关于画眉App开发的好处及相关功能介绍

爱美是女性的天性,所以她们会注重生活中的每一个细节,尤其是关于化妆方面。我们都知道化妆并不是只单独的涂涂口红,眉毛也是很重要的一部分,但是很多人并没有掌握真正的画眉技巧,所以在画眉上并未却有效的进展。面对传…

阅读更多 →
不会 PS 怎么做公众号封面?多款 AI 绘图工具对比推荐 2026/9/29 21:33:22

不会 PS 怎么做公众号封面?多款 AI 绘图工具对比推荐

在内容创作日益高频的今天,公众号封面作为文章的“门面”,直接影响打开率与传播效果。然而,许多运营者和创作者并不精通 Photoshop,面对复杂的设计软件往往望而却步。借助 AI 绘图工具,无需专业设计基础也能快速产出高…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉