🚀 为什么 PyTorch 里的动量公式和教材不一样?从理论到实践的一次踩坑复盘

大家好,我是书到用时方恨少!最近在学习动量梯度下降时,我发现了一个特别容易让人困惑的问题:

明明教材里的动量公式是vt=βvt−1+(1−β)gtv_t = \beta v_{t-1} + (1-\beta)g_tvt=βvt1+(1β)gt,但 PyTorch 里 torch.optim.SGD(momentum=0.9) 用的却是vt=βvt−1+gtv_t = \beta v_{t-1} + g_tvt=βvt1+gt

这两个公式看着差不多,结果却可能天差地别。今天,我就带着你从「问题发现」到「原因拆解」再到「代码验证」,把这个坑彻底填平,让你以后再也不会被这个细节搞懵!


🔍 一、问题发现:我的代码结果和理论公式对不上?

我写了一段超简单的 PyTorch 代码,测试带动量的 SGD:

import torch

def test_momentum():
    # 初始化权重 w=1.0,损失函数 L = w²/2,梯度 g = w
    w = torch.tensor([1.0], requires_grad=True, dtype=torch.float32)
    optimizer = torch.optim.SGD([w], lr=0.01, momentum=0.9)

    # 第1次更新
    loss = ((w ** 2) / 2.0).sum()
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    print(f"第1次:w.grad={w.grad.item()}, 更新后w={w.item()}")

    # 第2次更新
    loss = ((w ** 2) / 2.0).sum()
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    print(f"第2次:w.grad={w.grad.item()}, 更新后w={w.item()}")

test_momentum()

运行结果是:

第1次:w.grad=1.0, 更新后w=0.99
第2次:w.grad=0.99, 更新后w=0.9711

我当时就懵了:

  • 按教材公式 vt=0.9vt−1+0.1gtv_t = 0.9v_{t-1} + 0.1g_tvt=0.9vt1+0.1gt计算:
    • 第1次:v1=1.0v_1 = 1.0v1=1.0w=1.0−0.01∗1.0=0.99w = 1.0 - 0.01*1.0 = 0.99w=1.00.011.0=0.99 ✅(第一次因为没有前置项,所以v1v_1v1直接等于本次梯度,如果为了统一公式,大家也可以用动量公式来计算v1v_1v1)
    • 第2次的速度项 v2=0.9∗1.0+0.1∗0.99=0.999v_2 = 0.9 * 1.0 + 0.1 * 0.99=0.999v2=0.91.0+0.10.99=0.999,更新应该是 w=0.99−0.01∗0.999=0.98001w = 0.99 - 0.01*0.999 = 0.98001w=0.990.010.999=0.98001,和实际结果完全不符!
  • 但如果按 vt=0.9vt−1+gtv_t = 0.9v_{t-1} + g_tvt=0.9vt1+gt 计算:
    • 第1次:v1=1.0v_1 = 1.0v1=1.0w=1.0−0.01∗1.0=0.99w = 1.0 - 0.01*1.0 = 0.99w=1.00.011.0=0.99
    • 第2次:v2=0.9∗1.0+0.99=1.89v_2 = 0.9*1.0 + 0.99 = 1.89v2=0.91.0+0.99=1.89w=0.99−0.01∗1.89=0.9711w = 0.99 - 0.01*1.89 = 0.9711w=0.990.011.89=0.9711

经常多方验证才知道,原来,PyTorch 根本没按教材里的「带 1−β1-\beta1β」公式来实现!


📚 二、问题拆解:两种公式到底差在哪?

我们把两种公式摆在一起,一眼看清区别:

对比项 教材常见形式(EMA 形式) PyTorch 实现形式(标准动量)
速度项更新 vt=βvt−1+(1−β)gtv_t = \beta v_{t-1} + (1-\beta)g_tvt=βvt1+(1β)gt vt=βvt−1+gtv_t = \beta v_{t-1} + g_tvt=βvt1+gt
参数更新 wt+1=wt−ηvtw_{t+1} = w_t - \eta v_twt+1=wtηvt wt+1=wt−ηvtw_{t+1} = w_t - \eta v_twt+1=wtηvt
核心特点 速度项是梯度的指数移动平均,稳定梯度不变时,vtv_tvt 会收敛到 gtg_tgt 速度项是未归一化的累积梯度,梯度会被直接累加
等价关系 可通过调整学习率互相转换 可通过调整学习率互相转换

1. 为什么会有两种不同的写法?

(1)教材里的 EMA 形式:为了“归一化”而存在

教材里的 1−β1-\beta1β,本质上是为了让速度项 vtv_tvt 成为一个加权和为1的指数移动平均

举个例子,当 β=0.9\beta=0.9β=0.9 时:
vt=0.9vt−1+0.1gtv_t = 0.9v_{t-1} + 0.1g_tvt=0.9vt1+0.1gt
展开后,vt=0.1gt+0.1∗0.9gt−1+0.1∗0.92gt−2+...v_t = 0.1g_t + 0.1*0.9g_{t-1} + 0.1*0.9^2g_{t-2} + ...vt=0.1gt+0.10.9gt1+0.10.92gt2+...,所有系数加起来正好是1。这样,当梯度 gtg_tgt 稳定不变时,vtv_tvt 也会稳定在 gtg_tgt 附近,方便我们理解和可视化梯度的变化。

(2)PyTorch 里的标准动量:更贴近物理“惯性”直觉

PyTorch 去掉 1−β1-\beta1β,是为了实现更经典的「标准动量」:
vt=βvt−1+gt v_t = \beta v_{t-1} + g_t vt=βvt1+gt
这个公式的物理意义更直接:

  • βvt−1\beta v_{t-1}βvt1:保留上一步的“惯性”,让更新方向更稳定
  • gtg_tgt:加入当前梯度的“加速度”,修正方向
  • 整体效果:速度项会随着迭代不断累积梯度的“动量”,就像下山时越跑越快的球。

2. 关键结论:两种公式是等价的!

你可能会问:“公式不一样,结果肯定不一样吧?”

其实,它们只是线性缩放关系,只要调整学习率,就能得到完全相同的更新结果!

我们用你的例子来验证:

  • PyTorch 形式:η=0.01\eta=0.01η=0.01β=0.9\beta=0.9β=0.9,更新结果是w2=0.9711w_2=0.9711w2=0.9711
  • EMA 形式:我们把公式改成 vt=0.9vt−1+0.1gtv_t = 0.9v_{t-1} + 0.1g_tvt=0.9vt1+0.1gt,同时把学习率放大10倍η=0.1\eta=0.1η=0.1
    • 第1次:v1=1.0v_1 = 1.0v1=1.0w=1.0−0.01∗0.1=0.99w = 1.0 - 0.01*0.1 = 0.99w=1.00.010.1=0.99(第一次因为v1v_1v1没有前置项,所以大小直接等于本次梯度,并且因为第一次没用到动量公式,所以我们第一次的学习率就不用改了,当然为了统一,也可以用动量公式算v1v_1v1)
    • 第2次:v2=0.9∗0.1+0.1∗0.99=0.189v_2 = 0.9*0.1 + 0.1*0.99 = 0.189v2=0.90.1+0.10.99=0.189w=0.99−0.1∗0.189=0.9711w = 0.99 - 0.1*0.189 = 0.9711w=0.990.10.189=0.9711

你看,结果完全一样! 因为:
vtPyTorch=11−β⋅vtEMA v_t^{\text{PyTorch}} = \frac{1}{1-\beta} \cdot v_t^{\text{EMA}} vtPyTorch=1β1vtEMA
所以只要把学习率乘以 1−β1-\beta1β,两种形式的更新效果就完全等价了。


🛠️ 三、问题解决:如何在 PyTorch 中理解和使用动量?

搞懂了公式差异,我们再回到 PyTorch 的实现,看看怎么正确理解和使用它。

1. PyTorch 中 SGD + Momentum 的完整公式

PyTorch 的 torch.optim.SGD 还有两个参数:dampeningweight_decay,完整的速度项更新公式是:
vt=β⋅vt−1+(1−dampening)⋅gt v_t = \beta \cdot v_{t-1} + (1 - \text{dampening}) \cdot g_t vt=βvt1+(1dampening)gt

  • 默认 dampening=0,所以公式就变成了我们看到的 vt=βvt−1+gtv_t = \beta v_{t-1} + g_tvt=βvt1+gt
  • dampening 可以理解为“当前梯度的衰减系数”,当 dampening=0.1 时,当前梯度会被乘以0.9再加入速度项,进一步平滑更新

2. 如何查看 PyTorch 中动量的累积梯度?

很多人不知道,PyTorch 优化器内部维护了一个 momentum_buffer,也就是我们说的速度项 vtv_tvt。我们可以直接打印出来:

def test_momentum_with_buffer():
    w = torch.tensor([1.0], requires_grad=True, dtype=torch.float32)
    optimizer = torch.optim.SGD([w], lr=0.01, momentum=0.9)

    # 第1次更新
    loss = ((w ** 2) / 2.0).sum()
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    v1 = optimizer.state[w]['momentum_buffer']
    print(f"第1次:原始梯度w.grad={w.grad.item()}, 动量累积梯度v={v1.item()}, 更新后w={w.item()}")

    # 第2次更新
    loss = ((w ** 2) / 2.0).sum()
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    v2 = optimizer.state[w]['momentum_buffer']
    print(f"第2次:原始梯度w.grad={w.grad.item()}, 动量累积梯度v={v2.item()}, 更新后w={w.item()}")

test_momentum_with_buffer()

运行结果:

第1次:原始梯度w.grad=1.0, 动量累积梯度v=1.0, 更新后w=0.99
第2次:原始梯度w.grad=0.99, 动量累积梯度v=1.89, 更新后w=0.9711

这下彻底清晰了:

  • w.grad 永远是当前步骤的原始梯度,不会被动量修改
  • 动量的累积梯度 vtv_tvt 存在优化器的 state 字典里,参数更新时用的是 vtv_tvt,而不是直接用 w.grad

✅ 四、结论:给新手的3条关键建议

  1. 别被教材公式困住:PyTorch 用的是「标准动量」形式,没有 1−β1-\beta1β,但只要调整学习率,和 EMA 形式完全等价。
  2. 区分清楚两个梯度
    • w.grad:当前步骤的原始梯度,不会被优化器修改
    • momentum_buffer:优化器内部维护的累积梯度,参数更新用的是它
  3. 调参时注意学习率缩放:如果你习惯了教材里的 EMA 形式,换到 PyTorch 时,记得把学习率乘以 1−β1-\beta1β 来抵消公式差异。

💡 写在最后

很多时候,我们学习算法时,教材里的公式和工程实现总会有一些细节差异。这次关于动量公式的踩坑,让我意识到:只有把理论公式和底层实现对照着看,才能真正理解算法的本质。

如果你也在学习深度学习优化器,不妨自己写一段小代码,打印出 momentum_buffer 看看,相信你会对动量的理解更上一层楼!

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐