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

如何在PyTorch中实现《Attention Is All You Need》自定义学习率调度器

Transformer学习率调度器PyTorch实现方案

PyTorch官方没有内置和《Attention Is All You Need》论文完全匹配的学习率调度器,你需要自定义调度器类,实现逻辑和TensorFlow版本1:1对齐即可。

第一步:自定义调度器类

from torch.optim.lr_scheduler import LRScheduler
import torch

class TransformerLR(LRScheduler):
    def __init__(self, optimizer, d_model, warmup_steps=4000, last_epoch=-1):
        self.d_model = d_model
        self.warmup_steps = warmup_steps
        # 父类初始化要放在自定义参数之后,避免调用get_lr时参数未定义
        super().__init__(optimizer, last_epoch)
    
    def get_lr(self):
        # 取当前步数,避免step为0时报错
        current_step = max(1, self.last_epoch)
        # 和原TensorFlow逻辑完全对齐的计算
        arg1 = current_step ** (-0.5)
        arg2 = current_step * (self.warmup_steps ** (-1.5))
        computed_lr = (self.d_model ** (-0.5)) * min(arg1, arg2)
        # 对优化器的所有参数组应用相同的学习率
        return [computed_lr for _ in self.base_lrs]

第二步:补全你的现有代码

你只需要实例化上面的自定义调度器即可,优化器初始化时设置的lr参数不会生效,会被调度器计算的学习率覆盖:

import torch 

# 这里的lr参数可任意填写,会被调度器覆盖
optimizer = torch.optim.Adam(model.parameters(), lr=0.0, betas=(0.9, 0.98), eps=1e-9)
# 替换d_model为你实际使用的模型维度,warmup_steps可根据需求调整
scheduler = TransformerLR(optimizer, d_model=512, warmup_steps=4000)

使用注意事项

  • 该调度器是按训练步数更新的,需要在每次调用optimizer.step()之后执行scheduler.step(),不要放在epoch结束时才调用
  • 如果你需要中断后恢复训练,只需要在实例化调度器时传入last_epoch=已训练的步数,即可从对应位置继续计算学习率

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 10:45:03