如何在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
相关产品推荐
相关产品推荐

