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

如何在PyTorch中实现带Nesterov动量的全批量梯度下降?

实现支持Nesterov动量的全批量梯度下降优化器(PyTorch)

核心实现思路

全批量梯度下降(GD)的核心是每次参数更新基于整个数据集的梯度均值,优化器的职责是按照GD+Nesterov动量的公式完成参数更新。通过子类化torch.optim.Optimizer,我们可以实现完全兼容PyTorch生态的自定义优化器,同时严格遵循Nesterov动量的更新逻辑。

自定义优化器代码

import torch
from torch.optim.optimizer import Optimizer

class FullBatchGD(Optimizer):
    def __init__(self, params, lr=0.01, momentum=0.0, weight_decay=0.0):
        if lr < 0.0:
            raise ValueError(f"Invalid learning rate: {lr}")
        if momentum < 0.0:
            raise ValueError(f"Invalid momentum value: {momentum}")
        if weight_decay < 0.0:
            raise ValueError(f"Invalid weight_decay value: {weight_decay}")
        
        defaults = dict(lr=lr, momentum=momentum, weight_decay=weight_decay)
        super().__init__(params, defaults)

    def step(self, closure=None):
        loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            lr = group['lr']
            momentum = group['momentum']
            weight_decay = group['weight_decay']

            for p in group['params']:
                if p.grad is None:
                    continue
                grad = p.grad.data

                # 处理权重衰减
                if weight_decay != 0:
                    grad.add_(p.data, alpha=weight_decay)

                state = self.state[p]

                # 初始化动量缓存
                if len(state) == 0:
                    if momentum != 0:
                        state['momentum_buffer'] = torch.zeros_like(p.data)

                if momentum != 0:
                    buf = state['momentum_buffer']
                    # Nesterov动量:先更新缓存,再应用两次梯度修正
                    buf.mul_(momentum).add_(grad)
                    p.data.add_(buf, alpha=-lr)
                    p.data.add_(grad, alpha=-lr * momentum)
                else:
                    # 无动量的全批量GD
                    p.data.add_(grad, alpha=-lr)

        return loss

使用说明

  1. 确保全批量梯度计算:在调用optimizer.step()之前,必须基于整个数据集完成前向传播和反向传播。可以直接遍历全数据集累加梯度,或用DataLoader设置batch_size=len(full_dataset)一次性加载全量数据:
# 示例:用DataLoader加载全量数据
full_loader = torch.utils.data.DataLoader(full_dataset, batch_size=len(full_dataset))
model.train()
for inputs, targets in full_loader:
    outputs = model(inputs)
    loss = criterion(outputs, targets)
    loss.backward()  # 基于全量数据计算梯度
optimizer.step()
optimizer.zero_grad()  # 重置梯度缓存
  1. 启用Nesterov动量:初始化优化器时设置momentum>0即可,例如:
optimizer = FullBatchGD(model.parameters(), lr=0.001, momentum=0.9, weight_decay=1e-4)

关键细节说明

  • 与torch.optim.SGD的区别:本优化器严格遵循Nesterov动量的原始更新逻辑,无需依赖batch size设置模拟全批量更新,优化器本身的参数更新行为完全匹配全批量GD的定义。
  • 兼容性:继承自torch.optim.Optimizer,支持PyTorch生态的所有标准操作(如zero_grad()、参数分组、学习率调度器等)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 16:15:21