Adam梯度下降算法实战指南:从数学原理到PyTorch实现

1次阅读
没有评论

共计 2876 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

传统优化器的局限性

在深度学习模型训练中,随机梯度下降(SGD)是最基础的优化算法。但 SGD 存在几个明显问题:

  • 对学习率(learning rate)非常敏感,需要精心调参
  • 在非凸函数优化中容易陷入局部最优解
  • 所有参数使用相同的学习率,无法适应不同参数的特性
  • 遇到平坦区域时梯度接近于零,导致训练停滞

这些局限性使得 SGD 在实际应用中往往收敛缓慢,需要大量调参才能获得较好效果。

主流优化算法对比

优化算法 收敛速度 内存占用 超参数敏感性 适用场景
SGD 简单模型
Momentum 中等 中等 一般场景
RMSprop 中等 RNN/LSTM
Adam 最快 中等 大多数 DL 模型

Adam 算法原理详解

Adam(Adaptive Moment Estimation)结合了 Momentum 和 RMSprop 的思想,通过计算梯度的一阶矩估计和二阶矩估计来动态调整每个参数的学习率。

核心公式推导

  1. 计算梯度的一阶矩估计(动量):
    $$m_t = \beta_1 m_{t-1} + (1-\beta_1)g_t$$

  2. 计算梯度的二阶矩估计(自适应学习率):
    $$v_t = \beta_2 v_{t-1} + (1-\beta_2)g_t^2$$

  3. 偏差修正(针对初期迭代):
    $$\hat{m}_t = \frac{m_t}{1-\beta_1^t}$$
    $$\hat{v}_t = \frac{v_t}{1-\beta_2^t}$$

  4. 参数更新:
    $$\theta_t = \theta_{t-1} – \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon}$$

其中:
– $\beta_1$(默认 0.9):控制一阶矩估计的衰减率
– $\beta_2$(默认 0.999):控制二阶矩估计的衰减率
– $\epsilon$(默认 1e-8):数值稳定项

PyTorch 实现完整代码

import torch
import math

class CustomAdam(torch.optim.Optimizer):
    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
                 weight_decay=0, warmup_steps=4000):
        defaults = dict(lr=lr, betas=betas, eps=eps,
                        weight_decay=weight_decay, warmup_steps=warmup_steps)
        super(CustomAdam, self).__init__(params, defaults)

    def step(self, closure=None):
        loss = None
        if closure is not None:
            loss = closure()

        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad.data
                if grad.is_sparse:
                    raise RuntimeError('Adam does not support sparse gradients')

                state = self.state[p]

                # 初始化状态
                if len(state) == 0:
                    state['step'] = 0
                    state['m'] = torch.zeros_like(p.data)
                    state['v'] = torch.zeros_like(p.data)

                m, v = state['m'], state['v']
                beta1, beta2 = group['betas']
                state['step'] += 1

                # 梯度裁剪(可选)torch.nn.utils.clip_grad_norm_(p, max_norm=1.0)

                # 更新一阶和二阶矩估计
                m.mul_(beta1).add_(1 - beta1, grad)
                v.mul_(beta2).addcmul_(1 - beta2, grad, grad)

                # 偏差修正
                bias_correction1 = 1 - beta1 ** state['step']
                bias_correction2 = 1 - beta2 ** state['step']

                # 学习率 warmup
                warmup_factor = min(state['step'] / group['warmup_steps'], 1.0)
                current_lr = group['lr'] * warmup_factor

                # 参数更新
                denom = (v.sqrt() / math.sqrt(bias_correction2)).add_(group['eps'])
                step_size = current_lr / bias_correction1
                p.data.addcdiv_(-step_size, m, denom)

                # 权重衰减(L2 正则化)if group['weight_decay'] > 0:
                    p.data.add_(-group['lr'] * group['weight_decay'], p.data)

        return loss

常见问题及解决方案

  1. 训练初期震荡严重
  2. 原因:β1/β2 取值不当,导致动量估计不准确
  3. 解决:适当调高 β2(如 0.99→0.999)或增加 warmup 步骤

  4. 后期收敛变慢

  5. 原因:自适应学习率过度衰减
  6. 解决:尝试学习率线性衰减或余弦退火

  7. 梯度爆炸

  8. 原因:未做梯度裁剪
  9. 解决:添加 clip_grad_norm_,norm 阈值设为 1.0-5.0

性能对比实验

我们在 MNIST 数据集上对比了不同优化器的表现(使用相同的两层 CNN):

import matplotlib.pyplot as plt

# 训练代码省略...

plt.figure(figsize=(10,6))
plt.plot(sgd_losses, label='SGD')
plt.plot(momentum_losses, label='Momentum')
plt.plot(rmsprop_losses, label='RMSprop')
plt.plot(adam_losses, label='Adam')
plt.xlabel('Epochs')
plt.ylabel('Loss')
plt.legend()
plt.title('Optimizer Comparison on MNIST')
plt.show()

Adam 梯度下降算法实战指南:从数学原理到 PyTorch 实现

实验显示 Adam 在初期收敛速度明显快于其他优化器,最终 loss 也最低。

延伸思考

  1. Adam 是否适合所有网络结构?
  2. 有研究表明在 Transformer 等结构中,Adam 表现优异
  3. 但对于某些特定任务(如 GAN 训练),可能需要定制优化策略

  4. 如何证明 Adam 的收敛性?

  5. 原始论文通过构造 Lyapunov 函数证明
  6. 实际应用中可以通过监控梯度方差来判断

总结

Adam 通过结合动量和自适应学习率,在大多数深度学习任务中都能取得良好效果。本文从原理推导到代码实现,完整展示了如何应用 Adam 优化器。关键点包括:

  • 理解一阶 / 二阶矩估计的物理意义
  • 合理设置 β1/β2 和 warmup 策略
  • 注意梯度裁剪和学习率衰减
  • 根据任务特点灵活调整超参数

希望这篇指南能帮助初学者快速掌握 Adam 优化器的使用技巧。

正文完
 0
评论(没有评论)