跳到主要内容

一个优化器每步慢 25%,在 7 次梯度累积里怎么只剩 4%?

陈渊
陈渊

· 阅读约 3 分钟

一个优化器单步慢 25%,到七步累积的窗口里,整体开销只多 4%。这不是玄学,是摊薄。但摊薄具体发生在哪一行,多数人停在“累积步数大了自然就小”这种直觉上,没有再往下挖。这篇我们动手把 Newton-Schulz 正交化和梯度累积两层搭起来,跑一遍,看那个 25% 怎么被 7 次 micro-step 吞掉,然后再往上加一层 DDP——那才是真正反直觉的地方。

先把最朴素的 Newton-Schulz 版本写出来。它只做一件事:把梯度矩阵朝离它最近的正交矩阵拽。标准迭代骨架长这样:

def newton_schulz(mat, steps=5):
    a = mat / (mat.norm() + 1e-12)      # 先谱归一化,离 I 近一点
    for _ in range(steps):
        a = 1.5 * a - 0.5 * (a @ a.T) @ a
    return a

a @ a.T 这一项是矩阵平方级的乘法,每步五次迭代就是五次这种乘法。作者那个浅累积 benchmark 里测出的 25% 慢,基本就是这些迭代相对于 AdamW 逐元素更新的额外成本。真实的 Muon 不只是给矩阵参数用,还把 grad 乘 lr 之后喂进来,但骨架是这段。

但优化器不是每个 micro batch 都要跑。接下来把这层和梯度累积叠在一起:

num_micros = 7
for i, (x, y) in enumerate(train_loader):
    loss = model(x, y) / num_micros   # 把 loss 按累积步数分下去
    loss.backward()                    # gradient 只累加,不碰优化器
    if (i + 1) % num_micros == 0:
        optimizer.step()               # 这里才跑 newton_schulz
        optimizer.zero_grad()

追一次执行:第 1 个 micro batch 到第 6 个,loss.backward() 只把梯度累加到 param.grad,if 都不成立;到第 7 个,optimizer.step() 才进优化器,里面每个矩阵参数各跑一次 newton_schulz。七次 backward 配一次 Newton-Schulz,这就是摊薄。把单步 25% 的成本除以 7,平均下来就是约 3.6%——实际还会多一点,因为前向+反向每一步都不是零成本,但量级正好落在那 4% 附近。

到这一步,单卡上的算术已经清楚。但单机八卡训练不是八张卡各跑互不相干的单卡。原作者用 plain DDP,这意味着每个 rank 都有完整的一份 optimizer state。上面的 optimizer.step() 在每个 rank 上跑,Newton-Schulz 在每个 rank 上重算一遍。这层冗余画出来更直观:

# 8 个 rank,每个 rank 内都执行:
def optimizer_step_local(full_params, lr):
    for p in full_params.matrix_params:
        g = p.grad * lr
        p.data -= newton_schulz(g, steps=5)
    # 没有分片:8 个 rank 都在算同一批完整的 p

这里的 all_reduce_grads 只是同步梯度,真正的重算发生在 newton_schulz——每张卡都算,参数完全一样,没有分片。梯度通信在这个配置下根本不是瓶颈,3.8B 参数在单节点上同步用不了多少时间;冗余重算才是被摊薄掩盖掉的另一头。

我一开始也以为作者说 DDP 下 Muon 的更新成本就只是那 25% 的步慢。翻到他记录 DDP 那段才确认,真正多出来的是这八个 rank 各自重算并各自存状态。nanochat 的做法是把优化器状态分片,reduce-scatter、compute、all-gather 叠起来,只让每个 rank 算自己那 1/N 参数片,再把结果拼回去。这样八份重算就压成一份带通信的协同计算。梯度通信不构成约束,是因为它本来就便宜;优化器重算不构成伤害,不是因为它便宜,是因为梯度累积把它按 7 次摊了一遍,但它并没有被按 8 份摊——这个区别不能混。

所以这件事其实有两层:窗口内的摊薄,让每步慢 25% 变成整体多 4%;窗口外的放大,让 plain DDP 下的优化器重算在每张卡上冗余一遍,而梯度累积对此一点忙都帮不上。Muon 的代价不在牛顿-舒尔茨本身慢不慢,而在于你把它放在哪个循环外面、有没有给它分片。

想再深一层,有两个入口自己挖:一是去翻 nanochat 的 ZeRO-2 优化器分片实现,看它怎么把牛顿-舒尔茨的输入按参数张量切到各 rank 再 all-gather,怎么把通信和计算重叠起来;二是把本文这个最小 Newton-Schulz 实现加一个梯度累积计时器,自己测一遍 25% 怎么摊成约 4%——不是看出来的,是跑出来的。

陈渊
陈渊

据守底层,挑「会用却说不出为什么」的 CS 机制从第一性原理挖到底。

查看主页 →