资讯动态

训练侧显存测量与优化:从账单拆解到预算决策实战

发布时间:2026/10/4 11:13:38 来源:尧图企业网站定制
1. 训练侧显存测量到底在测什么显存优化这件事很多人一上来就想着怎么省结果省了半天发现根本没省到点子上。问题出在哪出在没搞清楚显存到底被谁吃掉了。训练侧的显存测量核心目标就一个把显存账单拆开看清楚每一笔开销的去向然后才能做预算决策。我见过太多人拿着nvidia-smi看一眼显存占用发现快满了就开始慌然后盲目上梯度检查点、盲目降batch size最后训练速度掉了一半显存也没省下多少。这种做法的问题在于nvidia-smi看到的只是一个总数它不会告诉你这8GB里有3GB是模型参数、2GB是优化器状态、1.5GB是激活值、剩下的是碎片和临时缓冲区。你不知道钱花在哪就没办法做预算。训练侧显存测量要回答的问题很具体模型参数占多少、梯度占多少、优化器状态占多少、激活值占多少、临时缓冲区占多少。这五块加起来才是真实账单。而且这五块的比例关系会随着模型规模、batch size、序列长度、精度策略的变化而剧烈变化。一个7B模型在FP16下参数占14GB但如果用Adam优化器优化器状态就要占28GBFP32的一阶动量和二阶动量各14GB再加上梯度14GB光这三项就56GB了还没算激活值。所以为什么大家说7B模型全量微调至少要80GB显存账就是这么算出来的。测量方法上最直接的是用PyTorch的显存分析工具。torch.cuda.memory_allocated()能拿到当前分配的显存torch.cuda.max_memory_allocated()能拿到峰值。但这两个数字只反映PyTorch分配器层面的情况不包括CUDA上下文本身占用的那几百MB。更细的拆解需要用torch.cuda.memory_summary()它会按分配块大小分类列出。不过这个输出比较原始我一般会自己写个hook在模型forward前后、backward前后分别打点算出每个阶段的增量。还有一个容易被忽略的点显存碎片。PyTorch的缓存分配器会预留一些显存不还给系统导致nvidia-smi看到的占用比memory_allocated()高不少。这个差值在长时间训练中会逐渐增大尤其是当你有动态shape的输入时。测量的时候要把这个差值也记录下来否则你按memory_allocated()做的预算到了实际训练时就会OOM。注意测量一定要在真实训练循环里做不能只跑一个forward就完事。因为激活值的峰值出现在backward阶段只测forward会严重低估。2. 显存账单的五大部分与计算逻辑2.1 模型参数与梯度的显存占用模型参数的显存占用最好算参数量乘以每个参数的字节数。FP32是4字节FP16和BF16是2字节INT8是1字节。一个7B模型在FP16下就是7B乘2等于14GB。梯度占用的字节数和参数一致因为梯度需要和参数同样的精度来保证更新时不丢信息。所以FP16训练时参数加梯度就是28GB。但这里有个坑很多人以为用FP16训练参数就是FP16。实际上PyTorch的AMP自动混合精度会保留一份FP32的master weight。也就是说参数实际上占了两份一份FP16用于前向和反向计算一份FP32用于优化器更新。这样算下来7B模型的参数占用是14GBFP16加28GBFP32 master总共42GB。梯度也是FP16一份但优化器更新时需要FP32的梯度所以梯度实际占用也是14GBFP16加28GBFP32又是42GB。这就是为什么AMP能省显存但省不了太多——它省的是激活值和部分计算缓冲区参数和优化器状态的大头省不掉。2.2 优化器状态的显存黑洞优化器状态是显存占用里最容易被低估的部分。以Adam为例它需要为每个参数维护一阶动量m和二阶动量v都是FP32。所以优化器状态的显存等于参数量乘以4字节乘以2也就是参数量乘以8字节。7B模型就是56GB。加上前面的参数和梯度已经98GB了。这就是为什么全量微调7B模型至少需要8张80GB的卡来做数据并行——单卡根本放不下。AdamW稍微好一点它把权重衰减和梯度更新解耦了但动量部分还是一样的。Adafactor通过分解二阶动量矩阵来省显存能把优化器状态降到参数量乘以4字节左右但收敛性会受一些影响。Sophia优化器用对角Hessian估计来替代二阶动量也能省不少但实现复杂度和调参难度都上去了。实际做预算的时候我一般按这个公式估算总显存等于参数量乘以2加2加8加激活值加缓冲区。前面的2是FP16参数第二个2是FP16梯度8是Adam的优化器状态。这个公式在AMP加AdamW的场景下比较准误差在10%以内。2.3 激活值的动态波动激活值是显存占用里最动态的部分它跟batch size、序列长度、模型层数、隐藏维度都相关。粗略估算的话激活值显存约等于batch size乘以序列长度乘以隐藏维度乘以层数乘以一个系数。这个系数取决于具体的网络结构Transformer里主要是注意力矩阵和FFN的中间激活。注意力矩阵的显存是序列长度的平方乘以batch size乘以头数乘以2字节。序列长度2048时这个平方项就是4M乘以batch size和头数后很容易上GB。序列长度拉到8192平方项变成64M直接爆炸。这就是长上下文训练显存吃紧的核心原因。FFN的中间激活是batch size乘以序列长度乘以4倍隐藏维度乘以2字节。这个和序列长度是线性关系比注意力矩阵温和一些。但层数一多累积起来也很可观。梯度检查点gradient checkpointing就是针对激活值的优化手段。它不保存中间激活而是在backward时重新计算。代价是计算量增加约30%但激活值显存能降到原来的平方根级别。对于层数很深的模型这个 trade-off 非常划算。2.4 临时缓冲区与碎片临时缓冲区包括CUDA kernel执行时的workspace、通信操作的缓冲区、以及PyTorch分配器的预留空间。这部分很难精确测量但可以通过对比nvidia-smi和memory_allocated()的差值来估算。一般来说这个差值在500MB到2GB之间模型越大、并行策略越复杂差值越大。碎片问题在动态shape场景下特别严重。比如你训练时序列长度不固定PyTorch的缓存分配器会按最大shape预留块导致实际占用远高于理论值。解决办法是设置PYTORCH_CUDA_ALLOC_CONF环境变量启用expandable_segments让分配器能合并碎片。这个设置在新版PyTorch里效果很明显我实测能把碎片率从20%降到5%以下。3. 预算决策从测量结果到优化方案3.1 显存预算的分配原则测完显存账单后下一步是做预算决策。预算的核心原则是先保训练稳定性再保吞吐量最后才考虑省显存。很多人搞反了顺序为了省显存把batch size降到1结果训练不稳定收敛慢得要命省下来的显存也没换来什么好处。我的预算分配一般是这样的参数和优化器状态是刚性支出没法省必须留足。激活值是弹性支出可以通过梯度检查点、序列并行、FlashAttention等手段压缩。临时缓冲区留10%到15%的余量防止峰值OOM。具体到数字上假设你有80GB显存7B模型AMP加AdamW参数加梯度加优化器状态大约98GB单卡放不下。这时候你有几个选择一是上ZeRO Stage 2把优化器状态和梯度分片到多卡每卡只存一部分二是上LoRA只训练低秩适配器参数量降到原来的百分之一三是上QLoRA把基座模型量化到4bit进一步压缩。3.2 ZeRO与FSDP的显存账ZeRO Stage 1只分片优化器状态每卡显存等于参数量乘以2加梯度乘以2加优化器状态除以N。7B模型8卡每卡优化器状态7GB加上参数14GB和梯度14GB总共35GB80GB卡能放下。Stage 2再分片梯度每卡梯度降到1.75GB总共约23GB。Stage 3分片参数每卡参数1.75GB总共约10GB但通信开销大幅增加。FSDP本质上是ZeRO Stage 3的PyTorch原生实现它把参数、梯度、优化器状态都分片每卡只存1/N。但FSDP在forward和backward时需要all-gather参数通信量很大。实际用下来FSDP在8卡A100上训练7B模型每卡显存约12GB但吞吐量比DDP低20%左右。这个 trade-off 要看你的瓶颈是显存还是算力。3.3 AMP的显存收益与代价AMP自动混合精度是显存优化的第一板斧但它的收益经常被高估。AMP省的主要是激活值和部分计算缓冲区参数和优化器状态的大头省不掉。实测下来AMP能把激活值显存降到FP32的60%左右总体显存节省约20%到30%。AMP的代价是数值稳定性。FP16的动态范围窄梯度容易下溢或上溢。PyTorch的GradScaler通过动态调整loss scale来缓解这个问题但在某些模型结构上仍然会出现NaN。BF16的动态范围和FP32一样不需要loss scaling但精度比FP16低。A100及以上支持BF16V100只支持FP16。选哪个取决于你的硬件和模型对精度的敏感度。实操心得用AMP时把LayerNorm和softmax强制保留在FP32能显著提升训练稳定性。PyTorch的torch.cuda.amp.autocast默认已经这么做了但如果你自己写kernel要注意手动指定。4. 实操测量流程与工具链4.1 测量脚本的编写要点我一般会写一个独立的测量脚本不掺在训练代码里这样干净、可复现。脚本的核心结构是初始化模型和优化器构造一个真实batch的输入然后分阶段打点。import torch from torch.cuda import memory_allocated, max_memory_allocated, reset_peak_memory_stats def measure(model, optimizer, input_ids, labels): reset_peak_memory_stats() base memory_allocated() # forward outputs model(input_ids, labelslabels) loss outputs.loss after_forward memory_allocated() # backward loss.backward() after_backward memory_allocated() # optimizer step optimizer.step() optimizer.zero_grad() after_step memory_allocated() return { base: base, forward_delta: after_forward - base, backward_delta: after_backward - after_forward, step_delta: after_step - after_backward, peak: max_memory_allocated() }这个脚本跑一次就能拿到四个关键数字。forward_delta主要是激活值backward_delta是梯度加激活值的峰值step_delta是优化器状态的增量。peak是整个过程的最大值用来做OOM判断。4.2 不同配置的对比测量测量不能只测一个配置要测一组配置做对比。我一般会测这几组FP32基线、AMP、AMP加梯度检查点、AMP加梯度检查点加ZeRO。每组跑三次取平均排除冷启动的影响。对比的时候重点看两个指标峰值显存和吞吐量。峰值显存决定你能不能跑起来吞吐量决定你跑得多快。有时候省显存的方案会把吞吐量砍半这时候就要算一笔账省下来的显存能不能换来更大的batch size如果能吞吐量可能反而更高。举个例子7B模型FP32训练峰值显存60GB吞吐量100 samples/s。AMP后峰值45GB吞吐量130 samples/s。AMP加梯度检查点后峰值30GB吞吐量90 samples/s。虽然梯度检查点让吞吐量降了但省下的15GB显存可以让你把batch size翻倍实际吞吐量变成180 samples/s。这就是预算决策的价值。4.3 显存碎片的手动清理测量过程中如果发现nvidia-smi和memory_allocated()差值越来越大说明碎片在累积。这时候可以手动调torch.cuda.empty_cache()但注意这个操作会释放缓存分配器预留的块可能导致后续分配变慢。更好的办法是设置PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True让分配器自己管理碎片。还有一个技巧是在训练循环里定期调用torch.cuda.reset_peak_memory_stats()把峰值统计重置这样能更准确地看到每个step的显存波动。如果不重置峰值会一直保留历史最大值掩盖了后期的显存增长。5. 常见问题与排查技巧实录5.1 测量结果与nvidia-smi对不上这是最常见的问题。memory_allocated()只统计PyTorch分配器管理的显存不包括CUDA上下文、cuDNN workspace、NCCL通信缓冲区。这些加起来可能有1到2GB。所以nvidia-smi看到的数字总是比memory_allocated()大。解决办法是用torch.cuda.memory_snapshot()拿到完整的内存快照它会列出所有分配块包括那些不在PyTorch管理范围内的。不过这个输出很大一般只在排查问题时用。5.2 激活值峰值出现在意想不到的地方有时候你会发现backward阶段的显存峰值比forward高很多甚至高出一倍。这通常是因为某些操作的backward需要保存forward的中间结果而这些结果在forward结束后并没有释放。比如attention的softmax输出forward时算完就丢了但backward需要它来计算梯度所以PyTorch会把它保留到backward。排查方法是给每个模块注册forward hook和backward hook记录每个模块的输入输出显存。这样能精确定位到哪个模块的backward显存开销异常。5.3 OOM发生在optimizer.step()optimizer.step()本身不分配大块显存但它会触发参数更新而参数更新需要读取梯度、写入参数。如果梯度是FP16而参数是FP32这里会有一个隐式的类型转换产生临时缓冲区。Adam的动量更新还会产生中间变量。这些加起来可能几百MB在显存已经接近上限时就是压死骆驼的最后一根稻草。解决办法是在step之前手动torch.cuda.empty_cache()或者把优化器状态分片到多卡。另一个办法是用foreach实现的优化器它把多个参数的更新合并成一个kernel减少临时缓冲区的分配次数。5.4 梯度检查点与AMP的兼容问题梯度检查点在backward时会重新计算forward如果和AMP一起用重计算时的精度策略要和原forward一致否则梯度会对不上。PyTorch的torch.utils.checkpoint默认会保留AMP的autocast状态但如果你自己实现了checkpoint逻辑要注意手动传递autocast上下文。还有一个坑是梯度检查点不能和某些自定义autograd函数一起用因为重计算时这些函数的forward可能不是确定性的。遇到这种情况要么把自定义函数排除在checkpoint范围外要么确保它是确定性的。5.5 多卡训练时的显存不均衡数据并行时每卡的显存占用应该基本一致。如果发现某张卡显存特别高通常是数据加载不均衡或者通信操作卡住了。检查DataLoader的num_workers和pin_memory设置确保每个进程拿到的batch大小一致。通信方面NCCL的all-reduce是同步操作如果某张卡算得慢其他卡会在通信点等待显存占用会暂时升高。排查方法是打印每张卡的memory_allocated()看差异是否超过5%。如果超过检查数据分片逻辑和通信配置。问题现象可能原因排查方法解决手段nvidia-smi比memory_allocated高2GBCUDA上下文和通信缓冲区memory_snapshot预留2GB余量backward峰值远高于forward中间激活未释放模块级hook梯度检查点step时OOM类型转换临时缓冲区逐步打点foreach优化器多卡显存不均衡数据或通信不均衡逐卡打印调整DataLoader碎片率持续增长动态shape监控差值expandable_segments6. 预算决策的实战案例6.1 单卡24GB训练7B模型的可行性分析有人问单卡24GB能不能训7B模型。按前面的公式算7B模型AMP加AdamW参数14GBFP16加梯度14GBFP16加优化器状态56GBFP32总共84GB远超24GB。所以全量微调不可能。但LoRA可以。LoRA只训练低秩矩阵参数量通常是原模型的0.1%到1%。7B模型的LoRA参数量约7M到70MFP16下占14MB到140MB。优化器状态按Adam算是参数量乘以8也就是112MB到1.1GB。加上基座模型的14GBFP16总共约15GB到16GB。24GB卡能放下还能留出8GB给激活值和缓冲区。QLoRA更进一步把基座模型量化到4bit7B模型只占3.5GB。加上LoRA参数和优化器状态总共约5GB。24GB卡能跑得很宽裕甚至能上更大的batch size。6.2 多卡场景下的并行策略选择如果你有4张24GB卡总共96GB显存想训7B模型全量微调。DDP每卡都要存完整的参数、梯度、优化器状态每卡84GB放不下。ZeRO Stage 2把优化器状态和梯度分片到4卡每卡参数14GB加梯度3.5GB加优化器状态14GB总共31.5GB还是放不下。ZeRO Stage 3把参数也分片每卡参数3.5GB加梯度3.5GB加优化器状态14GB总共21GB勉强能放下但激活值还没算。这时候要么上梯度检查点把激活值压到2GB以内要么上CPU offload把优化器状态放到内存。CPU offload的代价是step速度慢3到5倍但能省下14GB显存。实测下来4卡24GB加ZeRO Stage 3加梯度检查点加CPU offload能训7B模型但吞吐量只有DDP的十分之一。所以如果追求速度还是建议上80GB卡。6.3 预算决策的决策树我把预算决策整理成一个简单的决策树方便快速判断第一步算刚性支出参数量乘以2加2加8除以并行度。如果这个数字超过单卡显存的80%考虑LoRA或QLoRA。第二步算激活值batch size乘以序列长度乘以隐藏维度乘以层数乘以系数。如果超过剩余显存的50%上梯度检查点。第三步算碎片余量留10%到15%的显存给临时缓冲区和碎片。如果不够调expandable_segments或减小batch size。第四步测吞吐量在满足显存约束的前提下找吞吐量最大的配置。有时候大batch加梯度检查点比小batch不加检查点更快。这个决策树不是绝对的但能帮你快速缩小选择范围。实际做的时候还是要跑测量脚本验证因为不同模型结构、不同框架版本的显存行为差异很大。最后分享一个小技巧在训练脚本里加一个显存监控回调每个step记录峰值显存和吞吐量输出到TensorBoard。这样你能看到显存随训练进程的变化趋势及时发现碎片累积或激活值增长的问题。我靠这个回调抓到过好几次隐蔽的显存泄漏都是自定义层里缓存了不该缓存的东西。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑