优化器状态该放在哪?面向内存高效 MoE 训练的分层状态分配

  • 关联论文:2607.19058
  • 作者:Tom
  • 更新:2026-08-16

一句话结论

SkewAdam 通过识别 MoE 模型中稠密主干 / Expert / 路由器三类参数群体的梯度统计差异,为每类分配不同的优化器状态形式,将峰值训练内存从 81.4 GB 压缩至 31.3 GB(减少 61%),同时在相同 token 预算下取得更低的验证困惑度。


解决什么真问题

MoE(Mixture-of-Experts)训练长期以来面临一个被忽视的内存瓶颈:优化器状态是内存预算中最大的单项支出。以一个 6.78B 参数的 MoE 语言模型为例,AdamW 需要为 12.6 GB 的 bfloat16 权重维持 50.6 GB 的一阶和二阶动量——优化器状态是模型权重的 4 倍。

这一问题的根源在于:MoE 由三类参数群体构成——稠密主干(dense backbone)、Expert 和路由器(router)——它们在规模梯度统计上差异巨大,但在标准 AdamW 中被一视同仁地分配相同的优化器状态形式。作者的核心观察是:这种"一刀切"的状态分配策略,对不同群体而言既浪费又低效。

具体来说: - 路由器参数量极小(< 0.01%),但梯度信号高度动态,需要精确的二阶矩估计 - Expert 占参数量的 95%,但梯度方差相对较大,过精确的二阶矩反而浪费 - 稠密主干仅占 5% 的参数,却承载了大部分模型能力,需要最精确的优化状态

核心方法

SkewAdam 的设计哲学是:让优化器状态的形式与对应参数群体的梯度特性相匹配

分层状态分配策略

参数群体 参数占比 状态形式 理由
稠密主干 ~5% float32 动量 + 分解二阶矩 承载核心能力,需要高精度;分解形式省内存
Expert ~95% 分解二阶矩(无动量) 梯度方差大,动量收益有限;分解形式大幅省内存
路由器 <0.01% 精确二阶矩(无分解) 参数量极小;需要精确更新以维持负载均衡

分解二阶矩(Factored Second Moment)

传统 Adam 对每个参数维护完整二阶矩向量 v ∈ ℝ^(d),空间复杂度 O(d)。SkewAdam 对 Expert 采用分解形式:

v = reshape(γ, shape) * reshape(γ, shape).T

即维护两个秩-1 向量的外积近似二阶矩,将空间复杂度从 O(d) 降至 O(2√d)。这一近似对 Expert 群体效果良好,因为其梯度分布在秩-1 结构上有较好的低秩性。

状态体积对比(6.78B MoE)

优化器 状态总量 相对 AdamW
AdamW 50.6 GB 100%
SkewAdam 1.29 GB 2.6%

峰值训练内存从 81.4 GB 降至 31.3 GB,恰好落入 40 GB 加速器(如 A100 40GB)的内存预算。

关键伪代码逻辑

# 伪代码示意(本文验证环境无可用 PyTorch;实际使用请参考 GitHub nuemaan/skewadam)
def skewadam_update(param, grad, tier, state):
    if tier == "backbone":
        # 高精度:float32 动量 + 分解二阶矩
        state.m = beta1 * state.m + (1 - beta1) * grad
        v_hat = factored_estimate(grad)  # 分解近似
        param -= lr * state.m / (sqrt(v_hat) + eps)
    elif tier == "expert":
        # 中精度:仅分解二阶矩,无动量
        v_hat = factored_estimate(grad)
        param -= lr * grad / (sqrt(v_hat) + eps)
    else:  # router
        # 最高精度:精确二阶矩
        v_hat = beta2 * state.v + (1 - beta2) * grad ** 2
        param -= lr * grad / (sqrt(v_hat) + eps)

关键实验与数据

训练效率对比(82M tokens,相同初始化)

优化器 验证困惑度 相对基线
SkewAdam 108.4 −14.5%(相对 AdamW 126.8)
AdamW(untuned,原文直接对比值) 126.8 baseline
AdamW(tuned,Abstract 原文) 118.5 +7.6%(相对 SkewAdam)
Muon 120.2 +9.3%
Adafactor(tuned) 139.7 +22.5%
Lion 393.7 +232.2%

⚠️ 关键澄清:原 Abstract 中"SkewAdam 108.4,ahead of AdamW(126.8)"指的是未调参 AdamW的直接对比。同一 Abstract 后续消融部分提到的"tuned AdamW 118.5"是经过超参调优后的结果,出现在表格中。原文未明确说明两个 AdamW 实验是否在同一初始化下进行,此处两行并列仅供参照,解读时需注意这一未澄清的差异。

SkewAdam 在相同 token 预算下,困惑度显著低于所有调优后的基线。

Ablation 研究:谁在贡献性能?

消融实验揭示了一个反直觉的结论:分层设计本身并不提升模型精度。当把 SkewAdam 的 tier 分配打乱(backbone 用 Expert 的状态形式等),在携带 20 倍状态量的情况下,仍能达到相同的困惑度。这说明:

  • tier 机制的作用是节省内存,而非提升精度
  • 真正贡献精度的是 momentum(去除动量损失 31 个困惑度点)和 分解二阶矩的更新裁剪(损失 10 个点)

⚠️ 原文 Abstract 原句:"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)"。此处 31 和 10 均为 perplexity 差值;139.7 和 118.5 为对应 baseline 困惑度。

路由器负载均衡

SkewAdam 稳定将路由器负载均衡控制在均匀分布的 1% 以内,这对 MoE 的 Expert 并行稳定性至关重要。⚠️ 原文未给出不同 MoE 规模下的负载均衡数据表,仅有笼统的"within 1% of its uniform floor"表述。

亮点与局限

亮点

  1. 问题定位精准:首次系统性地量化了 MoE 训练中优化器状态的内存占比(4× 权重体积),提出了切实的工程痛点
  2. 跨优化器公平对比:在同一初始化、相同 token 预算下进行控制变量比较,实验设计严谨
  3. 工程可落地:峰值内存 31.3 GB 可用单卡 A100 40GB 训练,打破了"大 MoE 必须多卡"的前提假设
  4. 开源代码:GitHub nuemaan/skewadam 已公开,包含每轮训练日志

局限

  1. scale-up 未验证:实验仅在 6.78B 参数规模进行,更大模型(如 100B+)的梯度统计是否仍满足 tier 假设,原文未讨论
  2. Expert 类型依赖:分解二阶矩的低秩假设可能对不同 MoE 架构(SwiGLU、GLU 等)有效性不同,⚠️ 原文未做架构多样性实验
  3. 学习率敏感性:与标准 AdamW 相比,SkewAdam 的学习率超参数可能需要重新调优,增加了调参成本
  4. 非 MoE 模型适用性:对纯稠密模型,SkewAdam 的 tier 优势可能消失——因为只剩下一类参数群体

对工程落地的启发

对于正在训练或部署 MoE 的团队,SkewAdam 提供了一个几乎免费的内存优化

  • 单卡训练可行性:31.3 GB 峰值内存意味着可以在单卡 A100 40GB 上训练原本需要多卡的 MoE,大幅降低训练门槛
  • 长上下文场景直接受益:更低的内存占用 = 更多显存可用于 KV Cache,支持更长上下文
  • 与现有优化正交:可与 FP8 训练、梯度检查点、ZeRO 等技术叠加使用
  • 实现成本低:只需在参数初始化时按群体分类,分配不同的优化器状态,无需改动模型架构

⚠️ 注意:当前实现需要模型层面显式分离三类参数群体(backbone/expert/router),对黑盒模型或第三方预训练模型的适配成本未知。

与同方向工作的关系

SkewAdam 处于大模型训练优化这一活跃方向的核心节点:

  • 与 Adam-mini 的关系:v2 版本新增了对 Adam-mini 的讨论——后者通过类似原理(参数群体差异化状态)为纯稠密模型节省优化器内存。SkewAdam 将这一思想扩展到 MoE 的多群体场景
  • 与 QLoRA / GaLore 的关系:这些工作关注参数高效微调中的内存优化;SkewAdam 面向预训练阶段,层级不同但目标一致
  • 与 MoE 路由稳定性的关系:Soto et al. 等人研究了 MoE 负载均衡问题;SkewAdam 的路由器精确二阶矩设计为这一分支提供了新的优化器视角

适合谁读

  • LLM 训练工程师:正在管理 MoE 训练集群或优化显存占用的团队,直接受益
  • MoE 研究者:关注路由器动态、Expert 并行、或内存优化的学者
  • 推理系统工程师:对"训练时优化器状态如何影响推理显存"感兴趣的从业者
  • 大模型infra从业者:评估不同优化器在生产环境中的实际性价比

工程落地与核查(Jay)

实际系统怎么用

DATABASE 层:无直接数据库依赖

本文方法属于训练优化器算法,不涉及持久化存储。

BACKEND 层:PyTorch 自定义优化器

SkewAdam 本质是一个自定义 PyTorch 优化器模块,核心依赖仅为 PyTorch(无需额外 C++ 扩展)。

⚠️ PyTorch 环境未在本文验证环境实测;GitHub nuemaan/skewadam 有完整训练日志可参考。集成步骤:

# 伪代码示意(不可直接运行,依赖 PyTorch)
# 实际使用请参考 https://github.com/nuemaan/skewadam

# import torch
# from skewadam import SkewAdam

# 参数分组
backbone_params = [p for n, p in model.named_parameters() if "backbone" in n]
expert_params   = [p for n, p in model.named_parameters() if "expert" in n]
router_params   = [p for n, p in model.named_parameters() if "router" in n]

optimizer = SkewAdam([
    {"params": backbone_params, "tier": "backbone"},
    {"params": expert_params,   "tier": "expert"},
    {"params": router_params,    "tier": "router"},
], lr=1e-4)

CLOUD-NATIVE 层:单卡 A100 40GB 原生支持

峰值 31.3 GB 可直接运行于单卡 A100 40GB,无需多卡并行。对于多卡场景,建议配合 ZeRO Stage 2/3 使用——SkewAdam 的 tier 分组与 ZeRO 的分片策略可正交叠加: - Expert 参数(95%)经 SkewAdam 状态压缩后,ZeRO 分片负担大幅降低 - Backbone 参数(5%)的 float32 动量仍需一定显存,但总量可控

⚠️ FP8/ZeRO/梯度检查点叠加:三者与 SkewAdam 均无机制冲突,但建议按以下顺序验证: 1. 先单独集成 SkewAdam,确认收敛性 2. 再叠加梯度检查点(增加前向计算,换后向显存) 3. 再叠加 ZeRO(多卡分片) 4. FP8 作为最外层选项(需硬件支持 H100+)

CSDN 层:中文社区注意事项

  • 中文博客转述时,务必注明"Abstract 中的 AdamW 基线有两个数值(126.8 未调参 vs 118.5 调参后),容易混淆",否则会引发读者质疑
  • "分解二阶矩"(Factored Second Moment)国内常译为"低秩分解二阶矩"或"因子化二阶矩",需在正文中统一术语

REPRODUCTION 层:复现要点

# 克隆仓库
git clone https://github.com/nuemaan/skewadam.git
cd skewadam

# 依赖(纯 PyTorch,无 CUDA 扩展)
pip install torch transformers

# 训练日志已公开,可直接对照 perplexity 曲线
# 复现关键点:
# 1. 参数 tier 分配必须在初始化时完成,中途不可更改
# 2. 学习率建议从 1e-4 起(原文未提供详细 LR sweep 范围)
# 3. 82M tokens 实验约需单卡 A100 8-12 小时(估算)

坑在哪

坑点 描述 建议
tier 分配黑盒成本 现有预训练模型(如 Mixtral、FalconMOE)的 tier 结构未知,需手动逆向参数命名或从源码确认 ⚠️ 使用前务必确认参数分组正确,错误的 tier 分配(如把 expert 放进 backbone)会导致精度下降
低秩假设有效性未知 分解二阶矩对 SwiGLU、GLU 等复杂 Expert 结构是否仍有效,⚠️ 原文未验证 小规模消融验证后再上生产规模
学习率重调优 Backbone 组的 float32 动量引入额外超参组合空间 ⚠️ 建议使用 Learning Rate Finder 或 PreScalar 等自适应工具辅助调优
float32 动量边际占用 Backbone 占参数 5%,float32 动量为 4 字节/参数,约 0.2 GB(6.78B 模型) 相比 Expert 的状态节省(从 ~50 GB → ~1 GB 量级),这部分开销可忽略,但大规模模型需重新估算
调试复杂度增加 三套不同的更新路径,出问题定位成本高 建议先用单 tier(全部 backbone 设置)做 baseline 验证

核查记录

以下数字均经 arXiv Abstract(2607.19058v2)原文核实:

核查项 来源 Abstract
模型参数规模 6.78B "On a 6.78B-parameter MoE language model"
AdamW 优化器状态体积 50.6 GB "AdamW keeps 50.6 GB of first and second moments"
模型权重体积 12.6 GB bfloat16 "to update 12.6 GB of bfloat16 weights"
Backbone / Expert / Router 比例 5% / 95% / <0.01% "backbone (5% of parameters), the experts (95%) and the router (<0.01%)"
SkewAdam 优化器状态 1.29 GB (2.6% of AdamW) "state occupies 1.29 GB or 2.6% of AdamW's"
峰值训练内存 81.4 GB → 31.3 GB "peak training memory falls from 81.4 GB to 31.3 GB"
SkewAdam 困惑度 108.4 "SkewAdam reaches validation perplexity 108.4"
AdamW (untuned) 困惑度 126.8 "ahead of AdamW (126.8)"
AdamW (tuned) 困惑度 118.5 "tuned AdamW, 118.5"
Muon 困惑度 120.2 "Muon (120.2)"
Adafactor (tuned) 困惑度 139.7 "tuned Adafactor, 139.7"
Lion 困惑度 393.7 "Lion (393.7)"
路由器负载均衡 within 1% of uniform "router load balance to within 1% of its uniform floor"
Tier 消融结论 20x 状态量达到相同困惑度 "A tier ablation reaches the same value while carrying twenty times the state"
Momentum 消融损失 31 perplexity points "Removing momentum costs 31 perplexity points (tuned Adafactor, 139.7)"
分解二阶矩消融损失 10 perplexity points "replacing the factored second moment and its update clipping with a full second moment costs 10 (tuned AdamW, 118.5)"
GitHub 链接 nuemaan/skewadam Abstract meta + "GitHub nuemaan/skewadam"

特别标注

  • AdamW 双数值问题:本文 explainer 已在实验表格中对 126.8(未调参)和 118.5(调参后)两行并列并加注释说明差异,避免读者困惑。此为原文 Abstract 自身存在的表述结构问题,非 explainer 杜撰。
  • PyTorch 未实测:因验证环境无 PyTorch,代码段标注为"伪代码示意",实际运行请参考 GitHub 仓库。