引言

在本节中,我们会通过继承torch.optim.Optimizer,实现AdamW优化器

手写优化器的一般方法

先说一下继承torch.optim.Optimizer手写优化器的通用方法,首先就是所有优化器都要求在构造函数中传入进行优化的参数,统一命名为params,然后在构造函数中传入需要的超参数。传入的超参数需要通过审核,确认无误之后构造一个名为default的字典,里面包含了所有的超参数,然后将params和default传入父类的构造函数中,此时父类会自动生成param_groups。

接着工作便是通过step来更新参数,我们直接获得自身的param_groups,然后遍历每一个group,获得当前group超参数,这时我们可以遍历group中的每一个参数,获得每一个参数的状态,同时对每一个参数进行操作,优化。

具体实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
import torch
from torch import Tensor
from collections.abc import Iterable
class AdamW(torch.optim.Optimizer):
def __init__(self,params:Iterable[Tensor]|list[dict],
lr: float | Tensor = 1e-3,
betas: tuple[float | Tensor, float | Tensor] = (0.9, 0.999),
eps: float = 1e-8,
weight_decay: float = 0.01):
if lr<0.0:
raise ValueError("learning rate can't to be a negative number")
if betas[0]<0.0 or betas[0]>=1.0 or betas[1]<0.0 or betas[1]>=1.0:
raise ValueError("betas need to be in the [0.0,1.0)")
if eps<=0.0:
raise ValueError("eps need to be a positive number")
if weight_decay<0.0:
raise ValueError("weight_decay can't to be a negative number")
defaults={
"lr":lr,
"betas":betas,
"eps":eps,
"weight_decay":weight_decay
}
super().__init__(params,defaults)


@torch.no_grad
def step(self,closure=None):
loss=None
if closure is not None:
with torch.enable_grad():
loss=closure()

for param_group in self.param_groups:
lr=param_group['lr']
betas=param_group['betas']
eps=param_group['eps']
weight_decay=param_group['weight_decay']
for p in param_group['params']:
if p.grad is None:
continue
state=self.state[p]
# 获取状态
if len(state)==0:
step=state.get('step',0)
exp_avg=state.get('exp_avg',torch.zeros_like(p,dtype=p.dtype,device=p.device))
exp_avg_sq=state.get('exp_avg_sq',torch.zeros_like(p,dtype=p.dtype,device=p.device))
else:
step=state['step']
exp_avg=state['exp_avg']
exp_avg_sq=state['exp_avg_sq']
step+=1
state['step']=step

# 权重衰减
p.sub_(lr*weight_decay*p)

# 计算动量一阶矩和二阶矩
exp_avg=betas[0]*exp_avg+(1-betas[0])*p.grad
exp_avg_sq=betas[1]*exp_avg_sq+(1-betas[1])*p.grad*p.grad

state['exp_avg']=exp_avg
state['exp_avg_sq']=exp_avg_sq

# 修正动量
exp_avg=exp_avg/(1-betas[0]**step)
exp_avg_sq=exp_avg_sq/(1-betas[1]**step)

# 更新权重
p.sub_((lr*exp_avg)/(torch.sqrt(exp_avg_sq)+eps))
return loss

这里在进行step之前我们需要使用@torch.no_grad修饰step,防止pytorch的计算图扩张造成内存浪费

算法:

AdamW算法

注意这里是提前把修正因子更新到学习率上了,而代码的实现则是后面直接对一阶矩和二阶矩进行修正