
大模型训练绕不开两个硬骨头显存不够和精度怎么选。我见过太多团队在单卡上跑7B模型刚加载完权重就OOM也见过有人把FP16当万能药结果训练到一半loss直接炸成NaN。这篇就围绕显存估计和混合精度训练这两个核心问题把账算清楚把坑标明白。不管你是刚接触大模型训练的新手还是已经调过几轮参数的老手这里面的估算方法和精度选择逻辑都值得过一遍——因为这两个东西直接决定了你能不能在有限的硬件上把模型跑起来、跑稳。1. 显存到底被谁吃掉了很多人估显存的方式特别粗暴参数量乘以2FP16或者乘以4FP32然后买个对应显存的卡。这么算在推理场景下勉强能用但训练场景下会错得离谱。训练时的显存占用是一个多组件叠加的结果漏掉任何一块都可能导致实际跑起来直接爆掉。1.1 训练显存的四大消耗方把训练时的显存拆开看主要是四块模型参数本身。这部分最好理解FP32下每个参数占4字节FP16/BF16下占2字节。一个7B模型FP16加载权重需要约14GBFP32则需要约28GB。梯度。反向传播需要为每个可训练参数存储对应的梯度。梯度通常和参数保持相同的精度所以FP16训练时梯度也是2字节每参数7B模型约14GB。优化器状态。这是最容易被低估的部分。以Adam/AdamW为例它需要为每个参数维护一阶动量momentum和二阶动量variance如果优化器状态用FP32存储那就是每个参数8字节。7B模型光优化器状态就要56GB。这就是为什么很多人用Adam训练大模型时显存直接爆炸——优化器状态比模型本身还大。激活值Activations。前向传播过程中每一层的中间输出都需要保留供反向传播计算梯度使用。这部分的大小和batch size、序列长度、模型层数、隐藏维度都强相关而且往往是训练显存中占比最大、最难精确估计的一块。提示很多人只算参数和梯度忽略了优化器状态和激活值结果实际显存需求是估算值的3到5倍。这是新手最常踩的坑。1.2 一个可落地的显存估算公式基于上面的拆解我整理一个实操中比较靠谱的估算框架。假设模型参数量为 ( P )单位个训练精度为混合精度参数和梯度用FP16优化器状态用FP32则组件精度每参数字节数7B模型占用模型参数FP16214 GB梯度FP16214 GB优化器状态AdamFP32856 GB激活值FP16与batch/seq相关视配置而定合计不含激活-1284 GB激活值的估算更复杂一些。一个粗略的经验公式是激活值显存 ≈ batch_size × seq_len × hidden_size × num_layers × 系数这个系数取决于具体的模型架构是否有GQA、是否使用FlashAttention等通常在10到20之间。以7B模型hidden_size4096num_layers32、batch_size1、seq_len2048为例激活值大约在2.7GB到5.4GB之间。如果把batch_size提到8这部分就会涨到20GB以上。所以一个7B模型在混合精度下训练不含激活就需要约84GB显存加上激活值轻松超过100GB。这就是为什么单卡训练7B模型基本不现实必须上多卡并行或者用ZeRO之类的优化技术。1.3 激活重计算用时间换空间的经典操作激活值太大怎么办最直接的办法是激活重计算Activation Checkpointing也叫Gradient Checkpointing。它的思路很简单前向传播时不保存所有中间激活值只保存少数几个检查点的激活值反向传播需要用到某个激活值时从最近的检查点重新做一次前向计算把它算出来。这样做的好处是激活值显存可以降低到原来的 ( \sqrt{N} ) 左右N为层数代价是训练速度会慢20%到30%因为多了一次前向计算。在实际操作中如果你的显存刚好差一点不够开激活重计算是最省事的方案。PyTorch里几行代码就能开启from torch.utils.checkpoint import checkpoint # 在模型forward中对每个transformer block使用checkpoint def forward(self, x): return checkpoint(self._forward, x)注意激活重计算和FlashAttention可以叠加使用两者不冲突。FlashAttention本身已经大幅降低了注意力部分的激活值配合重计算能把整体激活值压到很低。2. FP16和BF16的本质区别搞清楚了显存去哪了接下来要解决精度选择的问题。FP16和BF16是混合精度训练中最常用的两种格式很多人知道BF16比FP16更稳但说不清楚为什么。这里把两者的底层表示掰开讲。2.1 从浮点数的位布局说起一个浮点数由三部分组成符号位、指数位、尾数位。FP16和BF16的总位数不同各部分的分配也不同格式总位数符号位指数位尾数位动态范围FP32321823约10^-38 到 10^38FP16161510约10^-5 到 65504BF1616187约10^-38 到 10^38关键差异在指数位。FP16只有5位指数能表示的数值范围很窄最大只能到65504。BF16有8位指数和FP32完全一样所以动态范围和FP32一致。尾数位决定的是精度。FP16有10位尾数BF16只有7位。这意味着FP16在表示同一个范围内的数时精度比BF16高。但BF16的精度损失在深度学习训练中通常可以接受因为神经网络对权重的精度本身就不敏感。2.2 为什么FP16容易溢出而BF16不容易训练过程中梯度值可能非常小比如10^-8也可能在某些层突然变得很大。FP16的最小正规数约为6×10^-5比这个还小的梯度直接变成0下溢。而FP16的最大值是65504超过这个值就变成inf上溢。BF16因为指数位和FP32一样能表示10^-38到10^38的范围几乎不会出现上下溢的问题。这就是为什么用BF16训练时loss更稳定——不是BF16更聪明而是它的数值范围足够宽不会因为梯度太小或太大而丢失信息。实际训练中FP16通常需要配合损失缩放Loss Scaling来防止梯度下溢。原理是在计算loss时乘以一个大的缩放因子比如1024反向传播得到的梯度也相应放大更新参数前再除回去。这样梯度在FP16范围内就不会下溢。但损失缩放本身需要调参缩放因子太小起不到作用太大又会导致上溢。2.3 硬件支持情况决定你的选择理论上BF16更好但能不能用BF16取决于你的硬件。BF16需要硬件原生支持目前主流的训练卡如A100、H100、RTX 4090等都支持BF16。但一些较老的卡如V100、T4只支持FP16不支持BF16。所以选择逻辑很清晰硬件支持BF16优先用BF16省心不需要调损失缩放硬件只支持FP16用FP16 损失缩放需要多调一个参数硬件两者都支持但追求极致精度可以试FP16 损失缩放但调参成本更高提示在PyTorch中torch.cuda.is_bf16_supported()可以快速检查当前显卡是否支持BF16。这个检查在代码里加一行就能避免运行时才发现不支持的尴尬。3. 混合精度训练的实操配置知道了原理接下来看怎么在实际训练中配置混合精度。PyTorch提供了torch.cuda.amp模块用起来不算复杂但有几个细节不注意就会出问题。3.1 标准混合精度训练代码模板先给一个可以直接用的模板import torch from torch.cuda.amp import autocast, GradScaler # 初始化 model MyModel().cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) # 如果使用FP16需要GradScalerBF16不需要 use_bf16 torch.cuda.is_bf16_supported() scaler GradScaler(enablednot use_bf16) for batch in dataloader: optimizer.zero_grad() # 前向传播在autocast上下文中进行 with autocast(dtypetorch.bfloat16 if use_bf16 else torch.float16): outputs model(batch) loss compute_loss(outputs, batch) if use_bf16: # BF16直接反向传播 loss.backward() optimizer.step() else: # FP16需要scaler scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这段代码的核心逻辑是autocast上下文管理器自动决定每个操作使用什么精度。矩阵乘法、卷积等计算密集型操作会用FP16/BF16加速而softmax、layer norm等对精度敏感的操作会保持FP32。3.2 autocast的精度决策逻辑autocast不是简单地把所有东西都转成半精度它维护了一个操作白名单和黑名单会转为半精度的操作矩阵乘法torch.mm、torch.bmm、卷积torch.nn.Conv2d、线性层torch.nn.Linear等。这些操作计算量大半精度带来的加速效果明显而且对精度损失不敏感。保持FP32的操作softmax、layer normalization、loss计算、指数运算等。这些操作涉及数值稳定性问题用半精度容易出问题。这个自动决策机制是混合精度训练能work的关键。你不需要手动指定每个操作的精度autocast帮你做了。3.3 损失缩放的动态调整机制FP16训练中的GradScaler不是固定缩放因子而是动态调整的。它的工作流程是初始缩放因子设为一个大值默认65536每次反向传播后检查梯度是否有inf或NaN如果连续多个step没有出现inf/NaN增大缩放因子如果出现inf/NaN跳过这个step的参数更新减小缩放因子这个机制的好处是不需要手动调缩放因子但有一个副作用如果模型本身有问题导致梯度经常溢出缩放因子会不断减小最终失去作用。所以如果发现GradScaler的缩放因子一直往下掉要检查模型或数据是否有问题而不是继续调scaler的参数。# 查看当前缩放因子 print(fCurrent scale: {scaler.get_scale()}) # 如果scale持续下降说明训练不稳定注意使用GradScaler时optimizer.step()必须通过scaler.step(optimizer)调用不能直接调optimizer.step()。否则缩放因子不会更新梯度也不会被正确还原。4. 显存优化的组合拳单靠混合精度和激活重计算显存还是可能不够。实际训练大模型时通常需要多种技术组合使用。这里梳理几个最常用的显存优化手段以及它们的适用场景。4.1 ZeRO系列分片存储优化器状态ZeROZero Redundancy Optimizer的核心思路是把优化器状态、梯度、参数分散到多张卡上而不是每张卡都存一份完整的。DeepSpeed实现了ZeRO的三个阶段阶段分片内容显存节省通信开销ZeRO-1优化器状态约4倍低ZeRO-2优化器状态梯度约8倍中ZeRO-3优化器状态梯度参数约N倍N为卡数高以7B模型为例单卡需要84GB不含激活8卡ZeRO-2可以把每卡的优化器状态和梯度分片每卡只需存1/8显存需求降到约20GB左右。ZeRO-3更激进连参数都分片但通信开销也更大。选择哪个阶段取决于你的卡数和互联带宽。如果卡间是NVLink高速互联ZeRO-3的通信开销可以接受如果是PCIe互联ZeRO-2通常更划算。4.2 梯度累积小显存模拟大batch梯度累积的思路很简单用小的batch size做多次前向和反向把梯度累加起来等累积到一定步数后再更新参数。这样等效于用了更大的batch size但显存占用不变。accumulation_steps 4 for i, batch in enumerate(dataloader): with autocast(dtypetorch.bfloat16): outputs model(batch) loss compute_loss(outputs, batch) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这里有个细节loss要除以accumulation_steps否则累积后的梯度会是正确值的N倍。这个除法看起来简单但很多人会忘导致训练效果异常。4.3 模型并行与流水线并行当单卡连模型参数都放不下时比如训练70B以上的模型就需要模型并行。常见的有两种张量并行Tensor Parallelism把单个矩阵乘法拆到多张卡上。比如一个大的线性层权重矩阵按列或按行切分到不同卡上计算时通过all-reduce汇总结果。这种方式通信频繁适合NVLink互联的场景。流水线并行Pipeline Parallelism把模型的不同层放到不同卡上数据像流水线一样依次经过各卡。这种方式通信量小但会有气泡bubble问题——前面的卡在计算时后面的卡在等待。通过微批次micro-batch可以减小气泡。实际训练超大模型时通常是张量并行流水线并行数据并行三种一起用这就是所谓的3D并行。4.4 CPU Offload把暂时不用的挪到内存ZeRO-Offload是DeepSpeed提供的一个功能把优化器状态和梯度放到CPU内存里需要时再搬到GPU。这样做的好处是显存需求大幅降低代价是CPU和GPU之间的数据传输会成为瓶颈训练速度会明显变慢。适合的场景是显存实在不够但CPU内存充足而且对训练速度要求不那么高。比如在单张消费级显卡上微调大模型CPU Offload几乎是唯一的选择。5. 精度选择的实际决策路径理论讲完了回到实际场景。面对一个具体的训练任务怎么决定用FP16还是BF16要不要开损失缩放显存怎么估这里给一条清晰的决策路径。5.1 先查硬件再定精度第一步永远是查硬件支持。在终端跑一行代码import torch print(fBF16 supported: {torch.cuda.is_bf16_supported()}) print(fGPU: {torch.cuda.get_device_name(0)}) print(fVRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB)如果BF16支持直接用BF16省去损失缩放的所有麻烦。如果不支持用FP16 GradScaler。这一步没有太多纠结的空间硬件决定了你的选择范围。5.2 显存估算的实操流程拿到硬件信息后按这个流程估显存算参数、梯度、优化器状态的基础占用参数量 × (2 2 8) 字节混合精度Adam估激活值batch_size × seq_len × hidden_size × num_layers × 系数10~20加总算总需求基础占用 激活值 框架开销约1~2GB对比可用显存如果总需求超过单卡显存考虑激活重计算、ZeRO、梯度累积等手段举个具体例子。假设要训练一个1.3B参数的模型hidden_size2048num_layers24batch_size4seq_len1024用BF16Adam参数梯度优化器状态1.3B × 12 15.6GB激活值4 × 1024 × 2048 × 24 × 15 ≈ 3GB框架开销约1.5GB总计约20GB一张24GB的卡如RTX 4090刚好能跑但余量不多。如果batch_size提到8激活值翻倍到6GB总计约23GB就非常紧张了。这时候开激活重计算可以把激活值降到1GB左右总需求降到18GB就比较舒服了。5.3 训练不稳定时的排查顺序混合精度训练中最常见的问题是loss变成NaN或者不收敛。遇到这种情况按这个顺序排查先看是不是精度问题。把混合精度关掉用纯FP32跑几百步如果loss正常说明是精度问题。这时候如果是FP16检查GradScaler的缩放因子是否正常如果是BF16检查是否有除零或log(0)之类的操作。再看是不是学习率太大。混合精度训练对学习率比FP32更敏感同样的学习率在FP16下可能就会发散。试试把学习率降一半。然后看梯度裁剪。混合精度训练中梯度值可能比FP32大梯度裁剪的阈值需要相应调整。通常设1.0是个安全的起点。最后看数据。如果数据里有异常值比如特别大的数或NaN混合精度下更容易触发溢出。检查一下数据预处理流程。提示BF16虽然动态范围大但不代表不会出问题。BF16的尾数位只有7位精度比FP16低在某些对精度敏感的操作如累加大量小数值中可能引入更大的误差。如果发现BF16训练效果不如FP16可以检查是否有大量的累加操作。6. 几个容易搞混的概念辨析最后澄清几个在实际交流中经常被混淆的概念这些点看似细节但理解错了会导致技术选型走弯路。6.1 FP16和BF16不是精度高低的关系很多人把FP16和BF16简单理解为FP16精度高、BF16精度低这个理解不完整。准确地说FP16在它可表示的范围内精度更高10位尾数 vs 7位尾数BF16能表示的范围大得多8位指数 vs 5位指数在深度学习训练中数值范围比精度更重要因为梯度溢出是比精度损失更致命的问题所以BF16在训练中通常表现更好不是因为精度低反而好而是因为它的动态范围避免了溢出问题。6.2 混合精度不是全部用半精度混合精度的混合二字很关键。它不是把模型全部转成FP16/BF16而是让不同的操作使用不同的精度。权重有一份FP32的master copy前向和反向用半精度计算参数更新时用FP32的master copy。这样既享受了半精度的速度又保持了FP32的更新精度。6.3 int8和bf16的区别这是最近被问得比较多的一个问题。int8和bf16是两种完全不同的东西维度int8bf16数据类型整数浮点数位数816表示范围-128到127约10^-38到10^38主要用途推理量化训练和推理精度损失较大需要校准较小硬件要求需要int8计算支持需要bf16计算支持int8主要用于推理阶段的量化把FP16的权重和激活值转成8位整数显存占用减半推理速度提升。但int8训练目前还不成熟因为整数的梯度传播很困难。bf16则是训练阶段的主流选择。两者不是替代关系而是分别服务于推理和训练两个不同场景。6.4 损失缩放不是万能的损失缩放解决的是FP16梯度下溢的问题但它解决不了上溢。如果某个梯度本身就超过了65504缩放后只会更大直接变成inf。这种情况下需要的是梯度裁剪而不是损失缩放。所以FP16训练中损失缩放和梯度裁剪通常要一起用。# FP16训练的标准配置 scaler.scale(loss).backward() scaler.unscale_(optimizer) # 先还原梯度 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 再裁剪 scaler.step(optimizer) scaler.update()注意unscale_的调用时机必须在scaler.step()之前否则裁剪的是缩放后的梯度阈值就不对了。我在实际训练中发现显存估计最准的方法不是套公式而是先用小模型跑一遍记录实际的显存占用然后按参数量线性外推。公式给的是数量级实际值受框架版本、CUDA版本、具体算子实现的影响可能有20%到30%的偏差。所以估算完之后留出至少30%的显存余量比精确计算更重要。另外BF16虽然省心但在一些老框架版本上支持不完善升级到PyTorch 2.0以上基本就没问题了。