首页 > 教程攻略 > ai资讯 >大模型:训练时GPU显存不足怎么办

大模型:训练时GPU显存不足怎么办

来源:互联网 时间:2026-08-03 14:12:27

前言

大模型:训练时GPU显存不足怎么办

大模型时代,显存压力与日俱增。之前BERT刚火的时候,有一篇《GPU显存不足怎么办?》的文章,这次咱们基于那篇内容做一次重构,专门聊聊大模型场景下显存不够用该怎么破。没看过旧文的朋友,直接看这篇就够了。

训练时显存占用分析

训练模型时,显存都消耗在哪儿了?主要可以拆成四块:

模型权重参数

优化器状态

梯度

激活值

。下面我们以模型本身大小为A(以fp32精度计)作为基准来展开说。

模型权重参数

假设模型大小为A,不同精度下权重占用的显存差异明显:

  • 纯fp32训练:

    4A

  • 混合精度(bf16/fp16):

    2A

简单说,精度砍一半,权重显存就缩一半。

优化器状态与梯度

不同的优化器,对显存的消耗天差地别。拿最常用的三种举例:

  • SGD

    :只保留梯度,优化器状态为0,显存占用 = 梯度 4A。
  • Momentum-SGD

    :多了一个动量项,优化器状态4A + 梯度4A = 8A。
  • Adam

    :当前梯度、梯度加权平均、梯度平方的加权平均都要存,优化器状态8A + 梯度4A = 12A。

如果采用混合精度训练,对应的显存占用还会有些变化,但整体规律不变——Adam是最吃显存的“大户”。

激活值

激活值这块,占用与

token长度

per_gpu_batch_size

hidden_size

transformer层数

正相关,数值也相当可观。技术细节相当复杂,这里就不展开细写了(坦白讲,也没完全算清楚,哈哈)。

训练时显存不足怎么办?

下面列出一些常见的省显存操作,按优先级从高到低排列:

  • 去掉compute_metrics

    :很多训练脚本在输出层后面会顺便计算rouge等指标,这会生成一个

    batch_size × vocab_size × seq_len

    的巨型张量,显存瞬间被吃掉一大块。训练时果断关掉它。
  • 采用bf16/fp16混合精度训练

    :现在大模型基本都用bf16训练,V100不支持bf16的话可以用fp16。显存占用直接降一半。
  • Flash attention

    :不仅省显存,还能加快训练速度,属于“白嫖”型优化,强烈推荐。
  • 降低batch size

    :batch size跟每层激活值显存直接挂钩,调小它立竿见影。如果担心全局batch size变小,可以配合梯度累积。
  • 梯度累积

    :global batch size = batch size × 梯度累积步数。降低batch size后,提高累积步数就能保持等效batch size不变。
  • 选择合适的上下文长度

    :序列长度与激活值显存正相关,适当缩短上下文长度能释放不少空间。
  • DeepSpeed Zero

    :显存消耗从高到低依次为 Zero 1 > Zero 2 > Zero 2 + offload > Zero 3 > Zero 3 + offload。建议最多试到

    Zero2 + offload

    ,再往下收益递减且复杂度飙升。
  • 选择更小的基座模型

    :在满足性能需求的前提下,小模型是最直接的解法。

几个需要慎重考虑的操作:

  • Lora

    :能跑全参数微调就别碰Lora或Qlora。不仅配置麻烦,效果也确有折扣。
  • Qlora

    :比Lora更慢,但显存需求更少。实在没资源时可以试试,但别指望速度。
  • Megatron-LM

    :支持流水线并行和张量并行,使用门槛较高,适合喜欢折腾的同学。
  • Pai-Megatron-LM

    :Megatron-LM的衍生版,支持Qwen的SFT和PT,坑不少,同样只建议爱折腾的玩家尝试。
  • 激活检查点(Activation Checkpointing)

    :不推荐,非常耗时。本质是用两次计算换一次存储,时间成本太高。

最后

好了,这篇文章主要是对之前的内容做了细化,并补充了大模型时代下的几种显存不足应对方案。大模型时代来了,显卡不够用的“乞丐玩家”是不是越来越多了?希望这些方法能帮你省下一块显存。

参考

【1】https://zhuanlan.zhihu.com/p/31558973