DeMo: Decoupled Momentum Optimization 精读文档
分布式训练通信压缩优化器的深度技术分析
DeMo: Decoupled Momentum Optimization 精读文档
一、论文概述
| 项目 | 内容 |
|---|---|
| 标题 | DeMo: Decoupled Momentum Optimization |
| 作者 | Bowen Peng, Lizhang Chen, Baiyu Su, Jeffrey Quesnelle, Diederik P. Kingma, Qiang Liu |
| 机构 | Krea AI, University of Texas at Austin |
| 论文 | arXiv:2411.19870 |
| 代码 | github.com/bloc97/DeMo |
| 发布 | 2024-11-29 (v1), 2026-02-06 (v2) |
| 许可 | CC BY 4.0 |
| 关键词 | 分布式训练, 通信压缩, 动量优化, 稀疏化, DCT |
二、问题背景与动机
2.1 分布式训练的通信瓶颈
大规模语言模型(LLM)训练依赖同步数据并行(Distributed Data Parallelism, DDP),每个 worker 持有完整的模型副本,独立计算梯度后通过 All-Reduce 操作同步。
核心问题:
- All-Reduce 通信量与模型参数量成正比
- 对于前沿模型(如 LLaMA-3 405B),每步通信量可达 TB 级别
- 要求高带宽互连(NVLink 900GB/s, InfiniBand 400Gbps)
- 地理共置集群成本高昂,限制了可扩展性
2.2 现有解决方案的局限性
| 方法 | 原理 | 局限性 |
|---|---|---|
| 梯度压缩(QSGD) | 标量量化 | 压缩比有限,需要特殊通信原语 |
| 低秩近似(PowerSGD) | 随机投影 | 随机性导致不稳定 |
| 稀疏化(Top-k) | 选择性传输 | 需要误差反馈机制 |
| 分层优化(DiLoCo) | 内/外层分离 | 需要多机协调,系统复杂 |
| 量化(1-bit Adam) | 二值化 | 需要特殊 All-Reduce 实现 |
2.3 DeMo 的核心洞察
关键观察:在标准 SGD with Momentum 中:
梯度同步和动量同步在理论上是等价的(线性性),但动量经过压缩后通信量可大幅降低。
三、技术架构详解
3.1 整体流程

DeMo 的压缩流水线包含三个阶段:
动量张量 M ∈ ℝ^{n₀×...×n_{d-1}}
↓
┌─────────────────────────────────┐
│ 1. 张量分块 (Tensor Chunking) │
│ 将 M 分割为大小为 s×s 的块 │
└─────────────────────────────────┘
↓
┌─────────────────────────────────┐
│ 2. 分块线性投影 (Blockwise LP) │
│ 对每个块应用 DCT 变换 │
│ Q_k = P₀ · B_k · P₁ᵀ │
└─────────────────────────────────┘
↓
┌─────────────────────────────────┐
│ 3. Top-k 稀疏化 │
│ 选择 k 个最大幅值系数 │
│ Q̄_k = TopK(Q_k, k) │
└─────────────────────────────────┘
↓
通信压缩后的稀疏表示
3.2 算法详解
步骤 1:解耦本地动量更新
标准 DDP 在每个 micro-batch 梯度计算后立即同步。DeMo 移除了 All-Reduce,让每个 worker 独立维护动量缓冲区:
其中:
- β ∈ (0,1) 是动量系数(默认 0.999)
- G_t^i 是 worker i 在时间步 t 的局部梯度
- M_t^i 是 worker i 的本地动量缓冲区
为什么可行? 由于线性性:
聚合动量等价于聚合梯度后的动量更新。
步骤 2:张量分块(Tensor Chunking)
给定动量张量 M ∈ ℝ^{n₀×…×n_{d-1}},将每个维度因式分解为 n_i = c_i · s_i:
其中每个块 B_k ∈ ℝ^{s₀×…×s_{d-1}}。
实际设置:
- 对于 Transformer 中的矩阵张量:s=64,即 64×64 块
- 对于向量张量(如 LayerNorm):s=64,即 64 维块
- 总块数 C = n₀/s₀ × n₁/s₁ × … × n_{d-1}/s_{d-1}
步骤 3:分块线性投影(Blockwise Linear Projection)
对每个块应用可分离多线性变换:
对于 2D 张量简化为双线性形式:
投影矩阵选择:
| 投影方式 | 公式 | 优点 | 缺点 |
|---|---|---|---|
| 恒等映射 | P_i = I | 计算简单 | 性能最差 |
| DCT | P_i = DCT(s_i) | FFT 快速实现,预计算一次 | 固定基底 |
| 随机正交投影 | P_i ~ N(0,I), 正交化 | 动态变化基底 | 需要每步生成 |
为什么 DCT 有效?
- 梯度/动量在频域具有能量集中特性
- 低频系数包含大部分信息
- DCT 将能量集中到少数系数,便于 top-k 选择
步骤 4:Top-k 稀疏化
在频域中选择 k 个最大幅值的系数:
压缩比计算:
- 块大小:s×s = 64×64 = 4096
- 保留系数:k(默认 k=8)
- 压缩比:4096/k ≈ 512 倍(per block)
步骤 5:动量重建与参数更新
每个 worker 从聚合的稀疏块重建动量张量:
由于 DCT 是正交变换,P_i^{-1} = P_i^T。
参数更新:
其中 ϕ(·) 是基础优化器的变换:
- SGD: ϕ(M) = M
- Signum: ϕ(M) = sign(M)(逐元素)
- Muon: ϕ(M) = M(M^TM + εI)^{-1/2}
步骤 6:动量减法(误差反馈)
这是 DeMo 的关键创新,用于解决稀疏化带来的信息丢失:
α ∈ (0,1] 是动量减法系数(默认 α=0.2)。
为什么需要动量减法?
- 没有减法(α=0):相同的 top-k 元素被重复选择和传输,导致退化
- 完全减法(α=1):过度衰减历史信息
- 部分减法(α=0.2):渐进演化 top-k 元素,部分衰减已通信值
与传统误差反馈的区别:
- 传统方法需要额外的误差累加器内存
- DeMo 重用动量缓冲区作为累加器,零额外内存开销
3.3 算法伪代码
Algorithm 4: DeMo: Decoupled Momentum Optimization
Input: 学习率 η, 动量系数 β, 权重衰减 λ, 稀疏预算 k, DCT 矩阵 {P_j}
Initialize: 参数 X₀, 全局动量 M₀ ← 0
for t = 1, 2, ... do
# 每个 worker i ∈ {0, ..., N-1}
G_t^i ← ∇L(X_{t-1}; ξ_t^i) # 计算局部梯度
M_t^i ← β·M_{t-1}^i + G_t^i # 更新本地动量
for each chunk M_t^{i,[ℓ]} of M_t^i do
Q_t^{i,[ℓ]} ← Top-k(DCT(M_t^{i,[ℓ]}; {P_j}), k) # 分块变换+稀疏化
M_t^{i,[ℓ]} ← M_t^{i,[ℓ]} - IDCT(Q_t^{i,[ℓ]}; {P_j}^T) # 原地残差更新
end for
send {Q_t^{i,[ℓ]}} to server # 发送稀疏化后的动量
# 参数服务器
for each chunk {Q_t^{i,[ℓ]}} do
M_t^{[ℓ]} ← IDCT(∑_i Q_t^{i,[ℓ]}; {P_j}^T) # 聚合并重建动量
end for
X_t ← X_{t-1} - η·(sgn(M_t) + λ·X_{t-1}) # 参数更新(以 Muon 为例)
end for
3.4 复杂度分析
| 操作 | 计算复杂度 | 内存开销 | 说明 |
|---|---|---|---|
| 无分块投影 | O(N³) | O(N²) | N 为张量总元素数 |
| 分块投影(C²块) | O(N³/C) | O(N²/C²) | C 为每维度分块数 |
| DCT 变换 | O(N log N) | O(N) | 使用 FFT 快速实现 |
| Top-k 选择 | O(N log k) | O(k) | 选择 k 个最大值 |
| 动量减法 | O(N) | O(1) | 原地操作 |
总开销:相比标准 DDP,额外计算开销可忽略不计。
3.5 通信量对比
| 方法 | 每步通信量(参数量 M) | 说明 |
|---|---|---|
| DDP (All-Reduce) | 2M × sizeof(dtype) | 全量梯度同步 |
| DeMo (k=8, s=64) | M/512 × sizeof(dtype) | 85x 压缩 |
| DeMo (k=16, s=64) | M/256 × sizeof(dtype) | 44x 压缩 |
| DeMo (k=32, s=64) | M/128 × sizeof(dtype) | 22x 压缩 |
四、理论分析
4.1 标准假设
假设 1(方差有界):
随机梯度估计满足:
假设 2(L-Smoothness):
目标函数 L-Smooth:
等价地:
假设 3(梯度有界):
4.2 收敛定理
定理 1(DeMo 收敛性):
在假设 1、2、3 下,DeMo 算法生成的序列 {X_t}_{t=1}^T 满足:
其中 D 是参数空间直径,N 是 worker 数量。
收敛率:O(1/√T) 的平均梯度范数收敛率。
4.3 误差分析
引理 2(Top-k 近似误差):
稀疏化引入的近似误差:
其中 k 是保留的系数数量,s² 是块大小。
引理 8.1(偏差分析):
稀疏化后的梯度估计偏差:
其中 ,ε 与 k/s² 成正比。
五、实验详解
5.1 实验设置
| 配置 | OLMo-300M | OLMo-1B |
|---|---|---|
| 非嵌入参数 | 320M | 1.18B |
| 层数 | 24 | 24 |
| 隐藏维度 | 1024 | 2048 |
| 注意力头数 | 16 | 16 |
| 序列长度 | 2048 | 2048 |
| 训练 tokens | 100B | 100B |
| GPU | 64 × NVIDIA H100 | 64 × NVIDIA H100 |
| 全局 batch size | 2048 | 2048 |
| 梯度累积步数 | 4 | 4 |
| 每 GPU batch size | 8 | 8 |
超参数设置:
- 学习率:线性 warmup + 余弦衰减
- AdamW β₁=0.9, λ=0.1
- DeMo 默认:β=0.999, α=0.2, s=64
5.2 主要结果
零样本下游任务评估

| 优化器 / k | HellaSwag ↑ | ARC-Easy ↑ | PIQA ↑ | 通信量 (MB/step) |
|---|---|---|---|---|
| OLMo-300M | ||||
| DeMo k=32 | 0.37 | 0.46 | 0.67 | 29.9 |
| DeMo k=16 | 0.38 | 0.50 | 0.67 | 14.9 |
| DeMo k=8 | 0.38 | 0.47 | 0.67 | 7.49 |
| DeMo k=4 | 0.37 | 0.47 | 0.67 | 3.74 |
| DeMo k=2 | 0.36 | 0.46 | 0.65 | 1.87 |
| DeMo k=1 | 0.35 | 0.45 | 0.65 | 0.93 |
| AdamW-DDP | 0.35 | 0.46 | 0.65 | 636.9 |
| OLMo-1B | ||||
| DeMo k=32 | 0.48 | 0.55 | 0.70 | 110.32 |
| DeMo k=16 | 0.47 | 0.53 | 0.70 | 55.16 |
| DeMo k=8 | 0.47 | 0.52 | 0.69 | 27.58 |
| DeMo k=4 | 0.45 | 0.52 | 0.70 | 13.79 |
| DeMo k=2 | 0.44 | 0.51 | 0.69 | 6.89 |
| DeMo k=1 | 0.41 | 0.52 | 0.69 | 3.44 |
| AdamW-DDP | 0.43 | 0.51 | 0.68 | 2416.6 |
关键发现:
- 300M 模型:DeMo k=8 将通信量降低 85 倍(7.5 MB vs 637 MB),精度无损
- 1B 模型:DeMo k=16 降低 44 倍(55 MB vs 2417 MB),HellaSwag 和 PIQA 还优于 AdamW
- k=2 即可超越 AdamW:即使极端压缩也能保持竞争力
Perplexity vs 通信量

- DeMo 在 Pareto 前沿上始终优于 AdamW
- 相同通信预算下,DeMo 提供更好的模型质量
- 相同模型质量下,DeMo 需要更少的通信
5.3 消融实验详解
5.3.1 线性变换选择

| 变换方法 | k=8 | k=16 | k=32 | 说明 |
|---|---|---|---|---|
| 恒等映射 | 最差 | 最差 | 最差 | 无变换,直接稀疏化 |
| DCT | 最优 | 最优 | 最优 | 固定基底,FFT 快速实现 |
| 随机正交投影 | 接近 DCT | 接近 DCT | 接近 DCT | 动态变化,计算开销大 |
结论:
- DCT 优于恒等映射:支持”参数更新为多个非稀疏向量的线性组合”
- DCT ≈ 随机投影:但 DCT 计算效率更高(FFT vs 矩阵乘法)
- 低 k 值时差距更明显:DCT 能更好地保留关键信息
5.3.2 动量减法系数 α

| α 值 | 训练损失 | 说明 |
|---|---|---|
| 0.0 | 显著退化 | 相同 top-k 元素被重复选择 |
| 0.1 | 较好 | 渐进演化 |
| 0.2 | 最优 | 平衡新旧信息 |
| 0.5 | 次优 | 过度衰减历史信息 |
| 1.0 | 较差 | 完全移除已通信信息 |
关键洞察:
- α=0 时,由于固定基底,top-k 元素演化缓慢
- 相同元素被重复选择,导致更新冗余
- α=0.2 通过部分衰减,让 top-k 元素逐渐演化
5.3.3 动量系数 β
| β 值 | 训练损失 | 说明 |
|---|---|---|
| 0.95 | 较差 | 动量累积不足 |
| 0.98 | 一般 | |
| 0.99 | 一般 | |
| 0.995 | 最优 | 最佳平衡点 |
| 0.999 | 接近最优 | 默认设置 |
结论:β=0.999 在动量减法开启时显著优于标准值 0.9。
5.3.4 分块大小 s
| s 值 | 压缩比 | 性能 | 说明 |
|---|---|---|---|
| 32 | 1024/k | 略差 | 块太小,DCT 效果有限 |
| 64 | 4096/k | 最优 | 默认设置 |
| 128 | 16384/k | 略差 | 块太大,边界效应 |
5.4 与其他优化器对比

| 优化器 | 训练损失 | 通信量 | 系统复杂度 |
|---|---|---|---|
| AdamW-DDP | 基线 | 基线 (1x) | 低 |
| Muon-DDP | 更优 | 基线 (1x) | 低 |
| DeMo | 接近 AdamW | 85x 压缩 | 低 |
| DiLoCo | 更差 | 可变 | 高(多机协调) |
| PowerSGD | 更差 | 10-50x | 中 |

关键发现:
- DeMo 在相同压缩比下始终优于 DiLoCo
- 高压缩比下差距更明显(DeMo 更稳定)
- DiLoCo 需要复杂的多机协调,DeMo 更简单
5.5 计算效率分析
每 GPU 通信量对比:
| 模型规模 | AdamW-DDP | DeMo k=8 | DeMo k=16 | 压缩比 |
|---|---|---|---|---|
| 300M | 637 MB | 7.5 MB | 14.9 MB | 85x / 44x |
| 1B | 2417 MB | 27.6 MB | 55.2 MB | 88x / 44x |
训练时间估算(假设 100Gbps 网络):
| 场景 | AdamW-DDP | DeMo k=8 | 加速比 |
|---|---|---|---|
| 300M, 64 GPU | ~5.1s/step | ~0.06s/step | 85x |
| 1B, 64 GPU | ~19.3s/step | ~0.22s/step | 88x |
六、代码实现分析
6.1 项目结构
DeMo/
├── demo/
│ ├── __init__.py
│ ├── demo_optimizer.py # DeMo 优化器核心实现
│ ├── demo_utils.py # DCT/IDCT 工具函数
│ └── configs/ # 配置文件
├── examples/
│ ├── olmo_300m.yaml # OLMo-300M 配置
│ └── olmo_1b.yaml # OLMo-1B 配置
├── scripts/
│ └── train.py # 训练脚本
└── README.md
6.2 核心实现
class DeMo(torch.optim.Optimizer):
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999),
weight_decay=0.01, k=8, chunk_size=64, alpha=0.2):
defaults = dict(lr=lr, betas=betas, weight_decay=weight_decay,
k=k, chunk_size=chunk_size, alpha=alpha)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
lr = group['lr']
beta = group['betas'][0]
weight_decay = group['weight_decay']
k = group['k']
chunk_size = group['chunk_size']
alpha = group['alpha']
for p in group['params']:
if p.grad is None:
continue
grad = p.grad
state = self.state[p]
# 初始化动量缓冲区
if 'momentum_buffer' not in state:
state['momentum_buffer'] = torch.zeros_like(p)
buf = state['momentum_buffer']
# 更新动量
buf.mul_(beta).add_(grad, alpha=1-beta)
# 分块 + DCT + Top-k
chunks = self._chunk_tensor(buf, chunk_size)
sparse_chunks = []
for chunk in chunks:
q = self._dct(chunk)
topk_vals, topk_idx = torch.topk(q.flatten(), k)
q_sparse = torch.zeros_like(q).flatten()
q_sparse[topk_idx] = topk_vals
q_sparse = q_sparse.view(q.shape)
sparse_chunks.append(q_sparse)
# 动量减法(误差反馈)
chunk_residual = self._idct(q_sparse)
chunk.sub_(alpha * chunk_residual)
# 通信(All-Gather 稀疏化后的动量)
aggregated = self._all_gather(sparse_chunks)
# 重建动量
buf_reconstructed = self._unchunk(aggregated)
# 参数更新
p.mul_(1 - lr * weight_decay)
p.add_(buf_reconstructed, alpha=-lr)
return loss
def _dct(self, x):
"""离散余弦变换"""
# 使用 torch.fft 实现快速 DCT
...
def _idct(self, x):
"""逆离散余弦变换"""
...
6.3 集成方式
-
禁用 DDP 梯度同步:
model = DDP(model, find_unused_parameters=False) # 禁用 All-Reduce model.require_forward_param_sync = False model.require_backward_grad_sync = False -
使用 DeMo 优化器:
optimizer = DeMo(model.parameters(), lr=1e-3, k=8, chunk_size=64) -
代码修改量:约 200 行核心实现,无需修改模型代码
七、相关工作对比
7.1 梯度压缩方法
| 方法 | 压缩方式 | 压缩比 | 额外内存 | 通信原语 |
|---|---|---|---|---|
| QSGD | 标量量化 | 4-32x | O(1) | All-Reduce |
| PowerSGD | 低秩近似 | 10-50x | O(rank) | All-Reduce |
| Top-k | 稀疏化 | 10-100x | O(1) | All-Reduce |
| 1-bit Adam | 二值化 | 32x | O(1) | 1-bit All-Reduce |
| DeMo | DCT + Top-k | 85x | O(1) | All-Gather |
7.2 分层优化方法
| 方法 | 原理 | 通信模式 | 系统复杂度 |
|---|---|---|---|
| DiLoCo | 内/外层分离 | 周期性同步 | 高(多机协调) |
| FedAvg | 聚合本地更新 | 周期性同步 | 中 |
| DeMo | 动量压缩同步 | 每步压缩同步 | 低 |
7.3 关键区别
| 特性 | DeMo | DiLoCo | PowerSGD |
|---|---|---|---|
| 通信频率 | 每步 | 每 H 步 | 每步 |
| 额外内存 | 0 | O(model) | O(rank) |
| 系统复杂度 | 低 | 高 | 中 |
| 拓扑无关性 | 是 | 部分 | 是 |
| 可与 Muon 结合 | 是 | 是 | 是 |
八、局限性与未来方向
8.1 当前局限性
- 超参数敏感:k、α、β、s 等参数需要针对不同模型调优
- 验证规模有限:仅在 300M 和 1B 参数模型上验证,未扩展到 10B+
- 固定稀疏模式:top-k 可能不是最优策略,自适应稀疏化可能更好
- 理论分析基础:收敛性分析基于标准假设,实际大规模训练可能有偏差
- 单机验证:实验在单节点 64 GPU 上进行,未验证跨节点场景
8.2 未来研究方向
- 大规模验证:在 10B+ 参数模型和跨数据中心场景中验证
- 自适应稀疏化:根据训练阶段动态调整 k 值
- 混合精度支持:结合 FP8/FP4 量化进一步压缩
- 异步训练:与异步 DDP 结合,进一步降低同步开销
- 硬件加速:设计专用硬件加速 DCT 和 Top-k 操作
- 理论深化:非凸优化下的收敛率分析
九、总结
核心贡献
- 提出 DeMo 优化器:通过解耦动量更新和结构化压缩,将分布式训练通信量降低 1-2 个数量级
- 结构化压缩流水线:分块 + DCT + Top-k 组合,利用梯度频域稀疏性实现高效压缩
- 动量减法误差反馈:创新性地重用动量缓冲区作为误差累加器,零额外内存开销
- 拓扑无关性:不依赖特定网络拓扑,支持跨数据中心和以太网训练
- 即插即用:可与任何基于动量的优化器结合,最小化代码修改
技术影响
- 降低训练成本:使低带宽互连(以太网)也能进行大规模 LLM 训练
- 扩展训练基础设施:支持跨数据中心的分布式训练
- 简化系统设计:无需复杂的通信拓扑优化
- 开源实现:提供完整代码和配置,便于复现和集成
关键数值总结
| 指标 | 数值 |
|---|---|
| 最大压缩比 | 85x (300M), 88x (1B) |
| 默认参数 | k=8, s=64, α=0.2, β=0.999 |
| 训练框架 | OLMo |
| GPU 规模 | 64 × H100 |
| 训练量 | 100B tokens |
| 代码量 | ~200 行 |
十、参考资源
- 论文: arXiv:2411.19870
- 代码: github.com/bloc97/DeMo
- 实验复现: anonymous.4open.science/r/DeMo-D3F1
- 相关工作:
- PowerSGD (Vogels et al., 2019) - 低秩梯度压缩
- DiLoCo (Douillard et al., 2023) - 分层分布式优化
- Muon (Jordan et al., 2024) - 动量优化器
- Deep Gradient Compression (Lin et al., 2018) - 梯度稀疏化
- 1-bit Adam (Tang et al., 2021) - 量化通信
- QSGD (Alistarh et al., 2017) - 标量量化
- OLMo (Groeneveld et al., 2024) - 开源语言模型框架