You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

torch.optim.SGD带动量优化效果异常问题问询

关于torch.optim.SGD带动量参数的异常表现疑问

我怀疑torch.optim.SGD在添加动量参数后存在异常。按我的理解,只需给momentum参数赋值就能实现带动量的SGD,示例代码如下:

torch.optim.SGD(params, lr=0.01, momentum=0.9)

我在复现PyTorch Lightning的优化器教程时,没有像教程那样从零实现优化器,而是直接调用torch.optim中的函数。具体来说,将教程里的:

SGDMom_points = train_curve(lambda params: SGDMomentum(params, lr=10, momentum=0.9))

替换为:

SGDMom_points = train_curve(lambda params: torch.optim.SGD(params, lr=10, momentum=0.9))

得到的结果远不如教程展示的效果。用以下代码实现的Nesterov加速梯度效果也不符合预期:

NAG_points = train_curve(lambda params: torch.optim.SGD(params, lr=10, momentum=0.9, nesterov=True))

我检查了代码,没发现可疑之处,希望有人能验证是否存在同样的差异。以下是完整代码(仅修改了计算SGD_points、SGDMom_points和Adam_points的几行):

完整代码

from matplotlib import cm
import seaborn as sns
from matplotlib import pyplot as plt
import torch
import numpy as np

def pathological_curve_loss(w1, w2):
    # 病态曲率示例,可自行尝试其他形式
    x1_loss = torch.tanh(w1) ** 2 + 0.01 * torch.abs(w1)
    x2_loss = torch.sigmoid(w2)
    return x1_loss + x2_loss

def plot_curve(
    curve_fn, x_range=(-5, 5), y_range=(-5, 5), plot_3d=False, cmap=cm.viridis, title="Pathological curvature"
):
    fig = plt.figure()
    ax = fig.add_subplot(projection='3d') if plot_3d else fig.gca()

    x = torch.arange(x_range[0], x_range[1], (x_range[1] - x_range[0]) / 100.0)
    y = torch.arange(y_range[0], y_range[1], (y_range[1] - y_range[0]) / 100.0)
    x, y = torch.meshgrid([x, y])
    z = curve_fn(x, y)
    x, y, z = x.numpy(), y.numpy(), z.numpy()

    if plot_3d:
        ax.plot_surface(x, y, z, cmap=cmap, linewidth=1, color="#000", antialiased=False)
        ax.set_zlabel("loss")
    else:
        ax.imshow(z.T[::-1], cmap=cmap, extent=(x_range[0], x_range[1], y_range[0], y_range[1]))
    plt.title(title)
    ax.set_xlabel(r"$w_1$")
    ax.set_ylabel(r"$w_2$")
    plt.tight_layout()
    return ax


# sns.reset_orig()
# _ = plot_curve(pathological_curve_loss, plot_3d=True)
# plt.show()

from torch import nn

def train_curve(optimizer_func, curve_func=pathological_curve_loss, num_updates=100, init=[5, 5]):
    """
    参数说明:
        optimizer_func: 优化器构造函数,仅接受参数列表作为输入
        curve_func: 损失函数(例如病态曲率函数)
        num_updates: 优化迭代的步数
        init: 参数初始值,需为包含两个元素的列表/元组,分别对应w_1和w_2
    返回:
        形状为[num_updates, 3]的numpy数组,其中[t,:2]为第t步的参数值,[t,2]为第t步的损失值
    """
    weights = nn.Parameter(torch.FloatTensor(init), requires_grad=True)
    optim = optimizer_func([weights])

    list_points = []
    for _ in range(num_updates):
        loss = curve_func(weights[0], weights[1])
        list_points.append(torch.cat([weights.data.detach(), loss.unsqueeze(dim=0).detach()], dim=0))
        optim.zero_grad()
        loss.backward()
        optim.step()
    points = torch.stack(list_points, dim=0).numpy()
    return points


# 以下是与教程不同的修改部分
SGD_points = train_curve(lambda params: torch.optim.SGD(params, lr=10))
SGDMom_points = train_curve(lambda params: torch.optim.SGD(params, lr=10, momentum=0.9))
Adam_points = train_curve(lambda params: torch.optim.Adam(params, lr=1))
# 修改部分结束

all_points = np.concatenate([SGD_points, SGDMom_points, Adam_points], axis=0)
ax = plot_curve(
    pathological_curve_loss,
    x_range=(-np.absolute(all_points[:, 0]).max(), np.absolute(all_points[:, 0]).max()),
    y_range=(all_points[:, 1].min(), all_points[:, 1].max()),
    plot_3d=False,
)
ax.plot(SGD_points[:, 0], SGD_points[:, 1], color="red", marker="o", zorder=1, label="SGD")
ax.plot(SGDMom_points[:, 0], SGDMom_points[:, 1], color="blue", marker="o", zorder=2, label="SGDMom")
ax.plot(Adam_points[:, 0], Adam_points[:, 1], color="grey", marker="o", zorder=3, label="Adam")
plt.legend()
plt.show()

内容的提问来源于stack exchange,提问作者user559678

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 04:46:48