优化器状态应该存放在哪里:面向内存高效混合专家训练的分层状态分配 Where Should Optimizer State Live? Tiered State Allocation for Memory-Efficient Mixture-of-Experts Training
按参数角色分层分配优化器状态,省 97% 显存且不掉精度
前置知识
AdamW 优化器及其状态
AdamW 为每个参数维护两个 float32 动量统计:一阶动量 $m$(梯度的指数移动平均)和二阶动量 $v$(梯度平方的指数移动平均),并用它们做逐坐标自适应缩放更新 $\theta_t = \theta_{t-1} - \eta \cdot \hat{m}_t/(\sqrt{\hat{v}_t}+\epsilon)$。对 bfloat16 权重而言,每个 2 字节权重背后有 8 字节状态影子,是训练显存的最大单项开销。
本文起点正是量化 AdamW 的状态成本——6.78B 参数 MoE 上 AdamW 要 50.6 GB 状态,而权重本身仅 12.6 GB。不理解 AdamW 状态如何占显存,就无法理解为什么要按层分配状态。
混合专家模型 (Mixture-of-Experts, MoE)
MoE 通过门控路由(router/gate)把每个 token 只送到少数专家(experts)处理,解耦总参数量与单 token 计算量。本文模型有 128 个 SwiGLU 专家、top-2 路由、6.78B 总参数但每 token 仅激活约 440M。三个参数群体(稠密主干 5%、专家 95%、路由器 <0.01%)在大小与梯度稀疏度上差异巨大。
分层分配的核心动机就是这三个群体梯度统计完全不同:专家每个只看到约 1/64 的 token,稠密主干每步都见全部 token。不懂 MoE 结构就无法理解为何给它们不同的优化器状态。
Adafactor 的因式分解二阶矩
Adafactor 把 $n\times m$ 矩阵的二阶矩分解为行均值 $R$ 和列均值 $C$ 两个向量,重构秩一估计 $\hat{V}=RC^\top/\bar{R}$,存储从 $nm$ 降到 $n+m$ 个浮点数。4096×4096 矩阵从 64 MB 降到 32 KB。SkewAdam 直接复用此估计器(用均值而非求和)与更新 RMS 裁剪。
SkewAdam 的因式分解部分就是 Adafactor 的估计器,本文不宣称这部分创新。理解 Adafactor 是读懂 Table 2 中 SkewAdam vs Adafactor(一个留动量、一个丢动量)对比的前提。
bfloat16 主权重与抖动随机舍入
纯 bfloat16 训练省显存但小更新会被舍入吃掉。本文主权重存 bfloat16,每步 float32 计算更新后,在类型转换前加一个 ULP 宽度的均匀噪声 $\xi\sim U(-\text{ulp}(w)/2,\text{ulp}(w)/2)$,近似无偏的随机舍入。所有对比优化器走完全相同的写回路径以保证公平。
这是公平对比的技术前提:四个优化器用相同精度路径,谁也不因精度处理占便宜,也解释了为何权重衰减在实际运行中变成空操作(变化量低于 bfloat16 ULP)。
研究动机
MoE 训练中优化器状态是显存预算里最大的一笔单项开销。论文给出具体数字:在 6.78B 参数的 MoE 语言模型上,AdamW 要为更新 12.6 GB 的 bfloat16 权重而保留 50.6 GB 的一阶/二阶动量状态,训练峰值显存高达 81.4 GB,远超 40 GB 级加速器(如 A100/A6000)容量。现存内存高效优化器(Lion、Muon、Adafactor)都把整个网络当成同质化一块:Lion 用一个动量缓冲并按符号更新,会丢弃所有梯度幅度(包括路由器上承载负载均衡信号的相对幅度);Muon 用单个缓冲并对每步更新做 Newton–Schulz 正交化,这是为稠密隐藏层设计的结构先验;Adafactor 对所有矩阵统一套用因式分解。没有任何方法问过 MoE 的不同部分是否该配不同状态,结果 Lion/Muon 状态减半但峰值仍约 57 GB,仍挤不进 40 GB 卡。
本文的目标是本文要设计名为 SkewAdam 的优化器,按参数角色逐层(tier)分配优化器状态:把 Adam 的每种成分只放在它物有所值的地方。具体目标是在 6.78B 参数 MoE 上把优化器状态从 AdamW 的 50.55 GB 大幅压缩,使训练峰值显存降到 40 GB 加速器预算内,同时验证分层分配在不损失(甚至提升)验证困惑度的前提下省下显存。论文还要厘清到底是分层策略带来的精度还是另有原因,并通过对照实验(共享初始化、相同数据顺序、相同 bfloat16 写回路径)和基线学习率扫描证明结论的稳健性。
与已有工作不同的是,现有方法的共同盲点是均匀压缩——要么全网络压缩(Adafactor 对所有矩阵因式分解、对动量一刀切),要么跨设备分片(ZeRO),要么量化(8-bit、SM3),要么统一低秩投影(GaLore),但都把 MoE 当成同质参数袋。SkewAdam 的独特切入角度来自一个观察:MoE 三个参数群体在大小和梯度统计上差异极大——稠密主干(5% 参数)每步都见全部 token,梯度稠密;专家(95% 参数)在 top-2/128 路由下每个平均只处理约 1/64 token,梯度稀疏且方差高;路由器(<0.01% 参数)决定所有 token 去向,其 logits 的相对幅度承载负载均衡信号。论文据此提出:与其问要压多少状态,不如问状态该放在哪里。
核心方法
整体思路先直觉后技术路线。直觉层面:MoE 不是同质参数袋,优化器不必假装它是。稠密主干占 5% 参数但每步都见全部 token,float32 动量便宜且有用;专家占 95% 参数但单步梯度极稀疏(top-2/128 下每个平均只处理约 1/64 token),给它配动量缓冲要花 24 GB 平滑稀疏高方差梯度得不偿失;路由器仅 0.5M 参数,给每个 logit 精确自适应缩放只需 2 MB 却能掌管全部流量。技术路线上,SkewAdam 把每个张量按角色分到三个 tier:BACKBONE(稠密主干:嵌入、注意力、稠密 FFN、norm)持有 float32 动量加因式分解二阶矩;EXPERT(专家)只持有因式分解二阶矩不存动量;ROUTER(路由器)持有完整未分解二阶矩且不做权重衰减。因式分解对 $n\times m$ 矩阵存行均值 $R_t$、列均值 $C_t$,重构秩一估计 $\hat{V}_t=R_tC_t^\top/\bar{R}_t$,最终状态有闭式模型 $S\approx 4N_{bb}$,约 1.29 GB。
核心创新点是分配问题本身——把状态预算斜着(skewed)投向它物有所值的地方,而非均匀压缩。与已有方法的本质区别:Adafactor 在所有矩阵上一刀切因式分解、对动量要么处处留要么处处不留;Lion/Muon 把整个隐藏层(专家与否)当同质对待。SkewAdam 把 Adam 每个成分拆开按 tier 发放:动量只给梯度稠密且便宜的主干;因式分解方差给数量庞大但稀疏激发的专家;精确方差给决定全局流量的路由器。名字即设计——状态预算斜向它该去的地方。论文坦诚这策略不是定理而是判断,并通过 tier ablation 证明:真正带来精度的是保留动量(主干那 1.27 GB),分层本身买的是内存而非精度,所以均匀分配但处处留动量也能达相同困惑度。
方法步骤详情
方法对应 Algorithm 1,输入为处于 tier $\tau$ 的张量 $W$ 及其 float32 梯度 $G$。第一步算 $G\leftarrow\nabla_W\mathcal{L}$(float32)。第二步按角色定二阶矩:若 $\tau=\text{ROUTER}$ 或 $W$ 是向量,用完整 $V\leftarrow\beta_2 V+(1-\beta_2)G\odot G$,分母 $D\leftarrow\sqrt{\max(V,\epsilon)}$;否则更新因式分解 $R,C$,分母 $D\leftarrow\sqrt{RC^\top/\bar{R}}$(带 $\epsilon$ 钳位)。第三步定动量:若 $\tau=\text{BACKBONE}$,$M\leftarrow\beta_1 M+(1-\beta_1)G$,更新 $U\leftarrow\frac{\sqrt{1-\beta_2^t}}{1-\beta_1^t}M\oslash D$;否则不用动量缓冲,$U\leftarrow\sqrt{1-\beta_2^t}G\oslash D$。第四步更新 RMS 裁剪 $U\leftarrow U/\max(1,\text{RMS}(U))$(阈值 1,Adafactor 式)。第五步 bfloat16 写回 $W\leftarrow\text{bf16}(W_{fp32}-\eta_t U+\xi)$,$\xi\sim U(\pm\text{ulp}(W)/2)$ 做抖动随机舍入。超参 $\beta_1=0.9,\beta_2=0.999,\epsilon=10^{-8}$,所有对比优化器共享此写回路径。
技术新颖性
技术新颖性三点。第一,首次把优化器状态分配当作 MoE 训练的直接研究对象,提出按参数角色的分层策略与闭式内存模型,给出可预测的显存账 $S\approx 4N_{bb}$(约 1.29 GB)。第二,把 Adam 的成分解耦——动量、因式分解方差、精确方差分别匹配不同 tier 的梯度统计,而非像 Adafactor 统一因式分解且对动量一刀切。第三,诚实的实验设计本身是新颖性的一部分:通过 tier ablation 精确定位精度来自动量而非分层,并用基线学习率扫描证明领先优势调参后仍在,把贡献严格限定在内存而非吹嘘成更好的优化器。此外作者还实测了一个工程边界——8-bit 优化器内核在单张量达 $2^{31}$ 元素时直接杀进程而非抛异常,而因式分解状态用原生 PyTorch 64 位索引能干净越过该边界,这对更大 MoE 有实际工程价值。
实验结果
核心发现逐表分析。Table 1(6.78B MoE,H200,10000 步共享初始化):SkewAdam 状态 1.29 GB、峰值 31.3 GB、5000 tokens/s、验证 PPL 108.4、均衡损失 0.0505;对比 AdamW 50.55/81.4/126.8,Muon 25.27/57.6/120.2,Lion 25.27/56.6/393.7。仅 SkewAdam 跨过 40 GB 线,吞吐比 AdamW 高 6.6%、仅比最快 Lion 低 1.5%,Muon 因每步对 6.4B 专家跑 Newton–Schulz 低 32%。收敛上前 3000 步落后(第 1000 步比 AdamW 高 114 PPL),第 4000 步反超,终值比 AdamW 低 14.5%。Table 2(H100):Adafactor 状态仅 12 MB 但 PPL 停在 149.5(差 40 点),两者共享因式分解估计器与裁剪,差别只在动量;rank-128 GaLore 式基线直接失败(PPL 1839.9)。Table 3(MI300X ablation):专家动量加回要多花 24 GB 却只移动 0.2 PPL(死重);均匀分配处处留动量在 20 倍状态下达到相同 PPL(108.3 vs 108.9),证明分层买的是内存而非精度。Table 4(学习率扫描):调参后 AdamW 改善到 118.5±0.5、Adafactor 到 139.7,但仍分别落后 SkewAdam 约 10/30 点(SkewAdam 故意未调)。路由上 SkewAdam 与 AdamW 从第 4000 步起在地板 0.05 的 1% 内,Muon 末千步跳到 0.0608(高 22%)。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| 验证困惑度(6.78B MoE,82M token 训练) | Validation Perplexity ↓ | 108.4(H200)/ 109.0(H100)/ 108.9(MI300X),三 GPU 在 0.6 PPL 内一致 | AdamW 126.8(未调);调优后 AdamW 118.5±0.5;Muon 120.2;Lion 393.7;Adafactor 149.5(未调),调优后 139.7 | 比未调 AdamW 低 14.5%,比调优后最佳 AdamW 低约 10 点(20 倍种子级标准差),比调优后 Adafactor 低约 30 点 |
| 峰值训练显存(6.78B MoE) | Peak GPU Memory (GB) | 31.3 GB,可装入 40 GB 加速器 | AdamW 81.4 GB;Muon 57.6 GB;Lion 56.6 GB;Adafactor 29.6 GB | 比 AdamW 降 61%(50 GB),且在精度上反而领先;相比 Adafactor 显存相近但 PPL 低 40 点 |
| 优化器状态显存 | Optimizer State (GB) | 1.29 GB(2.6% of AdamW) | AdamW 50.55 GB;Muon/Lion 25.27 GB;Adafactor 0.01 GB | 结构化节省,不依赖量化,且可与量化叠加 |
| 训练吞吐 | Tokens/s | 5000 tokens/s(H200) | Lion 5075;AdamW 4692;Muon 3409 | 比 AdamW 高 6.6%,仅比最快基线低 1.5%;Muon 因 Newton–Schulz 低 32% |
| 路由负载均衡损失(地板 0.05) | Balance Loss | 0.0505(地板 1% 以内) | AdamW 0.0502;Muon 0.0608(末期跳升 22%);Lion 0.0537 | SkewAdam 像全状态 AdamW 一样稳定接近均匀路由地板 |
局限与改进
作者承认的局限:模型仅两层深,刻意集中 95% 参数于单一专家库做压力测试,但多层 MoE 路由逐层复合时分层如何表现未验证;多数配置单种子(虽有部分多种子复现);82M token 时长和 128 token 上下文都小;权重衰减在所有运行中实际失效(变化量低于 bfloat16 ULP),规模化时需把衰减项融入 float32 更新再做随机舍入转换且此配置未长程验证;Lion/Muon 用单一未调学习率,SkewAdam 本身也未调参,AdamW 未探到 $10^{-4}$ 以下;tier ablation 每变体一种子,0.6 PPL 分布在噪声内无法判定变体是否可分离;下游 zero-shot 评测(PIQA/HellaSwag 等)82M token 下接近随机仅作完整性检查。我额外观察:所有结论建立在浅层小 token MoE 上,向更大模型外推时专家动量可能在更长程开始体现价值;路由器精确二阶矩被 ablation 证明非承重,说明 tier 某些细节可简化;权重衰减失效使结论实际是无正则化场景,与生产训练存在偏差。
独立分析的弱点
独立分析的弱点及改进方向:第一,模型深度过浅(仅 2 block),把 95% 参数集中于单一专家库是想要的压力测试,但掩盖了多层 MoE 路由逐层复合时分层策略的行为,应在 12–32 层多 MoE 层模型上验证,考察跨层专家梯度稀疏度差异是否需要更细 tier。第二,单种子为主,Table 3 的 ablation 四变体 PPL 差仅 0.6 落在噪声内,无法断定各 tier 决策是否可分离,应至少 3 种子并报置信区间。第三,专家动量被判死重基于 82M token 短程,更长训练(数百亿 token)专家动量可能翻身,应用 10×–100× 预算复跑。第四,权重衰减失效使实验处于无正则化状态与生产训练不符,应把衰减项融入 float32 更新再抖动舍入验证长程稳定性。第五,路由器精确二阶矩被证明非承重,可进一步因子化路由器降低实现复杂度。
未来方向
作者明确提出的方向:在更大规模上验证方案,待算力可用时推广到生产级时长;解决权重衰减在 bfloat16 抖动舍入下的融合写回配置。基于成果可延伸的方向:其一,把分层思想与量化优化器状态(8-bit)或 ZeRO 跨卡分片正交组合——论文指出因式分解与量化可叠加且能越过 8-bit 内核 $2^{31}$ 元素硬退出边界,对超大 MoE 有工程价值;其二,将 tier 划分从静态规则推广为依据实时梯度稀疏度/方差的动态分配,让状态预算自适应训练过程;其三,把分层分析框架推广到其他非同质架构(长短注意力混合、多模态模型中不同模态参数群体),验证状态该放哪里的设计原则普适性;其四,结合本文关于路由器需要精确二阶矩的发现,研究与负载均衡损失稳定性直接挂钩的路由器专用优化器设计。
复现评估
复现评估较好但仍需投入。论文公开关键超参($\beta_1=0.9,\beta_2=0.999,\epsilon=10^{-8}$,cosine 学习率 3% warmup,AdamW/SkewAdam $3\times10^{-4}$,Lion $10^{-4}$,Muon 0.02 矩阵+内部 Adam $10^{-3}$),数据为公开 OpenWebText(按文档哈希 95/5 切分),模型架构(2 block、宽 4096、GQA 32 query/8 KV 头、50304 词表、128 专家 top-2)描述清晰,给出闭式内存模型可解析复算。作者称脚本日志在仓库 experiments/ 目录(含 8-bit 边界测量),对工程细节复现有帮助。困难在于:算力门槛高(主实验单 H200,补充跨 H100/MI300X 三平台),多数配置单种子难严格复现噪声,GaLore 基线为作者自实现且声明不应视为 GaLore 判决,Adafactor 退火 $\beta_2$ 调度($1-t^{-0.8}$)需特别对齐才能复现 149.5。方法描述充分,但完整复现需多卡多平台投入。
论文图表