From 1f9caf2b529e99e2b6faf81c598e5787a2207069 Mon Sep 17 00:00:00 2001 From: shanyu Date: Sat, 26 Sep 2026 02:14:58 +0800 Subject: [PATCH] Keep a gradient second moment when AdaBelief beta1 is 0 With beta1 = 0 the first moment equals the current gradient, so (g - m) is 0 and the second moment stays at eps. The step is then lr * g / sqrt(eps). At eps 1e-16 a scalar quadratic reached -inf on step 10. Adam with beta1 = 0 still tracks g^2. Use that residual only when beta1 is 0. A positive beta1 still stores (1 - beta2) (g - m)^2 plus the eps that is added into the second-moment state. --- tests/test_adabelief_beta1.py | 41 +++++++++++++++++++++++++++++++++++ torch_optimizer/adabelief.py | 10 ++++++++- 2 files changed, 50 insertions(+), 1 deletion(-) create mode 100644 tests/test_adabelief_beta1.py 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 )