尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

模型优化器实战:从AdamW到ZeRO的显存与吞吐优化

发布时间:2026/9/29 19:35:02

资讯中心
01
ARTICLE

模型优化器实战:从AdamW到ZeRO的显存与吞吐优化

模型优化器实战:从AdamW到ZeRO的显存与吞吐优化
1. 模型优化器到底在解决什么问题第一次接触 Model-Optimizer 这个概念是在一个推荐系统的排序模型上。当时线上推理延迟卡在 180ms 下不去GPU 利用率却只有 30% 出头显存倒是先爆了。排查了一圈发现问题不在模型结构也不在数据管道而是优化器状态把显存吃掉了将近一半——Adam 的动量和方差两份状态参数量乘以 2 再乘以 4 字节一个 7B 参数的模型光优化器状态就接近 56GB。这个场景让我意识到优化器从来不只是“训练时选个 Adam 还是 SGD”这么简单它直接决定了你能不能训得动、训得快、训得省。Model-Optimizer 这个标题字面看是“模型优化器”但它背后其实覆盖了三个层次的东西。第一层是优化算法本身也就是 SGD、Adam、AdamW、Lion、Sophia 这些更新规则的数学形式与超参配置第二层是优化器的工程实现包括状态分片、混合精度、梯度累积、CPU offload 这些让优化器在真实硬件上跑得起来的手段第三层是面向部署的模型优化比如量化、剪枝、蒸馏、算子融合这些虽然不叫“优化器”但在工程语境里经常被归到同一个话题下讨论。我写这篇东西是想把这三层拆开讲清楚。适合谁看如果你正在训一个中等规模以上的模型发现显存不够、吞吐上不去、loss 震荡或者你准备把训练好的模型推到线上纠结量化掉点多少能接受那这篇内容应该能帮你少走一些弯路。如果你只是刚入门想搞清楚 Adam 和 SGD 到底差在哪我也会用生活化的类比把它讲明白。全文基于我自己的实操记录和踩坑经验参数和结论都尽量给出可复现的依据。2. 优化算法选型从 SGD 到自适应方法的取舍逻辑2.1 为什么 Adam 不是万能答案很多人上手就是optimizer Adam(lr1e-3)这几乎成了默认操作。Adam 的好处确实明显自适应学习率让它在稀疏梯度、不同量级的参数上都能比较稳地收敛对学习率不那么敏感。但它的代价也很实在。Adam 为每个参数维护一阶矩估计动量和二阶矩估计方差这意味着优化器状态的内存开销是参数量的两倍fp32 下。一个 1B 参数的模型参数本身 4GB优化器状态 8GB再加上梯度 4GB光这三样就 16GB还没算激活值。这就是为什么大模型训练里优化器状态往往是显存占用的头号大户。另一个问题是泛化。有大量实验表明在同等训练步数下Adam 收敛更快但最终测试集表现有时不如调好的 SGD Momentum。这在图像分类任务里尤其明显。我的理解是自适应方法容易在训练后期“过于自信”对平坦极小值的探索不如 SGD 充分。所以如果你追求的是极致精度而不是训练速度SGD 配合余弦退火仍然值得一试。2.2 AdamW 与权重衰减的正确姿势AdamW 的出现解决了一个长期被忽视的 bug原始 Adam 里的 L2 正则和真正的权重衰减并不等价。在 Adam 中加 L2梯度会被自适应学习率缩放导致大梯度的参数实际衰减更小正则效果被扭曲。AdamW 把权重衰减从梯度里拆出来直接作用在参数上公式上就是p p - lr * wd * p这一步独立于自适应部分。实操上我一般这样配from torch.optim import AdamW optimizer AdamW( model.parameters(), lr2e-5, betas(0.9, 0.999), eps1e-8, weight_decay0.01, )weight_decay0.01是 Transformer 类模型的常见起点CNN 可以试 5e-4 到 1e-4。eps在混合精度训练下建议放大到 1e-6 甚至 1e-4因为 fp16 的数值范围小太小的 eps 会导致除零或数值不稳定。这个细节很多教程不讲但我在 fp16 训练里踩过loss 直接变 NaN排查半天才发现是 eps 的问题。2.3 Lion、Sophia 这些新优化器值不值得换Lion 是 Google 提出的核心是用符号函数代替 Adam 的二阶矩只维护动量一份状态显存直接省一半。它的更新方向只取梯度的符号配合动量做平滑。实测下来Lion 在同等显存下能开更大的 batch或者在同等 batch 下省显存。但它的学习率通常要比 Adam 小 3 到 10 倍weight_decay 要大 3 到 10 倍超参搜索空间和 Adam 不一样不能直接套。Sophia 则是用二阶信息做预条件号称收敛更快。但它的实现复杂度高对 Hessian 估计的噪声敏感我在小规模实验里没跑出明显优势就暂时搁置了。我的建议是新优化器不要盲目上生产。先在中小规模上做对照实验固定其他变量只看优化器的影响。如果收益不明显用 AdamW 的成熟生态更稳妥。优化器换错导致的训练不稳定排查成本远高于那点理论收益。3. 工程实现让优化器在真实硬件上跑起来3.1 显存账本优化器状态到底占多少先算一笔账。假设模型参数量为 P混合精度训练下项目精度占用模型参数fp162P 字节模型参数副本fp324P 字节梯度fp162P 字节梯度副本fp324P 字节Adam 一阶矩fp324P 字节Adam 二阶矩fp324P 字节合计20P 字节一个 7B 模型20P 就是 140GB。单卡 80GB 根本放不下。这就是为什么必须做优化器状态分片。3.2 ZeRO 与优化器状态分片ZeROZero Redundancy Optimizer的思路很直接数据并行下每张卡都存一份完整的优化器状态是浪费因为每张卡算出的梯度最终要 all-reduce 成一样的。那不如把优化器状态切开放到不同卡上每张卡只更新自己负责的那部分参数更新完再广播给其他卡。ZeRO Stage 1 切优化器状态Stage 2 再切梯度Stage 3 连参数也切。Stage 1 的通信量和纯数据并行一样但显存省了 N 倍N 为卡数。Stage 3 省得最多但通信量最大。实操上如果你用 DeepSpeed配置大概长这样{ zero_optimization: { stage: 2, offload_optimizer: { device: cpu, pin_memory: true }, allgather_partitions: true, overlap_comm: true, reduce_scatter: true }, fp16: { enabled: true, loss_scale: 0, initial_scale_power: 16 } }offload_optimizer把优化器状态放到 CPU 内存进一步省显存代价是更新时要在 CPU 和 GPU 之间搬数据速度会慢一些。overlap_comm让通信和计算重叠能补回一部分性能。我实测在 8 卡 A100 上Stage 2 CPU offload 能把 13B 模型的训练塞进单机吞吐大约是纯 GPU 方案的 70%。3.3 梯度累积与学习率缩放显存不够时另一个常用手段是梯度累积跑多个 micro-batch把梯度累加起来再更新一次。这样等效于更大的 batch size但显存只按 micro-batch 算。这里有个容易错的点学习率要不要跟着放大。如果原来 batch size 是 32现在用 4 个 micro-batch 累积成 32那学习率不用变。但如果你是想通过累积把等效 batch 从 32 提到 256那学习率通常要按线性或平方根规则放大。线性规则是lr_new lr_old * (batch_new / batch_old)平方根规则是开根号。大 batch 训练里平方根更稳线性在 warmup 阶段配合使用。accumulation_steps 8 for i, batch in enumerate(loader): loss model(batch) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()注意loss要除以累积步数否则梯度会放大对应倍数。这个除法位置也有讲究放在 backward 前除和放在 loss 计算时除数值上等价但前者更清晰。4. 面向部署的模型优化量化、剪枝与蒸馏4.1 量化从 fp32 到 int8 的收益与代价训练完的模型要上线推理成本是大头。量化把 fp32 权重压到 int8模型体积直接小 4 倍推理速度在支持 int8 的硬件上能快 2 到 4 倍。但掉点多少取决于量化的粒度和方法。训练后量化PTQ最简单拿校准数据集跑一遍统计每层的激活范围算出 scale 和 zero_point。问题是激活值的分布往往有长尾一刀切会截断重要信息。量化感知训练QAT在训练时插入伪量化节点让模型适应量化误差掉点通常能控制在 1% 以内但需要重新训练。实操上我一般先用 PTQ 试如果掉点超过可接受范围比如分类任务 top-1 掉超过 1%再上 QAT。PyTorch 的torch.quantization和torch.ao.quantization都提供了现成接口但要注意算子支持情况不是所有 op 都能量化遇到不支持的会 fallback 到 fp32反而拖慢速度。4.2 剪枝结构化与非结构化的选择剪枝分两种。非结构化剪枝把单个权重置零稀疏度高但硬件不一定加速因为 GPU 对稀疏矩阵的支持有限。结构化剪枝直接砍掉整个通道或注意力头硬件友好但精度损失更大通常需要微调恢复。我的经验是非结构化剪枝适合研究场景追求极致压缩率结构化剪枝适合工程落地因为能真正减少计算量。剪枝比例不要一次到位逐步剪、逐步微调比如每次剪 10%微调几个 epoch再剪下一轮。一次性剪 50% 基本会崩。4.3 蒸馏让小模型学到大的蒸馏是用大模型teacher的软标签指导小模型student训练。软标签里包含了类别间的相对概率信息比硬标签信息量更大。温度参数 T 控制软标签的平滑程度T 越大分布越平通常 2 到 5 之间。蒸馏的坑在于 teacher 和 student 容量差距太大时student 学不动。一般建议 student 参数量不低于 teacher 的 10%。另外蒸馏 loss 和原始 CE loss 的权重也要调常见是 0.5 比 0.5 或 0.7 比 0.3看任务。5. 常见问题与排查技巧实录5.1 Loss 变 NaN 的排查顺序这是训练里最高频的问题。我的排查顺序是检查 eps混合精度下 eps 太小会导致除零先放大到 1e-6。检查 loss scalefp16 动态 loss scale 如果一直往下掉说明梯度溢出频繁可以手动设一个初始值。检查数据有没有异常值、空标签、除零操作。检查学习率太大直接炸先降到 1e-5 试。检查梯度裁剪加clip_grad_norm_(model.parameters(), 1.0)能挡掉大部分梯度爆炸。5.2 显存够但吞吐上不去这种情况通常是数据加载或通信成了瓶颈。先看 GPU 利用率如果忽高忽低多半是 dataloader 的num_workers不够或者预处理太重。把num_workers调到 CPU 核数的 2 到 4 倍试试。如果是多卡训练看通信占比NCCL 的 all-reduce 如果耗时超过计算的 30%考虑用梯度累积减少通信频率或者换更快的互联。5.3 优化器状态恢复后 loss 对不上断点续训时除了模型参数和优化器状态学习率调度器的状态、梯度缩放器的状态、随机数种子都要存。少存一个恢复后 loss 曲线就可能对不上。我习惯把所有这些打包成一个 dict用torch.save一起存恢复时一起 load。问题现象可能原因排查动作loss 变 NaNeps 太小、lr 太大、数据异常放大 eps、降 lr、查数据显存 OOM优化器状态、激活值、batch 太大开 ZeRO、梯度累积、减 batch吞吐低dataloader 瓶颈、通信瓶颈加 workers、overlap 通信恢复后 loss 跳变状态没存全存 scheduler、scaler、seed量化掉点多激活长尾、算子不支持换 QAT、检查 op 支持5.4 几个容易被忽略的实操心得第一优化器的 betas 不要乱改。(0.9, 0.999)是经过大量验证的默认值改成(0.99, 0.999)会让动量更平滑但响应变慢除非你有明确理由否则别动。第二warmup 不是可选项。Transformer 类模型没有 warmup 直接上大学习率前几百步基本必炸。warmup 步数一般是总步数的 1% 到 5%小模型取小值大模型取大值。第三梯度裁剪的阈值要按模型调。1.0是常见起点但 RNN 类模型可能需要更小比如 0.5而某些 CV 模型可以放到 5.0。看梯度范数的实际分布来定别照搬。第四混合精度下优化器状态建议保持 fp32。虽然 fp16 状态能省显存但更新时的数值误差会累积长期训练容易出问题。省显存优先用 ZeRO而不是降优化器状态的精度。6. 一套可复现的中等规模训练配置把上面的东西串起来给一套我实际用过的配置。场景是单机 8 卡 A100 80GB训练一个 1.3B 参数的 Transformer混合精度ZeRO Stage 2。import torch from torch.optim import AdamW from torch.cuda.amp import GradScaler optimizer AdamW( model.parameters(), lr1e-4, betas(0.9, 0.95), eps1e-6, weight_decay0.1, ) scheduler torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr1e-4, total_stepstotal_steps, pct_start0.03, anneal_strategycos, ) scaler GradScaler(init_scale2**16) for step, batch in enumerate(loader): with torch.cuda.amp.autocast(dtypetorch.float16): loss model(batch) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() scheduler.step() optimizer.zero_grad(set_to_noneTrue)几个关键点betas的第二项用 0.95 而不是 0.999是因为大模型训练里梯度噪声大0.95 响应更快OneCycleLR的pct_start0.03对应 3% 的 warmupset_to_noneTrue比zero_grad()省一点显存和时间因为直接把梯度置 None 而不是填零。这套配置在 1.3B 模型上8 卡能跑到大约 1200 tokens/s显存占用每卡 62GB 左右。如果换 Lion学习率要降到 3e-5weight_decay 提到 0.3显存能降到 45GB 左右但收敛步数会多一些。最后分享一个我踩过的坑有次训练中途换了优化器从 AdamW 换到 Lion结果 loss 直接飙上去。原因是 Lion 的更新幅度和 AdamW 完全不同学习率没跟着调。换优化器一定要重新做学习率扫描不能沿用旧值。这个教训让我后来养成了习惯任何优化器变更都当成一次独立的超参搜索来做而不是简单替换。
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。