Back to blog

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 中:

mt=βmt−1+(1−β)gt\bm{m}_t = \beta \bm{m}_{t-1} + (1-\beta) \bm{g}_t

xt+1=xt−ηmt\bm{x}_{t+1} = \bm{x}_t - \eta \bm{m}_t

梯度同步和动量同步在理论上是等价的(线性性),但动量经过压缩后通信量可大幅降低。

三、技术架构详解

3.1 整体流程

DCT与Top-k稀疏化流程

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 独立维护动量缓冲区:

Mti=βMt−1i+(1−β)Gti\bm{M}_t^i = \beta \bm{M}_{t-1}^i + (1-\beta) \bm{G}_t^i

其中:

  • β ∈ (0,1) 是动量系数(默认 0.999)
  • G_t^i 是 worker i 在时间步 t 的局部梯度
  • M_t^i 是 worker i 的本地动量缓冲区

为什么可行? 由于线性性: 1N∑i=1NMti=β⋅1N∑i=1NMt−1i+(1−β)⋅1N∑i=1NGti\frac{1}{N}\sum_{i=1}^N \bm{M}_t^i = \beta \cdot \frac{1}{N}\sum_{i=1}^N \bm{M}_{t-1}^i + (1-\beta) \cdot \frac{1}{N}\sum_{i=1}^N \bm{G}_t^i

聚合动量等价于聚合梯度后的动量更新。

步骤 2:张量分块(Tensor Chunking)

给定动量张量 M ∈ ℝ^{n₀×…×n_{d-1}},将每个维度因式分解为 n_i = c_i · s_i:

B(M)={Bk∣k∈[c0]×⋯×[cd−1]}\mathcal{B}(\bm{M}) = \{ \bm{B}_\mathbf{k} \mid \mathbf{k} \in [c_0] \times \dots \times [c_{d-1}] \}

其中每个块 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)

对每个块应用可分离多线性变换:

Qk=T(Bk;P0,…,Pd−1),Pi∈Rsi×si\bm{Q}_\mathbf{k} = \mathcal{T}(\bm{B}_\mathbf{k}; \bm{P}_0, \dots, \bm{P}_{d-1}), \quad \bm{P}_i \in \mathbb{R}^{s_i \times s_i}

对于 2D 张量简化为双线性形式: Qk=P0BkP1⊤\bm{Q}_\mathbf{k} = \bm{P}_0 \bm{B}_\mathbf{k} \bm{P}_1^\top

投影矩阵选择:

投影方式公式优点缺点
恒等映射P_i = I计算简单性能最差
DCTP_i = DCT(s_i)FFT 快速实现,预计算一次固定基底
随机正交投影P_i ~ N(0,I), 正交化动态变化基底需要每步生成

为什么 DCT 有效?

  • 梯度/动量在频域具有能量集中特性
  • 低频系数包含大部分信息
  • DCT 将能量集中到少数系数,便于 top-k 选择

步骤 4:Top-k 稀疏化

在频域中选择 k 个最大幅值的系数:

Qˉk=TopK(Qk,k)\bar{\bm{Q}}_\mathbf{k} = \text{TopK}(\bm{Q}_\mathbf{k}, k)

压缩比计算:

  • 块大小:s×s = 64×64 = 4096
  • 保留系数:k(默认 k=8)
  • 压缩比:4096/k ≈ 512 倍(per block)

步骤 5:动量重建与参数更新

每个 worker 从聚合的稀疏块重建动量张量:

Mt∗=B−1{T−1(Qˉk;P0−1,…,Pd−1−1)}\bm{M}_t^* = \mathcal{B}^{-1} \left\{ \mathcal{T}^{-1} \left( \bar{\bm{Q}}_\mathbf{k}; \bm{P}_0^{-1}, \dots, \bm{P}_{d-1}^{-1} \right) \right\}

由于 DCT 是正交变换,P_i^{-1} = P_i^T。

参数更新: Xt+1=Xt−ηt(ϕ(Mt∗)+λXt)\bm{X}_{t+1} = \bm{X}_t - \eta_t (\phi(\bm{M}_t^*) + \lambda \bm{X}_t)

其中 ϕ(·) 是基础优化器的变换:

  • SGD: ϕ(M) = M
  • Signum: ϕ(M) = sign(M)(逐元素)
  • Muon: ϕ(M) = M(M^TM + εI)^{-1/2}

步骤 6:动量减法(误差反馈)

这是 DeMo 的关键创新,用于解决稀疏化带来的信息丢失:

Mti←Mti−α B−1{T−1(Qk;P0−1,…,Pd−1−1)}\bm{M}_t^i \leftarrow \bm{M}_t^i - \alpha \, \mathcal{B}^{-1} \left\{ \mathcal{T}^{-1}(\bm{Q}_\mathbf{k}; \bm{P}_0^{-1}, \dots, \bm{P}_{d-1}^{-1}) \right\}

α ∈ (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(方差有界):

随机梯度估计满足: E[∥Gi(X)−∇L(X)∥F2]≤σ2nbatch\mathbb{E}\left[\left\lVert G^i(\bm{X}) - \nabla\mathcal{L}(\bm{X})\right\rVert_F^2\right] \leq \frac{\sigma^2}{n_{\text{batch}}}

假设 2(L-Smoothness):

目标函数 L-Smooth: ∥∇L(Y)−∇L(X)∥F≤L∥Y−X∥F\left\lVert\nabla\mathcal{L}(\bm{Y}) - \nabla\mathcal{L}(\bm{X})\right\rVert_F \leq L\left\lVert\bm{Y} - \bm{X}\right\rVert_F

等价地: L(Y)≤L(X)+⟨∇L(X),Y−X⟩+L2∥Y−X∥F2\mathcal{L}(\bm{Y}) \leq \mathcal{L}(\bm{X}) + \langle\nabla\mathcal{L}(\bm{X}), \bm{Y}-\bm{X}\rangle + \frac{L}{2}\left\lVert\bm{Y}-\bm{X}\right\rVert_F^2

假设 3(梯度有界): ∥∇L(X;ξ)∥1≤R\|\nabla\mathcal{L}(\bm{X};\xi)\|_1 \leq R

4.2 收敛定理

定理 1(DeMo 收敛性):

在假设 1、2、3 下,DeMo 算法生成的序列 {X_t}_{t=1}^T 满足:

1T∑t=1TE[∥∇L(Xt)∥1]≤E[L(X0)−L(XT)]Tη+2LDη1+σNnbatch\frac{1}{T}\sum_{t=1}^T \mathbb{E}\left[\|\nabla\mathcal{L}(\bm{X}_t)\|_1\right] \leq \frac{\mathbb{E}[\mathcal{L}(\bm{X}_0) - \mathcal{L}(\bm{X}_T)]}{T\eta} + \frac{2LD\eta}{1} + \frac{\sigma}{Nn_{\text{batch}}}

其中 D 是参数空间直径,N 是 worker 数量。

收敛率:O(1/√T) 的平均梯度范数收敛率。

4.3 误差分析

引理 2(Top-k 近似误差):

稀疏化引入的近似误差: E[∥Mt−Mˉt∥F2]≤ks2∥Mt∥F2\mathbb{E}\left[\|\bm{M}_t - \bar{\bm{M}}_t\|_F^2\right] \leq \frac{k}{s^2}\|\bm{M}_t\|_F^2

其中 k 是保留的系数数量,s² 是块大小。

引理 8.1(偏差分析):

稀疏化后的梯度估计偏差: E[gˉt]=gt+et\mathbb{E}[\bar{\bm{g}}_t] = \bm{g}_t + \bm{e}_t

其中 ∥et∥≤ϵ\|\bm{e}_t\| \leq \epsilon,ε 与 k/s² 成正比。

五、实验详解

5.1 实验设置

配置OLMo-300MOLMo-1B
非嵌入参数320M1.18B
层数2424
隐藏维度10242048
注意力头数1616
序列长度20482048
训练 tokens100B100B
GPU64 × NVIDIA H10064 × NVIDIA H100
全局 batch size20482048
梯度累积步数44
每 GPU batch size88

超参数设置:

  • 学习率:线性 warmup + 余弦衰减
  • AdamW β₁=0.9, λ=0.1
  • DeMo 默认:β=0.999, α=0.2, s=64

5.2 主要结果

零样本下游任务评估

训练损失对比

优化器 / kHellaSwag ↑ARC-Easy ↑PIQA ↑通信量 (MB/step)
OLMo-300M
DeMo k=320.370.460.6729.9
DeMo k=160.380.500.6714.9
DeMo k=80.380.470.677.49
DeMo k=40.370.470.673.74
DeMo k=20.360.460.651.87
DeMo k=10.350.450.650.93
AdamW-DDP0.350.460.65636.9
OLMo-1B
DeMo k=320.480.550.70110.32
DeMo k=160.470.530.7055.16
DeMo k=80.470.520.6927.58
DeMo k=40.450.520.7013.79
DeMo k=20.440.510.696.89
DeMo k=10.410.520.693.44
AdamW-DDP0.430.510.682416.6

关键发现:

  1. 300M 模型:DeMo k=8 将通信量降低 85 倍(7.5 MB vs 637 MB),精度无损
  2. 1B 模型:DeMo k=16 降低 44 倍(55 MB vs 2417 MB),HellaSwag 和 PIQA 还优于 AdamW
  3. k=2 即可超越 AdamW:即使极端压缩也能保持竞争力

Perplexity vs 通信量

Perplexity与通信量

  • DeMo 在 Pareto 前沿上始终优于 AdamW
  • 相同通信预算下,DeMo 提供更好的模型质量
  • 相同模型质量下,DeMo 需要更少的通信

5.3 消融实验详解

5.3.1 线性变换选择

压缩方法对比

变换方法k=8k=16k=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 值压缩比性能说明
321024/k略差块太小,DCT 效果有限
644096/k最优默认设置
12816384/k略差块太大,边界效应

5.4 与其他优化器对比

优化器对比

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

压缩比与Perplexity

关键发现:

  • DeMo 在相同压缩比下始终优于 DiLoCo
  • 高压缩比下差距更明显(DeMo 更稳定)
  • DiLoCo 需要复杂的多机协调,DeMo 更简单

5.5 计算效率分析

每 GPU 通信量对比:

模型规模AdamW-DDPDeMo k=8DeMo k=16压缩比
300M637 MB7.5 MB14.9 MB85x / 44x
1B2417 MB27.6 MB55.2 MB88x / 44x

训练时间估算(假设 100Gbps 网络):

场景AdamW-DDPDeMo k=8加速比
300M, 64 GPU~5.1s/step~0.06s/step85x
1B, 64 GPU~19.3s/step~0.22s/step88x

六、代码实现分析

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 集成方式

  1. 禁用 DDP 梯度同步:

    model = DDP(model, find_unused_parameters=False)
    # 禁用 All-Reduce
    model.require_forward_param_sync = False
    model.require_backward_grad_sync = False
  2. 使用 DeMo 优化器:

    optimizer = DeMo(model.parameters(), lr=1e-3, k=8, chunk_size=64)
  3. 代码修改量:约 200 行核心实现,无需修改模型代码

七、相关工作对比

7.1 梯度压缩方法

方法压缩方式压缩比额外内存通信原语
QSGD标量量化4-32xO(1)All-Reduce
PowerSGD低秩近似10-50xO(rank)All-Reduce
Top-k稀疏化10-100xO(1)All-Reduce
1-bit Adam二值化32xO(1)1-bit All-Reduce
DeMoDCT + Top-k85xO(1)All-Gather

7.2 分层优化方法

方法原理通信模式系统复杂度
DiLoCo内/外层分离周期性同步高(多机协调)
FedAvg聚合本地更新周期性同步中
DeMo动量压缩同步每步压缩同步低

7.3 关键区别

特性DeMoDiLoCoPowerSGD
通信频率每步每 H 步每步
额外内存0O(model)O(rank)
系统复杂度低高中
拓扑无关性是部分是
可与 Muon 结合是是是

八、局限性与未来方向

8.1 当前局限性

  1. 超参数敏感:k、α、β、s 等参数需要针对不同模型调优
  2. 验证规模有限:仅在 300M 和 1B 参数模型上验证,未扩展到 10B+
  3. 固定稀疏模式:top-k 可能不是最优策略,自适应稀疏化可能更好
  4. 理论分析基础:收敛性分析基于标准假设,实际大规模训练可能有偏差
  5. 单机验证:实验在单节点 64 GPU 上进行,未验证跨节点场景

8.2 未来研究方向

  1. 大规模验证:在 10B+ 参数模型和跨数据中心场景中验证
  2. 自适应稀疏化:根据训练阶段动态调整 k 值
  3. 混合精度支持:结合 FP8/FP4 量化进一步压缩
  4. 异步训练:与异步 DDP 结合,进一步降低同步开销
  5. 硬件加速:设计专用硬件加速 DCT 和 Top-k 操作
  6. 理论深化:非凸优化下的收敛率分析

九、总结

核心贡献

  1. 提出 DeMo 优化器:通过解耦动量更新和结构化压缩,将分布式训练通信量降低 1-2 个数量级
  2. 结构化压缩流水线:分块 + DCT + Top-k 组合,利用梯度频域稀疏性实现高效压缩
  3. 动量减法误差反馈:创新性地重用动量缓冲区作为误差累加器,零额外内存开销
  4. 拓扑无关性:不依赖特定网络拓扑,支持跨数据中心和以太网训练
  5. 即插即用:可与任何基于动量的优化器结合,最小化代码修改

技术影响

  • 降低训练成本:使低带宽互连(以太网)也能进行大规模 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) - 开源语言模型框架