def init_adamw_states(feature_dim):
m_w, m_b = d2l.zeros((feature_dim, 1)), d2l.zeros(1)
v_w, v_b = d2l.zeros((feature_dim, 1)), d2l.zeros(1)
return ((m_w, v_w), (m_b, v_b))
def adamw(params, states, hyperparams):
beta1, beta2, eps = 0.9, 0.999, 1e-6
for p, (m, v) in zip(params, states):
with torch.no_grad():
m[:] = beta1 * m + (1 - beta1) * p.grad
v[:] = beta2 * v + (1 - beta2) * torch.square(p.grad)
m_hat = m / (1 - beta1 ** hyperparams['t'])
v_hat = v / (1 - beta2 ** hyperparams['t'])
p[:] -= hyperparams['lr'] * (m_hat / (torch.sqrt(v_hat) + eps)
+ hyperparams['wd'] * p)
p.grad.zero_()
hyperparams['t'] += 1