AI 中文总结
研究MoE训练中优化器状态存储位置,提出SkewAdam方法,根据参数群体差异分层分配状态,大幅减少内存占用,提升训练效率,验证了优化器状态存储位置对训练效果的重要性。
AI 中文摘要
优化器状态是专家混合(MoE)训练内存预算中最大的单项:在一个67.8亿参数的MoE语言模型上,AdamW保留50.6GB的一阶和二阶矩来更新12.6GB的bfloat16权重。我们研究了SkewAdam,它基于MoE的三个参数群体(密集主干、专家和路由器)在大小和梯度统计上差异很大,不应接收相同状态这一观察构建。SkewAdam为骨干(5%的参数)保留float32动量加分解的二阶矩,为专家(95%)仅保留分解的二阶矩,为路由器(<0.01%)保留精确的二阶矩。最终状态占用1.29GB,为AdamW的2.6%,峰值训练内存从81.4GB降至31.3GB,在40GB加速器的预算内。在8200万个令牌的相同初始化控制比较中,SkewAdam达到验证困惑度108.4,领先于AdamW(126.8)、Muon(120.2)和Lion(393.7),并将路由器负载平衡稳定在其均匀下限的1%以内。分层分配并非是获得该困惑度的原因:分层消融用二十倍的状态与之匹配,而共享分解估计器但放弃动量的Adafactor则落后40分。分层在不损失准确性的情况下节省内存;准确性来自保留动量,这也是均匀优化器所共有的。扫描基线的学习率缩小了但未弥合差距:最佳调整的AdamW达到118.5,调整后的Adafactor为139.7。这些结果表明,优化器状态的存储位置至少与存储量一样重要。
英文摘要
Optimizer state is the largest single line item in the memory budget of mixture-of-experts (MoE) training. On a 6.78B-parameter MoE language model AdamW keeps 50.6 GB of first and second moments to update 12.6 GB of bfloat16 weights. We study SkewAdam, an optimizer built on the observation that the three parameter populations of an MoE differ enough in size and gradient statistics that they should not receive the same state. Those populations are the dense backbone, the experts and the router. SkewAdam keeps float32 momentum plus a factored second moment for the backbone (5% of parameters), a factored second moment alone for the experts (95%) and an exact second moment for the router (<0.01%). The resulting state occupies 1.29 GB or 2.6% of AdamW's and peak training memory falls from 81.4 GB to 31.3 GB, within the budget of a 40 GB accelerator. In a controlled comparison from identical initializations over 82M tokens, SkewAdam reaches validation perplexity 108.4, ahead of AdamW (126.8), Muon (120.2) and Lion (393.7), and settles router load balance to within 1% of its uniform floor. The allocation is not what earns that perplexity. A tier ablation reaches the same value while carrying twenty times the state, so the tiers buy memory rather than accuracy. Same-platform runs separate what does earn it. Removing momentum costs 31 perplexity points (tuned Adafactor, 139.7) and replacing the factored second moment and its update clipping with a full second moment costs 10 (tuned AdamW, 118.5), so neither tuned baseline reaches the untuned tiered policy. Where optimizer state lives, these results suggest, matters at least as much as how much of it there is.
Comments12 pages, 4 figures, 9 tables. v2: adds Adam-mini discussion, learning-rate sweeps with repeated seeds for the AdamW and Adafactor baselines, and a tier ablation; corrects the attribution of the perplexity advantage between momentum and the factored estimator. Code and per-run training logs: https://github.com/nuemaan/skewadam