Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions tests/test_lamb_debias.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
"""debias=True corrects the moments inside the trust ratio."""
import torch
import torch_optimizer as optim


def test_default_debias_is_the_uncorrected_step():
param = torch.nn.Parameter(torch.tensor([2.0, -1.0]))
opt = optim.Lamb([param], lr=0.1, weight_decay=0.0, betas=(0.9, 0.999))
grad = torch.tensor([1.0, 0.5])
param.grad = grad.clone()
opt.step()
# m = (1-0.9) g, v = (1-0.999) g^2, step = m / sqrt(v)
direction = (0.1 * grad) / (0.001 * grad * grad).sqrt()
weight_norm = torch.tensor([2.0, -1.0]).norm()
trust = weight_norm / direction.norm()
expected = torch.tensor([2.0, -1.0]) - 0.1 * trust * direction
assert torch.allclose(param.detach(), expected, atol=1e-5)


def test_debias_corrects_adam_norm():
# One step, scalar parameter. Stored moments must stay the raw EMA.
param = torch.nn.Parameter(torch.tensor([2.0]))
opt = optim.Lamb(
[param], lr=0.1, weight_decay=0.1, betas=(0.9, 0.999), debias=True
)
param.grad = torch.tensor([1.0])
opt.step()
state = opt.state[param]
assert torch.allclose(state["exp_avg"], torch.tensor([0.1]))
assert torch.allclose(state["exp_avg_sq"], torch.tensor([0.001]))
m_hat = 0.1 / (1 - 0.9)
v_hat = 0.001 / (1 - 0.999)
direction = m_hat / (v_hat ** 0.5) + 0.1 * 2.0
# scalar norms are absolute values; trust ratio is |p| / |direction|
updated = 2.0 - 0.1 * (2.0 / abs(direction)) * direction
assert torch.allclose(param.detach(), torch.tensor([updated]), atol=1e-5)
# Scaling the learning rate and leaving adam_norm uncorrected gives
# about 1.937 on this step, not the corrected 1.8.
assert abs(param.detach().item() - 1.936754) > 0.05
20 changes: 11 additions & 9 deletions torch_optimizer/lamb.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
import math

import torch
from torch.optim.optimizer import Optimizer

Expand Down Expand Up @@ -126,19 +124,23 @@ def step(self, closure: OptLossClosure = None) -> OptFloat:
# v_t
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)

# Paper v3 does not use debiasing.
# Paper v3 does not use debiasing. When it is requested, Adam's
# correction is applied to the moments that enter the trust
# ratio. Scaling the learning rate instead leaves adam_norm
# uncorrected and also scales the weight-decay term (issue 443).
# The stored moments are not divided in place.
if self.debias:
bias_correction = math.sqrt(1 - beta2 ** state["step"])
bias_correction /= 1 - beta1 ** state["step"]
exp_avg_hat = exp_avg / (1 - beta1 ** state["step"])
exp_avg_sq_hat = exp_avg_sq / (1 - beta2 ** state["step"])
else:
bias_correction = 1
exp_avg_hat = exp_avg
exp_avg_sq_hat = exp_avg_sq

# Apply bias to lr to avoid broadcast.
step_size = group["lr"] * bias_correction
step_size = group["lr"]

weight_norm = torch.norm(p.data).clamp(0, self.clamp_value)

adam_step = exp_avg / exp_avg_sq.sqrt().add(group["eps"])
adam_step = exp_avg_hat / exp_avg_sq_hat.sqrt().add(group["eps"])
if group["weight_decay"] != 0:
adam_step.add_(p.data, alpha=group["weight_decay"])

Expand Down