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
2 changes: 2 additions & 0 deletions bitsandbytes/optim/lars.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,8 @@ def step(self, closure=None):
update = d_p + buf * momentum
else:
update = buf
else:
update = d_p

update_scale = 1.0
if max_unorm > 0.0:
Expand Down
37 changes: 37 additions & 0 deletions tests/test_optim.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,43 @@ def assert_most_approx_close(a, b, rtol=1e-3, atol=1e-3, max_error_count=0):
torch.testing.assert_close(a, b, rtol=rtol, atol=atol)


def test_pytorch_lars_without_momentum():
parameter = torch.nn.Parameter(torch.tensor([3.0, 4.0]))
optimizer = bnb.optim.PytorchLARS([parameter])
parameter.grad = torch.tensor([4.0, -3.0])

optimizer.step()

torch.testing.assert_close(parameter, torch.tensor([2.9992, 4.0006]))


def test_pytorch_lars_mixed_momentum():
parameters = [torch.nn.Parameter(torch.tensor([1.0, 2.0])), torch.nn.Parameter(torch.tensor([3.0, 4.0]))]
reference_parameters = [torch.nn.Parameter(parameter.detach().clone()) for parameter in parameters]
optimizer = bnb.optim.PytorchLARS(
[{"params": parameters[:1], "momentum": 0.9}, {"params": parameters[1:]}],
weight_decay=0.1,
max_unorm=0.0,
)
reference_optimizer = torch.optim.SGD(
[{"params": reference_parameters[:1], "momentum": 0.9}, {"params": reference_parameters[1:]}],
lr=0.01,
weight_decay=0.1,
)

for direction in (1.0, -1.0):
for index, (parameter, reference_parameter) in enumerate(zip(parameters, reference_parameters)):
gradient = direction * torch.tensor([index + 1.0, index + 2.0])
parameter.grad = gradient.clone()
reference_parameter.grad = gradient.clone()

optimizer.step()
reference_optimizer.step()

for parameter, reference_parameter in zip(parameters, reference_parameters):
torch.testing.assert_close(parameter, reference_parameter)


str2optimizers = {}

## TODO: maybe remove these three.
Expand Down