☰
SGD与Adam显存占用差异解析:优化器状态如何影响深度学习训练
2026/10/2 2:44:21 网站建设 项目流程

在训练深度学习模型的过程中,显存不够用的场景基本人人都经历过。而显存消耗的构成里,除了模型参数、梯度和激活值,优化器的状态往往是被低估的大头。我接触过的不少新手同学,在调参阶段先把SGD换成Adam试了试,结果发现同一个模型、同一个batch size,显存突然就爆了,一时找不到原因。其实这就是优化器在内存占用上的差异在作怪。搞清楚SGD和Adam在内存占用上的区别,不仅有助于理解优化器的工作原理,也能帮你在设计训练方案时少走弯路,特别是当你跑大模型、大batch、长序列时,这个决策会直接影响能不能训练成功。

这篇文章会从优化器的状态存储机制入手,把SGD和Adam每一步到底在显存里放了什么东西讲清楚,配合计算示例和PyTorch实测方法,最后再给出一些显存紧张时的实操建议。适合正在做深度学习训练、对显存优化有困惑的同学,也适合需要为多卡训练做资源配置的工程师参考。

1. 内容整体设计与思路拆解:先弄清楚内存占用从哪来

1.1 优化器状态不是“存储”,而是训练的必要成本

想在内存占用这个问题上有个清晰的判断,先得明白优化器在训练过程中到底扮演什么角色。通俗点说,模型参数负责“表达”,梯度负责“指出改进方向”,而优化器状态负责“记住历史信息”。SGD和Adam最大的差异,恰恰就体现在“记住历史信息”这件事上。

梯度下降的核心动作是:计算出梯度后,沿着梯度的反方向更新参数。普通SGD只依赖当前时刻的梯度,它不需要记录历史梯度。但如果加上了动量(momentum),SGD就需要额外保存一个“速度向量”,用来累积历史的梯度方向。Adam则更进一步,它不仅保存一阶动量(也就是带指数衰减的历史梯度均值),还需要保存二阶动量(历史梯度平方的指数衰减均值),用于自适应地调整每个参数的学习率。

内存占用从直观上就可以预估:SGD几乎是零状态开销,而Adam每多一个参数就要多保存两份浮点状态。我见过很多工程同学在上手阶段不理解为什么Adam类优化器“那么贵”,本质上是没有意识到这个二阶动量需要为每个参数单独开辟存储空间。在参数量达到亿级甚至千亿级时,这部分开销就非常可观了。

具体到实现层面,以PyTorch为例,初始化一个Adam优化器时,它会为每个参数创建两个和参数形状完全相同的exp_avg和exp_avg_sq张量。如果你的模型是1亿参数,每个参数以fp32保存(4字节),那么Adam光是这两个状态张量就要占用800MB左右。加上模型参数本身400MB、梯度400MB,粗略一算就接近1.6GB了。而普通SGD(不带动量)只需要参数和梯度,合计800MB。这就是最核心、也最直观的差异来源。

1.2 内存占用为什么值得单独拎出来分析

有人可能会问:如果GPU显存够大,不在乎多占1GB,那这个区别重要吗?答案是:在单卡小模型上确实不重要,但在真实的业务场景里,显存往往是最稀缺的资源。

举一个实际场景:你在训练一个很大的Transformer模型,模型参数本身已经占据了大部分显存,剩下的空间要留给激活值(activation)和梯度。如果你的batch size受到显存限制,被迫降到很小,训练效率和稳定性都会受到影响。这种情况下,优化器状态每多占一分,训练配置就紧张一分。

另外,分布式训练场景中还有一个容易被忽略的点:优化器状态在数据并行(DDP)下是每张卡都保存一份完整副本的。也就是说,4卡训练时Adam状态占用是4份,8卡就是8份。虽然梯度通过all-reduce同步,但优化器状态从未被“均分”,除非你用ZeRO之类的方案进行分片。这也是为什么在训练大模型时,业界普遍会觉得Adam类优化器的显存压力大——它在多卡环境下会成倍放大。

把这层逻辑想清楚以后,再来看SGD和Adam的内存差异,就不会只停留在“Adam多占一倍显存”这种粗浅印象上了,而是能进一步思考:为什么是这样设计?有没有办法省?什么时候可以省?这些都会在后面的各个章节展开。

2. 核心细节解析与实操要点:Adam多占的那部分到底贵在哪

2.1 二阶动量v的不可替代性与内存成因

Adam的更新规则本质上是在对学习率做逐参数的调整。它维护的一阶动量m相当于梯度的平滑版本,用来决定更新的方向;二阶动量v则是梯度平方的平滑版本,用来衡量每个参数在过去一段时间内的梯度波动幅度。当某个参数的历史梯度一直很大时,v很大,更新步长就被压缩;当某个参数梯度很小且稳定时,v较小,更新步长就相对放大。这种自适应机制让Adam在稀疏梯度和非平稳目标上表现优异。

但是,v这个状态张量是有代价的。为了让每个参数都有独立的步长,优化器必须为每个参数保存独立的v值,这跟参数本身一样大。想象一下,如果是350亿参数的大模型,仅exp_avg_sq这一个张量,fp32下就是140GB。这也是为什么大批量训练大模型时,纯Adam几乎不可能直接在单卡上跑起来。

相比之下,带动量的SGD同样也有一份额外的历史信息m,但它不需要维护二阶的v。因此它的内存公式是“参数+梯度+动量buffer”,合起来是3倍参数内存,Adam则是4倍参数内存。如果再用fp16混合精度训练,状态张量仍然以fp32保存,这个差距在数值上还会进一步拉开。后面的实操章节会给出具体计算和示例。

2.2 优化器step阶段的峰值内存与临时张量

很多人计算内存时只看“参数+梯度+优化器状态”的稳态大小,却忽略了优化器参数更新那一步的峰值内存。PyTorch的优化器在调用optimizer.step()时,为了完成逐参数更新,会先读取参数、梯度、m、v这几个张量,然后计算新的参数值,再写回。在这个过程中,GPU会分配若干临时张量。

对于Adam来说,更新的计算表达式通常涉及多个中间步骤,例如grad的平方、m的更新、v的更新、分母的sqrt(v) + eps等。如果框架在实现时没有做很好的内存复用,峰值内存可能会比稳态多出几个临时张量的大小。虽然现代深度学习框架大多已经做了大量的in-place优化,但如果你在写自定义优化器或者使用一些比较“原始”的实现,这个峰值开销可能会非常明显。

我建议在实际评估显存时,不要只看nvidia-smi的现显存占用,而是使用PyTorch的torch.cuda.max_memory_allocated()来记录运行过程中的峰值。实测下来,用Adam训练时,峰值和稳态之间的差额通常会比用SGD训练时更大,因为Adam涉及的计算图更复杂、临时张量更多。这一点在极端的显存紧张情况下甚至会影响训练是否OOM。

2.3 变体优化器的内存影响:带动量SGD、AdamW与8bit优化器

讨论SGD和Adam的内存差异,还不能忽视它们各自的变体。SGD加上momentum后,内存从2倍参数上升到3倍;AdamW和Adam在状态存储上基本一致,但AdamW把weight decay从梯度中分离出来,不引入额外状态。LAMB则在Adam的基础上引入了层间自适应学习率调整,也没有增加新的状态张量。因此在内存维度上,AdamW和LAMB都可以视作与Adam“同等昂贵”。

还有一类值得关注的是8bit优化器。这类方案的核心思路是把优化器的状态张量压缩为8bit整数存储,用更小的数值范围换取一半甚至更多的内存节省。不过8bit状态在更新时通常需要反量化为fp32进行计算,这会引入额外的转换开销,也可能会对精度和稳定性带来影响。实测下来,对于大规模模型,8bit优化器效果不错,但新手使用时要格外注意数值溢出问题。这个话题在最后一章会专门展开。

3. 实操过程与核心环节实现:动手估算并实测验证

3.1 内存公式与计算示例

我先给出一个可以直接套用的内存计算公式。设模型参数量为 P,精度为fp32,那么:

  • 模型参数占用:4P 字节
  • 梯度占用:4P 字节
  • 普通SGD:参数 + 梯度 = 8P 字节
  • 动量SGD:参数 + 梯度 + 动量buffer = 12P 字节
  • Adam:参数 + 梯度 + 一阶动量 m + 二阶动量 v = 16P 字节

如果是混合精度训练,模型参数以fp16保存(2P字节),梯度通常也是fp16(2P字节),但Adam的两个状态仍然需要fp32精度(各4P字节)。如果把Adam的权重更新放到fp32的master weight上进行,还得额外多一份fp32参数副本(4P字节),总计就是 2P + 2P + 4P + 4P + 4P = 16P 字节。也就是说,在混合精度下,Adam的总状态依然大约是参数量的8倍(相对fp16参数而言)。这个数字对显存规划很有参考价值。

光看公式不够直观,我们来取一个例子:假设模型参数是2亿(200M),精度为fp32。

项目SGD(无动量)动量SGDAdam
参数800MB800MB800MB
梯度800MB800MB800MB
一阶状态0800MB800MB
二阶状态00800MB
合计1600MB2400MB3200MB

这个表格看起来很简单,但它说明了一个很关键的事实:在同等参数量和精度下,Adam比无动量SGD多占用1600MB,比动量SGD多占用800MB。如果你的模型来到10亿参数,Adam光优化器状态就要12.8GB,这已经接近很多消费级显卡的显存上限了。正因为如此,大模型训练中对Adam的优化器状态“动手脚”才成为一门重要的工程学问。

3.2 PyTorch实测:对比SGD和Adam的显存占用

公式终究是纸面计算,我建议你在自己的机器上跑一次实测,建立直观感知。下面这段代码可以在PyTorch中分别创建同一个小模型,使用SGD和Adam优化器,记录训练过程中的峰值显存。

import torch import torch.nn as nn class SimpleNet(nn.Module): def __init__(self, hidden=1024): super().__init__() self.fc1 = nn.Linear(1024, hidden) self.fc2 = nn.Linear(hidden, hidden) self.fc3 = nn.Linear(hidden, 10) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) return self.fc3(x) def test_optimizer_memory(optimizer_cls, steps=20, batch_size=64): torch.cuda.reset_peak_memory_stats() model = SimpleNet().cuda() opt = optimizer_cls(model.parameters(), lr=1e-3) x = torch.randn(batch_size, 1024).cuda() y = torch.randint(0, 10, (batch_size,)).cuda() loss_fn = nn.CrossEntropyLoss() for _ in range(steps): opt.zero_grad() out = model(x) loss = loss_fn(out, y) loss.backward() opt.step() return torch.cuda.max_memory_allocated() / 1024**2 sgd_mem = test_optimizer_memory(lambda params, lr: torch.optim.SGD(params, lr=lr)) adam_mem = test_optimizer_memory(lambda params, lr: torch.optim.Adam(params, lr=lr)) print(f"SGD peak memory: {sgd_mem:.1f} MB") print(f"Adam peak memory: {adam_mem:.1f} MB")

这段代码用了很小的网络和固定随机输入,目的不是训练出什么指标,而是观察峰值内存的差异。你拿到的结果里,Adam的峰值内存通常会比SGD高出一大截,高出来的部分基本对应Adam维护的一阶和二阶状态,以及在step时的临时张量开销。如果你把hidden调大,比如改成4096,差值会更加明显。

还需要提醒的是,峰值内存和稳态内存不是同一个概念。上面统计的max_memory_allocated()是运行过程中的最高值,包括激活值、临时张量等。在真实训练中,一个大的batch所产生的激活值可能会远超优化器状态,所以你在判断内存瓶颈时,要分别统计各部分,不要直接认为优化器状态就是全部占用。

3.3 使用torch profiler定位内存分配细节

如果你想进一步确认哪些张量占了最多显存,可以使用PyTorch的torch.profiler来查看内存分配事件。例如在训练循环外层包一层torch.profiler.profile(profile_memory=True),然后通过prof.key_averages().table(sort_by="self_memory_footprint", row_limit=20)查看占用最多的操作。

我实际用下来,Adam的step操作通常会在内存分配表中占很大一块,因为它在更新时同时读取和写入了多个与参数等形状的张量。SGD的step则相对简单,动态分配的内存少很多。这个profiler工具对分析显存瓶颈非常有用,比单纯看nvidia-smi要精细得多,建议遇到OOM问题时先跑一下。

3.4 混合精度AMP下的内存账本

现在工程上普遍使用自动混合精度(AMP)来加速训练和减少显存。AMP带来的一个有趣现象是:模型参数和梯度是fp16,但优化器状态通常仍然保持fp32精度,因此Adam状态的相对占比反而更高了。

用前面那个200M参数模型举例。纯fp32训练时,Adam总占用是3200MB,其中参数和梯度占1600MB,优化器状态占1600MB。换成AMP之后,参数和梯度各占400MB,master weight占800MB,Adam的两个状态各占800MB,总计约3200MB。可以看到,整体并没有比fp32省太多,原因正是master weight和两个动量状态都必须保留fp32。SGD配合AMP的情况也类似,参数和梯度变成fp16,但master weight和动量buffer仍然是fp32。这也是为什么很多官方大模型训练教学里都会强调:AMP省下的显存,一部分被优化器状态又吃回去了。

理解了这一层,你就明白了为什么ZeRO、分片优化器、8bit优化器要在优化器状态上做文章。因为训练越是往大模型走,优化器状态就越显眼,甚至超过参数本身好几倍。

4. 常见问题与排查技巧实录:显存告急时该怎么应对

4.1 显存爆掉,不要第一时间换优化器

训练中途OOM,很多人的第一反应是“Adam太吃内存了,换成SGD试试”。这个思路可以理解,但未必是代价最小的方案。显存占用里除了优化器,还有激活值、临时张量、分布式通信缓冲等。单纯换优化器可能只是把峰值从爆掉降到刚好能跑,但训练效率可能下降,得不偿失。

我建议按下面的顺序排查:

  • 先看激活值能不能通过activation checkpointing(激活检查点)压缩,这通常能省下非常可观的显存,代价是增加少量计算。
  • 再确认batch size是否需要那么大,如果显存接近满载,梯度累积可以在不减少每个batch样本数的情况下,降低每一步的实际显存需求(但要注意梯度累积和batch size对BatchNorm等层的影响不同)。
  • 然后检查模型实现中是否有不必要的中间变量留在计算图中,例如把不需要的hidden state及时释放,或者用del配合torch.cuda.empty_cache()做临时清理。
  • 最后才考虑优化器状态的压缩,包括切换到动量SGD、使用8bit优化器或者做ZeRO分片。

这个顺序的核心理念是:先处理那些对收敛影响小的部分,再考虑对训练行为影响大的部分。

4.2 换用SGD会牺牲什么

把Adam换成无动量SGD,内存确实立刻减少一半,但代价不是免费的。SGD对学习率极其敏感,如果学习率设置不合理,收敛会非常慢甚至发散;Adam则凭借自适应学习率,对初学者友好得多。很多视觉模型和推荐模型用SGD配合精细调参后可以达到不错的泛化效果,但调参成本显著上升。

如果显存不是极度紧张,我更推荐的折中方案是使用带动量的SGD,因为它保留了历史梯度信息,训练稳定性比纯SGD好,而内存只比Adam少一份二阶状态。当然,这仍然不是单独决定模型精度的关键,你需要结合数据集规模和训练技巧综合判断。

4.3 8bit优化器的实测体验与坑

现在不少同学会尝试用bitsandbytes库把Adam优化器变成8bit版本,号称能省一半内存。我实测下来确实有效,比如一个原本需要4GB优化器状态的模型,8bit化之后能压到2GB出头。但要注意几点:第一,8bit状态在数值范围上远小于fp32,如果某些参数的梯度方差极大,容易溢出,导致训练不稳定;第二,这类优化器通常需要在更新时临时反量化到fp32,会引入额外计算,训练速度可能略有下降;第三,如果模型本身参数比较少,8bit优化器省下的绝对内存可能不值得引入这个额外依赖。

我建议在使用之前跑一个短时间的对比实验:分别用Adam(fp32)和8bit版本训练500步,观察loss曲线是否保持一致。如果差异很小,再用于正式训练。

4.4 分布式训练中的ZeRO与优化器状态分片

在多卡训练时,还有一种非常有效的手段是分片优化器状态。DeepSpeed的ZeRO Stage 1就是一个典型实现:将优化器状态切分到多张卡上,每张卡只保存自己负责的那一份,到更新参数时通过通信聚合。这样虽然训练通信量略有增加,但每张卡的内存占用可以大幅下降。

ZeRO Stage 1对Adam的效果特别明显,因为Adam的状态占比高。举个例子,如果一个模型的Adam状态需要12.8GB,8卡时原本需要8份也就是102.4GB的总内存开销,但ZeRO Stage 1分片后,每卡只保留1.6GB左右。这个差距在工程上是决定能否跑起来的因素。不过ZeRO的配置相对复杂,一般用在多卡或者Megatron等框架的训练中,单卡场景下暂时用不上。

4.5 一个经验:上线前先做“内存预演”

最后分享一个我在实际项目中养成的习惯:在正式大规模训练之前,先写一个很小的脚本,初始化一个小模型(比如原模型的1%参数),分别用SGD和Adam跑几十步,记录显存占用曲线。然后用公式放大到全模型规模,估算正式训练所需的显存。这个“内存预演”过程用不了多少时间,但能帮你在开会的时候直接给出明确结论,比如“这个模型用Adam至少要4卡A100,如果改成动量SGD,3卡就够了”。

这个习惯也让我在实际调参时少了很多措手不及的情况。很多OOM问题,其实在你开始写训练代码之前就能提前预判,关键还是得对优化器的内存账本有足够清晰的认知。

SGD和Adam在内存占用上的区别,本质上是“是否记住历史梯度信息”的工程代价。Adam用更多内存换来了更鲁棒的收敛过程,而SGD用更少内存换来了更可控的传统动量行为。在显存充裕时,我倾向直接使用AdamW省心;在显存紧张时,则需要仔细算出这笔账,选择分片、8bit转换或者回归动量SGD。只有把优化器状态这部分开销和模型参数、梯度、激活值放在同一个脑子里统筹规划,才算是真正把显存利用好。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询