SGD与Adam显存占用差多少?优化器状态内存详解
训练深度学习模型时最让人抓狂的报错之一就是“CUDA out of memory”。同一个模型用SGD训练明明还能跑换成Adam就立刻爆显存。不少人一开始都以为是batch size没调好或者模型本身太大了其实问题很可能出在优化器身上——优化器不是简简单单帮你更新参数的工具它自己也会在显存里占一块地方而且SGD和Adam要占的内存完全不是一个级别。这篇文章就把这件事彻底讲明白两者差在哪、具体差多少、用代码怎么实测、以及显存吃紧的时候有哪些真正有效的降压手段。顺带解决一个容易混淆的概念系统里那些“内存占用高但是实际没开程序”“Edge浏览器内存占用高”的问题和深度学习训练里的显存占用完全是两码事不要搞混排查方向。1. 内存到底差在哪先看两个优化器各自要存几份张量1.1 朴素SGD最省但别漏了梯度这一份最原始的SGD更新公式非常简单θ_{t1} θ_t - lr * ∇θ_t。它只需要知道当前的参数和当前的梯度算出结果后直接覆盖更新不需要记住任何历史信息。所以朴素SGD在训练阶段“必须驻留”的张量其实只有两个一份模型参数一份当前梯度。这两个张量的形状和模型参数完全一致如果都用float32存储每个参数占4字节那内存就是2 * 参数量 * 4字节。很多人算内存时只盯着模型参数这一步就漏了梯度。实际上在训练过程中反向传播算出来的梯度会保留在parameter.grad这个字段里并不会立刻释放直到调用optimizer.zero_grad()才会被清掉。也就是说即使是最省内存的朴素SGD实际训练时也要同时容纳“参数梯度”两份数据。1.2 动量SGD多出一份“历史梯度”SGD加了动量之后更新公式就从一步计算变成了“先更新动量再更新参数”v_t momentum * v_{t-1} ∇θ_tθ_{t1} θ_t - lr * v_t这里的v就是“历史梯度的指数移动平均”你需要为每个参数额外维护一个和参数形状完全相同的动量张量。PyTorch里的torch.optim.SGD(momentum0.9)就是这种实现优化器的state里会保存这个动量向量。于是带动量SGD的内存构成变成了三份参数、梯度、历史梯度动量。对应字节数就是3 * 参数量 * 4字节。这里顺便说一句Nesterov动量nesterovTrue在PyTorch里也只要一份动量张量所以内存占用和普通动量SGD基本一致不会像Adam那样翻倍多。1.3 Adam直接两份额外状态一阶矩和二阶矩Adam和SGD最大的不同在于Adam除了存储模型参数和梯度还会为每个参数额外维护两个状态张量一阶动量估计mexp_avg和二阶动量估计vexp_avg_sq。PyTorch的torch.optim.Adam源码里体现得很明确第一次step()时会根据参数的形状创建两个和参数同尺寸的零张量分别保存exp_avg和exp_avg_sq。所以Adam在训练时需要同时驻留四份同尺寸张量模型参数、梯度、一阶动量、二阶动量。按float32算就是4 * 参数量 * 4字节。这也是很多训练框架里提到“Adam吃显存”的根源。AdamW呢AdamW只是在解耦权重衰减上做了改进但存储结构和Adam一模一样同样需要两份动量状态内存占用没有变化。2. 数据说话一个模型在SGD和Adam下的内存差距能有多大2.1 优化器更新公式决定了存储规格要准确计算优化器状态的内存不能只停在“几份”这个层面还得看更新公式里每一步到底保存了什么。Adam每一步的更新是这么做的m_t β1 * m_{t-1} (1 - β1) * g_tv_t β2 * v_{t-1} (1 - β2) * g_t^2m̂_t m_t / (1 - β1^t)v̂_t v_t / (1 - β2^t)θ_{t1} θ_t - lr / (√v̂_t ε) * m̂_t这里的m_t和v_t分别对应一阶矩和二阶矩它们必须在两次迭代之间持续保存否则下一次迭代无法计算。所以无论你的业务代码里写了多少行优化器配置Adam最终都会占用两份额外张量。这就是公式层面的必然结果。从存储量级看参数量为Nfloat32占4字节时各种优化器的理论占用可以整理成下面这张表优化器需要驻留的张量体积计算公式相对朴素SGD的倍数SGD无动量参数、梯度2 * N * 4B1.0xSGD Momentum参数、梯度、动量3 * N * 4B1.5xAdam / AdamW参数、梯度、一阶动量、二阶动量4 * N * 4B2.0x这里算的是“优化器相关”的显存实际训练时还会叠加激活值、中间计算结果、通信缓冲等后面再说。2.2 算一下700万参数模型的账用一个直观的例子。假设模型参数量N 7,000,000大概是BERT-tiny级别或一个小型CNN的规模。float32下每份参数的体积是7,000,000 * 4 / 1024 / 1024 ≈ 26.7 MB朴素SGD参数26.7MB 梯度26.7MB ≈ 53.4MB动量SGD再加一份26.7MB ≈ 80.1MBAdam参数26.7MB 梯度26.7MB 一阶26.7MB 二阶26.7MB ≈ 106.8MB所以一个700万参数的模型Adam比朴素SGD多约53MB比动量SGD多约27MB。放在8GB、12GB显存上看这个差距确实不痛不痒。难怪很多小模型训练时大家根本不关心优化器状态毕竟几十MB相对于几个GB显存完全是毛毛雨。2.3 放大到上亿参数差距才开始吓人但模型规模一旦上去情况完全不同。现在做LLM微调动辄就是7B、13B参数一张推理卡甚至都放不下模型本身。如果拿7B参数来算7,000,000,000 * 4 / 1024 / 1024 / 1024 ≈ 26.1 GB朴素SGD需要约52.2GB动量SGD需要约78.3GBAdam需要约104.4GB也就是说一个7B参数的模型用Adam训练时优化器状态加参数、梯度就要100GB级别单张40GB卡根本放不下换成朴素SGD则只需要52GB左右压力小很多。这还只是“和优化器直接相关”的部分实际要算上激活值只会更夸张。所以大模型训练时为什么普遍流行SGD/Momentum、8bit Adam、Adafactor、Lion这类节省优化器状态的方案就是因为Adam在这种规模下真的撑不住。3. 实操用PyTorch量一量优化器状态到底占多少字节3.1 step之前state是空的先看状态字典很多人在调试时盯着optimizer.state_dict()看发现里面什么都没有以为自己的代码写错了。其实这是PyTorch的延迟初始化机制exp_avg、exp_avg_sq这些状态张量是在第一次调用optimizer.step()时才创建的。所以想看到Adam的真实状态需要先做一次forward、backward、step的完整流程。可以先跑一段很简单代码把优化器状态字典内容打出来import torch import torch.nn as nn import torch.optim as optim model nn.Linear(8, 8).cuda() adam optim.Adam(model.parameters(), lr1e-3) # 此时 state_dict 基本是空的 print(Step 前:, adam.state_dict()) x torch.randn(2, 8).cuda() loss model(x).sum() loss.backward() adam.step() # 第一次 step 后exp_avg 和 exp_avg_sq 才会出现 for name, val in adam.state_dict().items(): print(name, -, val.keys() if isinstance(val, dict) else type(val))输出里你会看到state这一项已经从空字典变成了包含exp_avg、exp_avg_sq和step的结构。这个设计平时不会影响使用但在做内存统计时是个非常容易踩的坑。3.2 统计优化器状态内存的实测代码既然知道了state里保存什么写一个统计函数就不难。核心思路是遍历optimizer.state中每个参数对应的状态张量累加numel() * element_size()def optimizer_state_bytes(model, optimizer): total 0 for p in model.parameters(): if p in optimizer.state: for _, value in optimizer.state[p].items(): if isinstance(value, torch.Tensor): total value.nelement() * value.element_size() return total然后同一模型分别用SGD、momentum SGD、Adam做一次完整更新看看数值对不对得上def run_with(optimizer): model nn.Sequential( nn.Linear(1024, 1024), nn.ReLU(), nn.Linear(1024, 1024), ).cuda() opt optimizer(model.parameters()) x torch.randn(8, 1024).cuda() loss model(x).sum() loss.backward() opt.step() return optimizer_state_bytes(model, opt) print(SGD:, run_with(lambda params: optim.SGD(params, lr0.01))) print(SGD Momentum:, run_with(lambda params: optim.SGD(params, lr0.01, momentum0.9))) print(Adam:, run_with(lambda params: optim.Adam(params, lr1e-3)))以两层1024的全连接层为例参数量大概是1024*1024 1024 1024*1024 1024 ≈ 2,099,200。理论值分别是SGD无状态0字节因为state为空、动量SGD约8MB、Adam约16MB。实测输出会和这个量级一致。要注意的是这段代码统计的只是优化器状态本身并不包含模型参数和梯度如果你想看训练全过程的显存峰值需要配合下面说的方法。3.3 用CUDA内存接口看整体显存变化只看优化器状态还不够实际上训练中显存是多种因素交织在一起的。推荐用PyTorch自带的内存统计接口来看torch.cuda.reset_peak_memory_stats() torch.cuda.reset_max_memory_allocated() # 跑一段训练逻辑 model nn.Sequential(nn.Linear(2048, 2048), nn.ReLU(), nn.Linear(2048, 2048)).cuda() opt optim.Adam(model.parameters(), lr1e-3) x torch.randn(16, 2048).cuda() loss model(x).sum() loss.backward() opt.step() print(当前分配显存:, torch.cuda.memory_allocated() / 1024**2, MB) print(历史峰值显存:, torch.cuda.max_memory_allocated() / 1024**2, MB)更详细一点的报告可以用torch.cuda.memory_summary()它会按张量段位打印显存分配情况包括缓存池、使用率等。看到这份报告你就会明白一个很小的模型里激活值往往比优化器状态大得多。我之前在调试一个序列长度很长的Transformer时把优化器从Adam换成SGD发现峰值显存几乎没变化原因就是激活值占了90%以上的空间优化器状态在里面根本排不上号。这也是很多人“换了优化器感觉没用”的原因。4. 显存吃紧时的应对策略从换优化器到压状态精度4.1 能跑就尽量换SGD/Momentum说句实在话如果你的模型规模在一千万参数以内Adam和SGD的内存差异通常不会成为瓶颈该用哪个用哪个别为了省显存牺牲调参便利性。但当你面对的是上亿甚至几十亿参数模型就得认真盘算优化器状态这笔账了。一个最直接的做法就是在不降低效果的前提下换成带动量的SGD。很多CV分类任务里SGDMomentum本身就有非常成熟的学习率策略效果并不比Adam差甚至泛化性更好。NLP任务里Adam更稳但在部分微调场景下Momentum SGD经过仔细调参也能逼近Adam的效果。不过切换优化器不是改一行代码就完事学习率几乎必须重调。Adam的自适应学习率让它可以容忍比较大的初始学习率波动而SGD对学习率非常敏感我见过很多人直接把Adam的1e-4套到SGD上结果loss直接发散或者收敛慢到怀疑人生。比较靠谱的做法是先用Adam做小规模实验确定大致的量级和训练策略再切换到SGD并配合cosine或warmup调度从头调一遍学习率。4.2 混合精度其实省的主要是激活不是优化器状态说起省显存很多人第一反应就是开混合精度AMP。这个概念要分清楚PyTorch AMP主要把前向和反向计算中的很多张量变成FP16或BF16从而大幅减少激活值和中间计算结果的占用。但优化器状态通常还是保留FP32尤其Adam里的一阶、二阶动量依然以FP32存储因为直接压成FP16很容易影响训练稳定性。所以混合精度真正省的是“模型计算过程中的动态内存”优化器状态那部分并不能靠AMP直接砍掉。对大模型训练来说更彻底的做法是让优化器操作一个单独维护的FP32主权重副本而真正做前向和反向的是FP16/BF16的模型副本。这在DeepSpeed、Megatron-LM里很常见能同时压住激活值和参数副本的显存。但如果你只是在单卡上用PyTorch自带的torch.cuda.amp和GradScaler请不要指望它把Adam的state省掉那是另一套机制。4.3 8bit优化器把Adam的动量状态压到1/4如果既想要Adam的自适应能力又想大幅降低内存可以考虑8bit优化器。这个方案最早由Dettmers等人提出核心思路是把一阶动量、二阶动量分块量化成8bit整数存储更新参数时再临时反量化成FP32做计算。论文里的结论是8bit Adam在大部分场景下能接近32bit Adam的效果而优化器状态内存直接降到原来的四分之一左右。在PyTorch生态里用起来也简单bitsandbytes库就提供了现成实现import bitsandbytes as bnb # 替代 torch.optim.Adam opt bnb.optim.Adam8bit(model.parameters(), lr1e-3) # 或 AdamW8bit opt bnb.optim.AdamW8bit(model.parameters(), lr1e-3)如果你在用Hugging Face Transformers的训练器也可以直接传optimizerbnb.optim.AdamW8bit。实测下来8bit Adam确实能把优化器状态这部分显存压得很低尤其在7B以上的模型微调时省出的显存十分可观。需要提醒的是8bit优化器在安装时对环境有依赖部分老显卡或特殊硬件上可能出现兼容问题另外量化本身会带来少量精度损失对异常敏感的模型需要先做小规模验证。4.4 梯度累积与激活检查点减少峰值显存很多时候显存爆掉并不全是Adam的锅而是激活值峰值太高造成的。假设一个Transformer模型前向传播要保存每一层的中间输出batch size一大激活值就会以非常夸张的速度增长。这种情况换什么优化器都救不了应该先做下面几件事第一梯度累积。把一个大batch拆成几个小batch每个小batch独立做前向和反向累积梯度到一定步数后再做一次优化器更新。这样等效batch size不变但峰值激活内存被压到了单个小batch的水平。代码上就是循环里多写一行loss.backward()设置accumulation_steps之后再optimizer.step()。第二激活检查点activation checkpointing。这是以时间换空间的技术它不保存每一层全部激活而是在反向传播时重新计算中间值。PyTorch里可以用torch.utils.checkpoint.checkpoint包住某些模块Hugging Face模型的gradient_checkpointing_enable()也是同一套路。它的效果非常显著代价是训练时间可能会明显变长。第三别忘了用torch.cuda.memory_summary()看清楚显存到底花在哪。先诊断再做优化这是我踩了几次坑后学到的经验。否则你可能花了一晚上换优化器调参结果发现支配显存的是激活值方向完全错了。5. 常见问题速查与避坑清单5.1 五个高频问题一次说清很多人看完上面的分析后还是会遇到一些具体困惑我整理了几个问得最多的问题直接给结论。问题原因与对策为什么我的优化器state一开始是空的PyTorch延迟初始化首次step()才会创建状态张量。调试时先跑一次完整的前反向和step。SGD不带动量和Adam差多少同样是float32下SGD约2份张量Adam约4份张量后者是前者的2倍。但带动量SGD约3份差距缩小到1.33倍。AdamW和Adam内存占用一样吗一样。它们都维护exp_avg和exp_avg_sq只是权重衰减实现方式不同。8bit Adam真的效果不差吗大部分场景效果接近但不绝对。模型越大、数据越充足时通常越稳敏感任务先小规模验证。推理阶段SGD和Adam有区别吗没有。推理时只用已经训练好的参数优化器状态根本不参与计算。这个差异只存在于训练阶段。第5个问题经常被忽略但它解释了为什么很多人下载一个模型直接推理时无论训练时用SGD还是Adam显存占用都一样。所谓的“内存差异”始终是训练过程专属的。5.2 关于“内存占用高但是实际没开程序”这个热点的边界提示最近经常看到“内存占用高但是实际没开程序”“Edge浏览器内存占用高”“Windows Modules Installer占用内存”之类的讨论这类问题和深度学习训练里的显存占用完全是两个世界。系统ROM中的内存占用通常涉及系统服务、浏览器扩展、缓存机制甚至可能是某个进程的后台服务在搞鬼排查方向是任务管理器、启动项和系统服务而不是去看PyTorch优化器状态。深度学习场景里说的“内存占用”绝大多数时候是指GPU显存顶多包括CPU侧存储参数的RAM但原理和Windows系统内存管理并不一样。看这类内容时不要把两者混在一起否则很容易被无关的优化思路带偏。我在实际训练中的体会是内存这件事一定要先弄清楚瓶颈在哪一个维度。优化器状态只是训练内存里的一块激活值、参数副本、梯度和通信缓冲同样重要。最稳妥的做法是写一个十几行的内存统计脚本在换策略前后分别跑一次用数据判断方向而不是凭感觉认为“Adam一定更耗显存”就去随便换优化器。最后分享一个小技巧当你决定从Adam切换成SGD/Momentum时先别急着动其他超参只用小数据集跑几个step对比一下loss曲线的下降速度再决定学习率该调高还是调低。这一步能替你省下一堆无效训练时间。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →