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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 13:17:21