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
相关产品推荐
相关产品推荐

