大模型:训练时GPU显存不足怎么办
来源:互联网
时间:2026-08-03 14:12:27
前言

大模型时代,显存压力与日俱增。之前BERT刚火的时候,有一篇《GPU显存不足怎么办?》的文章,这次咱们基于那篇内容做一次重构,专门聊聊大模型场景下显存不够用该怎么破。没看过旧文的朋友,直接看这篇就够了。
训练时显存占用分析
训练模型时,显存都消耗在哪儿了?主要可以拆成四块:
模型权重参数
优化器状态
梯度
激活值
模型权重参数
假设模型大小为A,不同精度下权重占用的显存差异明显:
- 纯fp32训练:
4A
- 混合精度(bf16/fp16):
2A
简单说,精度砍一半,权重显存就缩一半。
优化器状态与梯度
不同的优化器,对显存的消耗天差地别。拿最常用的三种举例:
- :只保留梯度,优化器状态为0,显存占用 = 梯度 4A。
SGD
- :多了一个动量项,优化器状态4A + 梯度4A = 8A。
Momentum-SGD
- :当前梯度、梯度加权平均、梯度平方的加权平均都要存,优化器状态8A + 梯度4A = 12A。
Adam
如果采用混合精度训练,对应的显存占用还会有些变化,但整体规律不变——Adam是最吃显存的“大户”。
激活值
激活值这块,占用与
token长度
per_gpu_batch_size
hidden_size
transformer层数
训练时显存不足怎么办?
下面列出一些常见的省显存操作,按优先级从高到低排列:
- :很多训练脚本在输出层后面会顺便计算rouge等指标,这会生成一个
去掉compute_metrics
的巨型张量,显存瞬间被吃掉一大块。训练时果断关掉它。batch_size × vocab_size × seq_len
- :现在大模型基本都用bf16训练,V100不支持bf16的话可以用fp16。显存占用直接降一半。
采用bf16/fp16混合精度训练
- :不仅省显存,还能加快训练速度,属于“白嫖”型优化,强烈推荐。
Flash attention
- :batch size跟每层激活值显存直接挂钩,调小它立竿见影。如果担心全局batch size变小,可以配合梯度累积。
降低batch size
- :global batch size = batch size × 梯度累积步数。降低batch size后,提高累积步数就能保持等效batch size不变。
梯度累积
- :序列长度与激活值显存正相关,适当缩短上下文长度能释放不少空间。
选择合适的上下文长度
- :显存消耗从高到低依次为 Zero 1 > Zero 2 > Zero 2 + offload > Zero 3 > Zero 3 + offload。建议最多试到
DeepSpeed Zero
,再往下收益递减且复杂度飙升。Zero2 + offload
- :在满足性能需求的前提下,小模型是最直接的解法。
选择更小的基座模型
几个需要慎重考虑的操作:
- :能跑全参数微调就别碰Lora或Qlora。不仅配置麻烦,效果也确有折扣。
Lora
- :比Lora更慢,但显存需求更少。实在没资源时可以试试,但别指望速度。
Qlora
- :支持流水线并行和张量并行,使用门槛较高,适合喜欢折腾的同学。
Megatron-LM
- :Megatron-LM的衍生版,支持Qwen的SFT和PT,坑不少,同样只建议爱折腾的玩家尝试。
Pai-Megatron-LM
- :不推荐,非常耗时。本质是用两次计算换一次存储,时间成本太高。
激活检查点(Activation Checkpointing)
最后
好了,这篇文章主要是对之前的内容做了细化,并补充了大模型时代下的几种显存不足应对方案。大模型时代来了,显卡不够用的“乞丐玩家”是不是越来越多了?希望这些方法能帮你省下一块显存。
参考
【1】https://zhuanlan.zhihu.com/p/31558973
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名