如何在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
使用说明
- 确保全批量梯度计算:在调用
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() # 重置梯度缓存
- 启用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
相关产品推荐
相关产品推荐

