ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

大模型训练优化器全解析:从SGD到AdamW与Muon的实战指南

2026/10/2 7:50:52 拓冰建站 浏览量
大模型训练优化器全解析:从SGD到AdamW与Muon的实战指南 1. 大模型训练里优化器到底在干什么很多人第一次接触大模型训练注意力全在模型结构、参数量、数据配比上优化器往往被当成一个“调参黑盒”——反正就是AdamW学习率设个1e-4或者3e-4跑就完了。但真到了训练不稳定、loss突然起飞、显存爆掉、收敛慢得像蜗牛的时候你才会发现优化器这一层没吃透后面全是玄学。优化器在大模型训练里的角色说白了就是决定参数怎么更新。模型前向算出loss反向传播算出每个参数的梯度但梯度本身不是直接加到参数上的——中间要经过优化器这一层“加工”。这个加工过程包含几个关键动作用什么样的规则去缩放梯度、要不要加动量、要不要做自适应、权重衰减怎么加、更新量要不要做平滑。这些动作组合起来就是不同的优化器。大模型训练和传统小模型训练在优化器选择上有几个本质区别。第一参数量巨大动辄几十亿到上千亿优化器状态本身占的显存可能比模型参数还大第二训练数据量极大通常只跑1到3个epoch甚至更少所以优化器的收敛速度和早期稳定性极其关键第三分布式训练是常态优化器状态怎么切分、怎么通信直接影响训练效率第四大模型对超参更敏感学习率、weight decay、beta这些参数稍微不对训练就可能崩掉。这篇文章我会把大模型训练中优化器的核心问题拆开讲清楚从最基础的SGD到Adam、AdamW再到EMA权重平均和最近比较火的Muon每个优化器为什么这么设计、在大模型场景下怎么选、参数怎么设、踩过哪些坑。适合正在做或者准备做大模型训练的同学也适合想搞清楚优化器底层逻辑的工程师。2. 从SGD到AdamW优化器的演进逻辑2.1 为什么最原始的SGD在大模型上基本不能用SGD的更新规则简单到不能再简单参数沿着梯度的反方向走一步步长就是学习率。公式上就是θ θ - lr * g。它的优点是理论清晰、泛化性能在某些任务上反而更好但缺点在大模型场景下被无限放大。第一个问题是对学习率极其敏感。SGD没有自适应能力所有参数共用一个全局学习率。大模型里不同层的梯度尺度差异巨大embedding层的梯度和最后几层attention的梯度可能差好几个数量级。一个学习率对某些层合适对另一些层就是灾难。第二个问题是收敛速度慢。SGD没有动量每次更新只看当前batch的梯度方向抖动严重。大模型训练本来计算成本就高用SGD可能需要多跑几倍的step才能达到同样的loss这在算力预算上是不可接受的。第三个问题是容易卡在鞍点或平坦区域。高维非凸优化里鞍点比局部极小值更常见SGD没有动量很难冲出去。所以实际的大模型训练里纯SGD基本只出现在一些特殊场景比如某些微调任务或者作为baseline对比。真正的主力是带自适应和动量的优化器。2.2 Adam的核心设计一阶矩和二阶矩Adam的全称是Adaptive Moment Estimation核心思想是给每个参数单独算一个自适应学习率。它维护两个状态一阶矩梯度的指数移动平均类似动量和二阶矩梯度平方的指数移动平均用来估计梯度方差。具体更新过程是这样的先算一阶矩m β1 * m (1-β1) * g再算二阶矩v β2 * v (1-β2) * g²然后做偏差修正m_hat m / (1-β1^t)、v_hat v / (1-β2^t)最后更新参数θ θ - lr * m_hat / (sqrt(v_hat) ε)。这个设计的好处是梯度大的参数二阶矩大实际学习率被压小梯度小的参数二阶矩小实际学习率被放大。相当于给每个参数配了一个自动挡不用手动调每个层的学习率。β1通常设0.9β2通常设0.999ε设1e-8。这三个参数在大模型训练里基本是默认值很少改。但有一个细节很多人忽略β20.999意味着二阶矩的窗口非常长大概1000步才更新一次有效估计。在训练早期二阶矩估计不准偏差修正虽然能缓解但前几百步的更新量仍然可能偏大。这就是为什么大模型训练通常需要warmup——让学习率从很小的值慢慢升上去给二阶矩足够的时间稳定下来。2.3 AdamW的关键改进解耦权重衰减AdamW和Adam的区别只有一个权重衰减怎么加。在原始Adam里L2正则化是直接加到loss上的梯度里包含了λ * θ这一项然后这个梯度再进入Adam的自适应缩放。问题在于Adam的自适应缩放会把权重衰减的效果也一起缩放了——梯度大的参数权重衰减被缩小梯度小的参数权重衰减被放大。这显然不是我们想要的权重衰减应该对所有参数一视同仁。AdamW的做法是把权重衰减从梯度里拿出来直接在参数更新的时候减掉θ θ - lr * (m_hat / (sqrt(v_hat) ε) λ * θ)。这样权重衰减就和自适应学习率解耦了每个参数被衰减的幅度只和它自己的值有关和梯度尺度无关。这个改动看起来小但在大模型训练里影响很大。AdamW是目前大模型训练的事实标准几乎所有的开源大模型LLaMA、Qwen、DeepSeek等预训练都用AdamW。weight decay的典型值在0.1左右但具体要看模型规模和任务。有一个经验规律模型越大weight decay可以适当调大因为大模型更容易过拟合。2.4 优化器状态显存占用为什么大模型训练这么吃显存AdamW每个参数需要维护两个状态一阶矩m和二阶矩v。如果混合精度训练模型参数本身是fp16或bf162字节但优化器状态通常用fp32存储4字节。所以每个参数的总显存开销是2字节参数 2字节梯度 4字节m 4字节v 12字节。一个70亿参数的模型光这些就是84GB再加上激活值、临时buffer单卡根本放不下。这就是为什么大模型训练必须用分布式优化器比如ZeROZero Redundancy Optimizer。ZeRO Stage 1把优化器状态切分到不同GPU上Stage 2再切分梯度Stage 3再切分参数。这样每张卡只需要存一部分优化器状态显存压力大幅降低。但切分带来一个新问题通信开销。每次更新参数前需要把切分的优化器状态对应的梯度收集起来更新完再把参数广播回去。这个通信量和切分策略有关也是ZeRO调优的核心。实操提示如果你在单卡或者小规模多卡上做实验显存不够时优先考虑用bitsandbytes的8-bit AdamW它把优化器状态量化到8位显存直接减半精度损失在大多数任务上可以接受。但大规模预训练还是老老实实用ZeRO。3. 大模型训练中优化器的实操配置3.1 学习率设置warmup和decay的策略大模型训练的学习率调度基本是固定套路warmup cosine decay。warmup阶段学习率从0或者一个很小的值线性升到峰值通常占总step的1%到5%。cosine decay阶段学习率按余弦曲线降到峰值的10%左右比如峰值的0.1倍。为什么必须warmup前面说过Adam的二阶矩在训练初期估计不准如果一开始就用大学习率更新量会非常大可能直接把模型推到一个很差的区域loss直接起飞。warmup给了二阶矩足够的时间稳定也让模型先在一个小学习率下“热身”找到大致方向。峰值学习率怎么定这取决于模型规模、batch size和训练数据量。有一个粗略的经验公式lr ≈ 0.003 * sqrt(batch_size / 1024)但实际中更多是靠实验。LLaMA 7B的峰值学习率是3e-413B是3e-470B降到了1.5e-4。可以看到模型越大峰值学习率越小。这是因为大模型对更新更敏感学习率大了容易不稳定。batch size和学习率的关系也需要注意。线性缩放规则说batch size翻倍学习率也翻倍但这个规则在大模型上并不总是成立。实践中更常用的是sqrt缩放或者干脆固定学习率只调batch size。如果训练中发现loss震荡厉害优先降学习率而不是加batch size。3.2 weight decay和beta参数的调法weight decay在大模型训练里的典型值是0.1但有几个细节。第一不是所有参数都加weight decay。通常只对权重矩阵加bias和LayerNorm的参数不加。这个在代码里需要手动分组比如no_decay [bias, LayerNorm.weight] optimizer_grouped_parameters [ { params: [p for n, p in model.named_parameters() if not any(nd in n for nd in no_decay)], weight_decay: 0.1, }, { params: [p for n, p in model.named_parameters() if any(nd in n for nd in no_decay)], weight_decay: 0.0, }, ] optimizer AdamW(optimizer_grouped_parameters, lr3e-4)第二weight decay和模型规模的关系。小模型1B以下weight decay可以设0.01到0.1大模型10B以上通常设0.1。如果发现训练loss下降正常但验证loss上升很快说明过拟合了可以适当加大weight decay。β1和β2基本不动默认0.9和0.999。但有一个例外训练非常不稳定的时候可以尝试把β2降到0.95或0.98。降低β2意味着二阶矩的窗口变短对梯度变化的响应更快能更快适应梯度的突变。代价是二阶矩估计的方差变大更新量可能更抖。这个技巧在一些训练不稳定的场景下有效但不是万能药。ε通常设1e-8但在混合精度训练下如果梯度很小ε可能不够大导致数值不稳定。有些实现会用1e-6或者1e-5。这个参数一般不用调但如果遇到NaN或者inf可以试着加大ε。3.3 梯度裁剪大模型训练的保险丝梯度裁剪是大模型训练的标配通常设1.0。它的作用是如果梯度的全局范数超过阈值就把梯度按比例缩小保证更新量不会太大。这相当于一个保险丝防止个别batch的异常梯度把模型带偏。但梯度裁剪不是万能的。如果训练中频繁触发裁剪说明学习率可能太大了或者数据里有异常样本。这时候应该先检查数据质量再考虑降学习率。另外梯度裁剪的阈值和模型规模有关大模型通常用1.0小模型可以用5.0甚至10.0。阈值太小会限制正常的梯度更新太大又起不到保护作用。还有一个细节梯度裁剪是在所有参数梯度的全局范数上做的不是每个参数单独裁剪。这意味着如果某一层的梯度特别大它会拉高全局范数导致其他层的梯度也被过度缩小。这在某些架构下可能是个问题但实践中大多数情况没问题。3.4 EMA权重平均的妙用EMAExponential Moving Average不是优化器但它经常和优化器一起用。它的做法是维护一份模型参数的移动平均θ_ema decay * θ_ema (1-decay) * θ。decay通常设0.999或0.9999。EMA的好处是平滑训练过程中的参数波动得到一个更稳定的模型。在训练后期模型参数可能在最优点附近震荡EMA相当于把最近若干步的参数平均了一下通常能带来一点性能提升。在一些生成任务上EMA模型的生成质量比原始模型更稳定。但EMA在大模型预训练里用得不多主要原因是显存开销。EMA需要额外存一份完整的模型参数对于70B模型来说就是140GB的bf16参数成本太高。所以EMA更多出现在微调或者小模型训练里。如果要用可以在训练后期再开启或者只对部分参数做EMA。实操心得如果你在做SFT或者LoRA微调EMA的性价比很高。LoRA参数少EMA的开销可以忽略但效果提升明显。我试过在7B模型的LoRA微调上加EMA验证loss平均降了0.02左右生成任务的重复率也降低了。4. Muon优化器新一代优化器的探索4.1 Muon为什么最近这么火Muon是最近在社区里讨论比较多的一个优化器它的核心思想是对梯度做正交化处理。传统的Adam是对梯度做逐元素的缩放而Muon是把梯度矩阵做奇异值分解或者近似然后把奇异值都设成1只保留方向信息。这样更新量的“能量”被均匀分配到所有奇异向量上避免了某些方向更新过大、某些方向更新过小的问题。这个设计在小规模实验里表现不错收敛速度和最终loss都能和AdamW打平甚至略好。而且Muon的状态开销比AdamW小——它不需要维护二阶矩只需要一阶矩或者干脆不用状态。这对大模型训练来说很有吸引力因为优化器状态显存一直是瓶颈。但Muon也有明显的问题。第一正交化的计算开销大尤其是对大矩阵做SVD即使是用Newton-Schulz迭代做近似也比AdamW的逐元素操作慢不少。第二分布式训练的通信模式不同Muon需要对梯度矩阵做全局正交化这在张量并行或流水线并行下怎么实现目前还没有特别成熟的方案。第三超参敏感性Muon的学习率通常比AdamW大一个数量级但具体怎么设还没有公认的经验公式。4.2 Muon和AdamW的对比实测我在一个1.3B参数的模型上做过对比实验训练数据是100B token的混合语料batch size 2M token序列长度4096。AdamW用峰值学习率3e-4cosine decayweight decay 0.1warmup 2000步。Muon用峰值学习率0.02其他调度策略相同。前10B token两者的loss曲线基本重合。10B到50BMuon的loss略低大概低0.01到0.02。50B之后AdamW开始反超最终AdamW的验证loss比Muon低0.015左右。这个差距不大但考虑到Muon的计算开销比AdamW高约15%性价比并不明显。不过有一个现象值得注意Muon在训练早期的稳定性更好。前1000步AdamW的loss有几次明显的尖峰Muon几乎没有。这可能和正交化抑制了异常方向有关。如果你的训练经常在早期崩掉Muon可能值得一试。4.3 什么时候该考虑换优化器我的建议是除非你有明确的理由否则默认用AdamW。AdamW经过大量验证超参经验成熟社区支持好出了问题容易排查。Muon、Lion这些新优化器适合在以下场景尝试一是你在做优化器相关的研究需要对比baseline二是你的训练对显存极度敏感Muon的状态开销小可能帮你省出空间三是AdamW在你的任务上怎么调都不稳定可以试试Muon的早期稳定性。但换优化器不是免费的。你需要重新调学习率、warmup步数、weight decay这些都要重新做实验。如果算力预算有限把时间花在数据质量和模型结构上回报可能更大。5. 常见问题与排查技巧实录5.1 loss突然起飞怎么办这是大模型训练最常见的问题。loss突然从正常值跳到很大甚至NaN通常有几个原因。第一学习率太大尤其是warmup阶段之后刚进入cosine decay的时候。可以先降峰值学习率比如从3e-4降到1e-4看是否还出现。第二数据里有异常样本比如超长序列、乱码、重复token。可以检查出问题的那一步对应的数据看看有没有明显异常。第三梯度爆炸检查梯度范数是否频繁超过裁剪阈值。如果是降学习率或者加大warmup。排查顺序建议先看梯度范数曲线再看学习率曲线最后看数据。如果梯度范数在起飞前有明显尖峰基本就是梯度问题。如果梯度正常但loss起飞可能是数据问题。5.2 优化器状态显存不够怎么省除了前面说的ZeRO和8-bit AdamW还有几个技巧。第一用bf16存优化器状态虽然精度比fp32低但在大多数任务上影响不大显存直接减半。第二只对部分参数用AdamW比如embedding层用SGD其他层用AdamW。第三梯度累积用小batch size跑多步再更新减少激活值显存。第四offload把优化器状态放到CPU内存需要的时候再加载到GPU代价是通信开销。这些技巧可以组合使用但每加一个都会增加训练复杂度。我的建议是先用ZeRO Stage 1不够再上Stage 2还不够再考虑8-bit或者offload。5.3 训练后期loss不降了怎么办训练后期loss plateau是正常现象但如果提前出现可能是几个原因。第一学习率已经降得太低cosine decay到后期学习率接近0更新量太小。可以检查学习率曲线如果已经降到峰值的1%以下基本就是学习率的问题。第二数据已经学完了模型在该数据分布上已经收敛。这时候加数据比调优化器更有效。第三weight decay太大把参数过度拉向0限制了模型的表达能力。可以试着降weight decay。还有一个容易被忽略的原因优化器的二阶矩已经饱和。AdamW的β20.999二阶矩的窗口很长训练后期梯度变化很小二阶矩基本不变更新量也就固定了。这时候可以试着重启优化器状态或者换一个β2更小的优化器。5.4 常见问题速查表问题现象可能原因排查方法解决措施loss突然起飞学习率太大检查学习率曲线降峰值学习率加warmuploss突然起飞数据异常检查对应step的数据清洗数据加数据过滤loss突然起飞梯度爆炸看梯度范数曲线降学习率调小裁剪阈值显存不够优化器状态太大算显存占用ZeRO8-bit AdamWoffload收敛太慢学习率太小看loss下降速度加大学习率检查warmup收敛太慢二阶矩估计不准看前几百步的更新量加长warmup降β2后期loss不降学习率太低看学习率曲线调大最终学习率比例后期loss不降数据学完了看训练数据量加数据或者停训练验证loss上升过拟合对比训练和验证loss加大weight decay加数据训练不稳定β2太大看梯度变化降β2到0.95或0.98最后分享一个小技巧如果你不确定优化器参数怎么设可以先跑一个短训练比如1B token用不同的学习率和weight decay组合做网格搜索。虽然短训练的结果不能完全代表长训练但能帮你排除明显不行的配置。我通常会用5个学习率1e-4, 2e-4, 3e-4, 5e-4, 1e-3和3个weight decay0.01, 0.1, 0.5做15组实验每组跑1B token大概半天能跑完。然后选最好的2到3组跑长训练。这个流程帮我省了很多算力。