def init_adamw_states(feature_dim):
m_w, m_b = jnp.zeros((feature_dim, 1)), jnp.zeros(1)
v_w, v_b = jnp.zeros((feature_dim, 1)), jnp.zeros(1)
return [(m_w, v_w), (m_b, v_b)]
def adamw(params, grads, states, hyperparams):
beta1, beta2, eps = 0.9, 0.999, 1e-6
for i, (p, (m, v), g) in enumerate(zip(params, states, grads)):
m = beta1 * m + (1 - beta1) * g
v = beta2 * v + (1 - beta2) * jnp.square(g)
m_hat = m / (1 - beta1 ** hyperparams['t'])
v_hat = v / (1 - beta2 ** hyperparams['t'])
params[i] = p - hyperparams['lr'] * (
m_hat / (jnp.sqrt(v_hat) + eps) + hyperparams['wd'] * p)
states[i] = (m, v)
hyperparams['t'] += 1
return params[0], params[1]