PyTorch自定义优化器实现求助:基于导数符号调整步长
PyTorch自定义优化器实现
首先修正你初始代码中的拼写错误(super调用里的slef应为self),然后按照需求补全实现逻辑:
import torch.optim as optim class MyOpt(optim.Optimizer): def __init__(self, params, lr=1.0): # 初始化默认参数:学习率、当前导数符号d1、前一次导数符号d2 defaults = dict(lr=lr, d1=None, d2=None) super(MyOpt, self).__init__(params, defaults) def step(self, closure=None): loss = None if closure is not None: loss = closure() # 遍历每个参数组 for group in self.param_groups: # 取出当前组的参数状态 current_lr = group['lr'] prev_prev_sign = group['d2'] prev_sign = group['d1'] # 遍历组内每个可优化参数 for p in group['params']: if p.grad is None: continue grad = p.grad.data current_sign = torch.sign(grad) # 第一次迭代时,仅初始化前一次符号记录,不调整步长 if prev_prev_sign is None: group['d2'] = current_sign else: # 比较连续两次导数符号是否不同 if not torch.all(prev_sign == prev_prev_sign): group['lr'] /= 2.0 # 可选:添加步长下限,防止步长过小 # group['lr'] = max(group['lr'], 1e-6) # 执行参数更新(梯度下降方向) p.data -= group['lr'] * grad # 更新符号记录,为下一次迭代做准备 group['d2'] = prev_sign group['d1'] = current_sign return loss
关键逻辑说明:
- 符号跟踪机制:
d1存储当前迭代的导数符号,d2存储上一次迭代的导数符号,每次迭代后完成两者的状态更新。 - 步长调整规则:仅当连续两次导数符号不一致时,将当前参数组的学习率减半;第一次迭代时直接初始化符号记录,不调整步长。
- 兼容性处理:跳过无梯度的参数,保留
closure参数以兼容需要重新计算损失的训练场景。
内容的提问来源于stack exchange,提问作者Michal
相关产品推荐
相关产品推荐

