diff --git a/tests/test_sgdw.py b/tests/test_sgdw.py new file mode 100644 index 00000000..a7abcc0e --- /dev/null +++ b/tests/test_sgdw.py @@ -0,0 +1,34 @@ +"""SGDW decay is a scale of the parameter, as in arXiv 1711.05101.""" +import torch +import torch_optimizer as optim + + +def _step(theta, grad, lr, weight_decay, momentum=0.0): + param = torch.nn.Parameter(torch.as_tensor(theta, dtype=torch.float64).clone()) + opt = optim.SGDW( + [param], lr=lr, weight_decay=weight_decay, momentum=momentum + ) + param.grad = torch.as_tensor(grad, dtype=torch.float64).clone() + opt.step() + return param.detach() + + +def test_zero_gradient_scales_the_parameter(): + # θ = 4, μ = 0.1, λ = 0.2. The old update subtracted λμ and landed + # on 3.98. The paper scale (1 − λμ) θ is 3.92. + got = _step([4.0, -2.0], [0.0, 0.0], lr=0.1, weight_decay=0.2) + assert torch.allclose(got, torch.tensor([3.92, -1.96], dtype=torch.float64)) + + +def test_decay_precedes_the_gradient_step(): + # Scale first, then subtract μ g: 4*(1-0.02) - 0.1*1.5 = 3.77. + got = _step([4.0], [1.5], lr=0.1, weight_decay=0.2) + assert torch.allclose(got, torch.tensor([3.77], dtype=torch.float64)) + + +def test_no_decay_matches_plain_sgd(): + theta = torch.tensor([1.5, -0.25, 3.0], dtype=torch.float64) + grad = torch.tensor([0.4, -1.0, 0.0], dtype=torch.float64) + got = _step(theta, grad, lr=0.1, weight_decay=0.0, momentum=0.9) + # one momentum step from an empty buffer stores g and subtracts μ g + assert torch.allclose(got, theta - 0.1 * grad) diff --git a/torch_optimizer/sgdw.py b/torch_optimizer/sgdw.py index 62e7d998..24dd7f7c 100644 --- a/torch_optimizer/sgdw.py +++ b/torch_optimizer/sgdw.py @@ -113,10 +113,11 @@ def step(self, closure: OptLossClosure = None) -> OptFloat: else: d_p = buf - # Apply momentum - p.data.add_(d_p, alpha=-group["lr"]) - - # Apply weight decay + # Decoupled weight decay scales the parameter. The paper + # (Loshchilov and Hutter, arXiv 1711.05101) writes + # θ ← (1 − λ μ) θ, then subtracts the step. Adding −λ μ to + # every coordinate subtracts a constant instead. if weight_decay != 0: - p.data.add_(weight_decay, alpha=-group["lr"]) + p.data.mul_(1 - group["lr"] * weight_decay) + p.data.add_(d_p, alpha=-group["lr"]) return loss