diff --git a/tests/test_adabelief_beta1.py b/tests/test_adabelief_beta1.py new file mode 100644 index 00000000..c1431c50 --- /dev/null +++ b/tests/test_adabelief_beta1.py @@ -0,0 +1,41 @@ +"""beta1 = 0 must not collapse the AdaBelief second moment to eps.""" +import torch +import torch_optimizer as optim + + +def test_positive_beta1_uses_the_residual(): + param = torch.nn.Parameter(torch.tensor([1.0, -0.5])) + opt = optim.AdaBelief( + [param], + lr=1e-3, + betas=(0.9, 0.999), + eps=1e-8, + rectify=False, + weight_decouple=False, + ) + grad = torch.tensor([0.4, -0.2]) + param.grad = grad.clone() + opt.step() + exp_avg = opt.state[param]["exp_avg"] + residual = grad - exp_avg + # s = (1 - beta2) residual^2, then eps is added into the state + expected = (1 - 0.999) * residual * residual + 1e-8 + assert torch.allclose(opt.state[param]["exp_avg_var"], expected) + + +def test_beta1_zero_stays_finite(): + param = torch.nn.Parameter(torch.tensor([0.5])) + opt = optim.AdaBelief( + [param], + lr=1e-2, + betas=(0.0, 0.999), + eps=1e-16, + rectify=False, + weight_decouple=False, + weight_decay=0.0, + ) + for _ in range(30): + param.grad = (2 * param.detach()).clone() + opt.step() + assert torch.isfinite(param.detach()).all() + assert param.detach().abs().item() < 0.5 diff --git a/torch_optimizer/adabelief.py b/torch_optimizer/adabelief.py index 7131f6f6..2ce331d6 100644 --- a/torch_optimizer/adabelief.py +++ b/torch_optimizer/adabelief.py @@ -161,7 +161,15 @@ def step(self, closure: OptLossClosure = None) -> OptFloat: # Update first and second moment running average exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1) - grad_residual = grad - exp_avg + # beta1 == 0 sets m_t = g_t, so (g_t - m_t) is 0 and the + # second moment stays at eps. The step is lr * g / sqrt(eps), + # which overflows when eps is the paper's 1e-16. Adam with + # beta1 = 0 still tracks g^2. Keep (g - m) for every + # positive beta1. + if beta1 == 0: + grad_residual = grad + else: + grad_residual = grad - exp_avg exp_avg_var.mul_(beta2).addcmul_( grad_residual, grad_residual, value=1 - beta2 )