From 80e170e7dff98ec5d3b2062175e80d301c741bb5 Mon Sep 17 00:00:00 2001 From: mottopanikeiku <176798723+mottopanikeiku@users.noreply.github.com> Date: Wed, 9 Sep 2026 23:21:17 -0700 Subject: [PATCH] Fix PytorchLARS updates when momentum is zero (#2079) --- bitsandbytes/optim/lars.py | 2 ++ tests/test_optim.py | 37 +++++++++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+) diff --git a/bitsandbytes/optim/lars.py b/bitsandbytes/optim/lars.py index c2f5aa784..f3ebcb8a0 100644 --- a/bitsandbytes/optim/lars.py +++ b/bitsandbytes/optim/lars.py @@ -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: diff --git a/tests/test_optim.py b/tests/test_optim.py index 29736311d..331c6ac36 100644 --- a/tests/test_optim.py +++ b/tests/test_optim.py @@ -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.