ARTICLE DETAIL

资讯详情

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

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

大模型训练显存优化与混合精度实战指南 显存不够这件事几乎每个做大模型训练的人都撞过。你兴冲冲写好模型代码数据集也清洗完了结果一跑起来就给你甩一个CUDA out of memory连第一个 step 都没迈过去。更气人的是有时候你把 batch size 调到 1 了它还是炸。这时候很多人第一反应是去加卡、换大显存机器但其实相当一部分情况下问题根本不在硬件而在于你没算清楚显存到底花在哪了以及没用对精度格式。这篇就来把大模型训练的显存估计和混合精度训练这两件事讲透。我会从显存的实际构成拆起告诉你每一块显存被谁吃掉了、怎么用公式估出来然后再讲 FP16、BF16、INT8 这几种精度格式到底差在哪、混合精度训练是怎么省显存的、为什么现在大家更推荐 BF16 而不是 FP16。不管你是刚入门想搞明白原理还是已经在训练中遇到 OOM 想找解决方案这篇都能给你可直接上手的方法。1. 显存到底被谁吃掉了四大部分拆解很多人估显存的方式特别粗糙就是拿参数量乘个系数比如 7B 模型就估个 7B × 4 28GB然后发现实际跑起来要 80GB 甚至更多直接懵了。问题在于模型训练时的显存消耗远不止参数本身它至少由四大部分组成而且每一部分的量级都可能超出你的直觉。1.1 模型参数最容易被低估的基础开销模型参数就是权重本身占的显存。一个 7B70 亿参数的模型如果用 FP32单精度浮点每个数占 4 字节存储光权重就要 7B × 4 28GB。如果用 FP16 或 BF16每个数占 2 字节那就是 14GB。注意这里说的是存储一份的量实际训练中往往不止一份后面讲优化器状态时会展开。这里有个常见的认知误区很多人以为用 FP16 训练参数就只占 2 字节。实际上在混合精度训练里为了数值稳定性通常会保留一份 FP32 的 master weight主权重所以参数相关的显存可能是 FP16 副本加 FP32 主副本加起来是 2 4 6 字节每参数。这个细节后面会详细讲。1.2 梯度和参数同量级的影子开销反向传播会为每个可训练参数计算一个梯度梯度的数量和参数量一一对应。所以梯度的显存开销和参数是同一量级的。如果梯度用 FP32 存7B 模型就是 28GB用 FP16/BF16 存就是 14GB。梯度这块容易被忽略的点是如果你用了梯度累积gradient accumulation梯度本身不会成倍增长但如果你做了梯度分片或者用了某些并行策略梯度的存储方式会变。另外梯度在反向传播过程中是逐层产生、逐层释放的所以峰值显存和平均显存会有差异这也是为什么有时候你看到显存曲线是锯齿状的。1.3 优化器状态真正的显存杀手这是最容易被低估的部分。以最常用的 Adam 优化器为例它需要为每个参数维护两个状态一阶矩估计momentum动量和二阶矩估计variance方差。这两个状态通常都用 FP32 存储所以每个参数需要额外 4 4 8 字节。把前面几项加起来用 Adam FP32 训练一个 7B 模型每个参数的显存开销是组成部分精度每参数字节数模型参数FP324梯度FP324Adam 一阶矩FP324Adam 二阶矩FP324合计-167B × 16 字节 112GB。这还没算激活值就已经超过单张 80GB 卡的容量了。这就是为什么全量微调大模型需要多卡或者用 ZeRO 这类优化技术——不是模型本身大是优化器状态太占地方。如果你用的是 SGD 优化器只有动量没有二阶矩那优化器状态就是 4 字节每参数总开销降到 12 字节每参数。但 SGD 在大模型上收敛效果通常不如 Adam所以实际中很少为了省显存而换 SGD。1.4 激活值和 batch size、序列长度强相关激活值是前向传播过程中每一层的输出它们需要被保存下来供反向传播计算梯度用。激活值的显存开销和 batch size、序列长度、隐藏层维度、层数都相关公式比较复杂但核心规律是激活值随 batch size 和序列长度线性增长。粗略估算激活值的经验公式是激活值显存 ≈ batch_size × seq_len × hidden_size × num_layers × 系数这个系数取决于具体的模型结构和实现通常在 10 到 20 之间。举个例子一个 7B 模型hidden_size 约 4096num_layers 约 32batch size 为 1序列长度为 2048 时激活值大约在 2 到 4GB 这个量级。如果把序列长度拉到 8192激活值就会涨到 8 到 16GB。激活值这块有个很实用的优化手段叫激活重计算activation checkpointing / gradient checkpointing思路是不保存所有中间激活值只保存部分关键节点反向传播时重新计算需要的激活值。这样能把激活值显存降到原来的几分之一代价是增加约 30% 的计算时间。在显存紧张时这是性价比很高的取舍。2. 手把手估算从公式到实际数字光知道有哪几部分还不够你得能算出一个具体数字才能判断自己的卡够不够、batch size 能开多大。这一节给出一套可复用的估算流程。2.1 参数、梯度、优化器状态的固定开销这部分是固定的和 batch size 无关只和参数量、精度、优化器类型有关。我整理了一个对照表你可以直接查训练配置每参数字节数7B 模型开销13B 模型开销FP32 SGD444 1284GB156GBFP32 Adam448 16112GB208GB混合精度 Adam22444 16112GB208GB混合精度 Adam ZeRO-1约 16/卡数按卡数分摊按卡数分摊注意混合精度那一行参数有 FP16 副本2 字节和 FP32 主副本4 字节梯度有 FP16 副本2 字节Adam 状态是 FP328 字节加起来是 2428 16 字节。看起来和纯 FP32 一样是的混合精度省的主要是激活值和计算速度对参数相关的固定开销节省有限除非配合 ZeRO 做分片。2.2 激活值的动态估算与 batch size 反推激活值才是决定你能开多大 batch size 的关键。实际估算时我一般用一个简化公式先粗算激活值 ≈ batch_size × seq_len × hidden_size × num_layers × 2 × 精度字节数以 7B 模型为例hidden_size 4096num_layers 32seq_len 2048batch_size 1用 FP162 字节1 × 2048 × 4096 × 32 × 2 × 2 ≈ 2.1GB这个数字和实测比较接近。如果你有 80GB 的卡固定开销混合精度 Adam是 112GB 已经超了所以必须上 ZeRO 或者用 LoRA 这类参数高效微调方法。假设用了 ZeRO-3 把固定开销分摊到 8 张卡上每卡 14GB那还剩 66GB 给激活值理论上 batch size 可以开到 30 左右但实际还要留出通信 buffer、临时变量等空间保守开 16 到 24 比较稳。提示这个公式是粗估实际激活值还受注意力机制实现是否用 FlashAttention、是否有 dropout、是否用梯度检查点等因素影响。建议先用小 batch size 跑起来用torch.cuda.max_memory_allocated()看实际峰值再逐步往上加。2.3 用工具实测别只靠算要验证公式估算只能给你一个起点真正靠谱的是实测。PyTorch 提供了几个很实用的工具import torch # 训练循环中每个 step 后打印峰值显存 print(f峰值显存: {torch.cuda.max_memory_allocated() / 1024**3:.2f} GB) # 重置峰值统计 torch.cuda.reset_peak_memory_stats()另外nvidia-smi可以看到整体显存占用但它包含 CUDA context、缓存等开销通常比max_memory_allocated大 1 到 2GB。判断 OOM 风险时以max_memory_allocated为准更准确。我自己的习惯是先用 batch size 1 跑通记录峰值显存然后按线性关系估算目标 batch size 的显存留 20% 余量再试。如果超了优先考虑开梯度检查点其次降序列长度最后才考虑加卡。3. FP16、BF16、INT8三种精度格式的本质区别搞清楚了显存构成接下来讲精度格式。这是混合精度训练的核心也是很多人搞混的地方。FP16、BF16、INT8 这三种格式名字看起来差不多但设计目标和适用场景完全不同。3.1 浮点数的存储结构符号位、指数位、尾数位要理解这三种格式的区别得先知道浮点数是怎么存的。一个浮点数由三部分组成符号位sign表示正负1 位指数位exponent决定数值的范围能表示多大、多小尾数位mantissa决定数值的精度有效数字有多少位FP32 是 1 位符号 8 位指数 23 位尾数总共 32 位。FP16 是 1 位符号 5 位指数 10 位尾数总共 16 位。BF16 是 1 位符号 8 位指数 7 位尾数总共 16 位。关键差异在这里BF16 的指数位和 FP32 一样是 8 位所以它的数值范围和 FP32 相同但尾数位只有 7 位精度比 FP16 低。FP16 的尾数位有 10 位精度更高但指数位只有 5 位数值范围小得多。3.2 为什么 BF16 比 FP16 更适合大模型训练这个差异直接决定了两者在训练中的表现。FP16 的数值范围大约是 ±65504超出这个范围就会溢出变成 inf。大模型训练中梯度值、激活值很容易出现极端值用 FP16 就经常溢出需要配合 loss scaling损失缩放来把数值拉回安全范围。loss scaling 是个麻烦事需要动态调整缩放因子调不好就训练不稳定。BF16 因为指数位和 FP32 一样数值范围和 FP32 相同基本不会溢出所以不需要 loss scaling训练稳定性好很多。代价是精度低一些但大模型训练对精度的敏感度没有想象中那么高BF16 的 7 位尾数在实践中够用。这就是为什么现在主流的大模型训练包括很多开源模型的训练配置都默认用 BF16 而不是 FP16。BF16 是 Google 在 TPU 上先推的后来 NVIDIA 的 A100、H100 都原生支持 BF16硬件层面也跟上了。格式符号位指数位尾数位数值范围精度是否需要 loss scalingFP321823±3.4e38高否FP161510±65504中是BF16187±3.4e38中低否INT81-7-128~127低不适用3.3 INT8 在训练中的角色量化和推理不是训练主力INT8 是 8 位整数它没有指数位和尾数位的概念就是一个 -128 到 127 的整数。INT8 在训练中主要用于量化感知训练QAT和推理加速不是训练的主力格式。为什么训练不用 INT8因为训练需要梯度更新梯度值通常很小INT8 的精度根本表示不了。INT8 量化一般用在推理阶段把训练好的 FP16/BF16 模型量化成 INT8显存占用降到 1/2推理速度提升明显精度损失通常在可接受范围内。如果你看到有人说用 INT8 训练大模型那多半指的是 QAT——在训练过程中模拟量化误差让模型适应低精度但实际计算还是用浮点。纯 INT8 训练目前还不成熟不是主流做法。4. 混合精度训练怎么混、混哪里、为什么有效混合精度训练不是简单地把所有东西都换成 FP16而是该用低精度的地方用低精度该用高精度的地方用高精度。理解这个混的逻辑比记住配置参数更重要。4.1 混合精度的核心思路计算用低精度更新用高精度混合精度训练的基本流程是这样的前向传播用 FP16/BF16 计算激活值和中间结果都是低精度省显存、算得快反向传播也用 FP16/BF16 计算梯度但保留一份 FP32 的 master weight主权重用 FP16/BF16 的梯度更新 FP32 的主权重再把 FP32 主权重转成 FP16/BF16 供下一轮前向使用为什么要保留 FP32 主权重因为如果直接用 FP16 更新权重每次更新的量可能很小FP16 的精度不够更新会被吃掉舍入误差累积训练会停滞。FP32 主权重保证了更新的精度这是混合精度能稳定训练的关键。4.2 PyTorch 中的混合精度实现autocast 和 GradScalerPyTorch 从 1.6 开始内置了混合精度支持主要用两个组件from torch.cuda.amp import autocast, GradScaler model MyModel().cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-4) scaler GradScaler() # 用于 FP16 的 loss scalingBF16 不需要 for data, target in dataloader: optimizer.zero_grad() # 前向传播用 autocast 自动选择精度 with autocast(dtypetorch.bfloat16): # 或 torch.float16 output model(data) loss loss_fn(output, target) # 反向传播 scaler.scale(loss).backward() # FP16 需要 scaleBF16 可省略 scaler.step(optimizer) scaler.update()用 BF16 的话可以简化成with autocast(dtypetorch.bfloat16): output model(data) loss loss_fn(output, target) loss.backward() optimizer.step()不需要 GradScaler因为 BF16 不会溢出。这就是 BF16 比 FP16 省心的地方。4.3 哪些层必须保持 FP32数值敏感操作的避坑清单autocast 会自动决定哪些操作用低精度、哪些用 FP32但有些操作它可能判断不准或者你需要手动控制。以下这些操作建议保持 FP32Softmax涉及指数运算FP16 容易溢出LayerNorm / BatchNorm统计量计算对精度敏感Loss 计算特别是交叉熵涉及 log 运算小数值的累加比如 attention 的累加FP16 精度不够autocast 默认会把大部分逐元素操作和矩阵乘法转成低精度但会把 reduction 类操作如 sum、mean和归一化层保持 FP32。如果你发现训练不稳定可以检查是不是某些层被错误地转成了低精度。注意用 autocast 时模型参数本身还是 FP32autocast 只是在计算时临时转成低精度。所以显存节省主要来自激活值参数和优化器状态的显存不会因为 autocast 而减少。要减少参数显存需要用model.half()或者配合 DeepSpeed/FSDP 这类框架。5. 实战中的显存优化组合拳知道了原理最后落到实操。显存优化不是单一手段而是一套组合拳。我按性价比从高到低排个序你可以根据实际情况选用。5.1 梯度检查点用时间换空间的第一选择梯度检查点gradient checkpointing是我最推荐的显存优化手段没有之一。它的原理是前向传播时不保存所有中间激活值只保存少数几个检查点反向传播时从最近的检查点重新计算需要的激活值。PyTorch 里开启很简单from torch.utils.checkpoint import checkpoint # 在模型 forward 中对每个 transformer block 用 checkpoint 包裹 def forward(self, x): for layer in self.layers: x checkpoint(layer, x, use_reentrantFalse) return x或者用 HuggingFace Transformers 的话直接设model.gradient_checkpointing_enable()就行。代价是计算时间增加约 30%但激活值显存能降到原来的 1/3 到 1/5。对于显存紧张的场景这个取舍非常划算。我实测过一个 7B 模型开梯度检查点后batch size 能从 4 提到 16训练速度只慢了不到 25%。5.2 ZeRO 与 FSDP多卡场景下的分片策略如果你有多张卡ZeROZero Redundancy Optimizer和 FSDPFully Sharded Data Parallel是必须了解的。它们的核心思路是把参数、梯度、优化器状态分片到不同卡上每张卡只存一部分需要时再通信获取。ZeRO 分三个阶段ZeRO-1只分片优化器状态显存节省约 4 倍ZeRO-2分片优化器状态 梯度节省约 8 倍ZeRO-3分片优化器状态 梯度 参数节省与卡数成正比FSDP 可以理解为 PyTorch 原生的 ZeRO-3 实现。用 DeepSpeed 的话配置 ZeRO-2 通常性价比最高因为 ZeRO-3 的通信开销比较大训练速度会明显下降。5.3 参数高效微调LoRA 为什么能省这么多如果你不是要从头预训练只是想做微调那 LoRALow-Rank Adaptation这类参数高效微调方法是首选。LoRA 的思路是不更新原模型参数而是在旁边加一对低秩矩阵只训练这对小矩阵。一个 7B 模型LoRA 的可训练参数可能只有几百万到几千万优化器状态和梯度的显存开销直接降了两个数量级。原本需要 80GB 卡的全量微调用 LoRA 在 24GB 卡上就能跑。代价是效果可能略逊于全量微调但对大多数下游任务来说差距很小。5.4 精度选择决策树什么时候用 BF16什么时候用 FP16最后给一个精度选择的决策建议硬件支持 BF16A100、H100、RTX 30/40 系等优先用 BF16省心稳定硬件只支持 FP16V100、T4 等用 FP16 GradScaler注意监控 loss scale 变化推理部署可以考虑 INT8 量化显存减半速度提升数值敏感任务如某些科学计算保持 FP32或者混合精度但关键层用 FP32判断硬件是否支持 BF16可以跑这段代码import torch print(torch.cuda.is_bf16_supported())返回 True 就放心用 BF16。显存估计和混合精度这两件事说到底是一个算清楚账的问题。你得知道每一块显存花在哪才能有针对性地优化。我见过太多人一遇到 OOM 就加卡结果加了卡还是 OOM因为瓶颈根本不在卡的数量上而在 batch size 和序列长度的组合上。先把公式算一遍再用实测验证最后按梯度检查点、ZeRO、LoRA 的顺序逐个尝试大部分显存问题都能在不加卡的前提下解决。BF16 能上就上省心不是一点半点。
返回列表