From 82dca579f5e36469df26fbe51ea3c4277277c5e8 Mon Sep 17 00:00:00 2001 From: Caio Theodoro <48462974+caiotheodoro@users.noreply.github.com> Date: Tue, 8 Sep 2026 10:12:53 -0400 Subject: [PATCH] Make AdEMAMix32bit inherit AdEMAMix so it allocates the double state buffer AdEMAMix32bit and PagedAdEMAMix32bit subclassed Optimizer2State directly, so they never used AdEMAMix.init_state and allocated state1 with the shape of the parameter instead of (2, *p.shape). The ademamix kernels read m2 from the second half of state1, so on the CPU and default backends the first step raised RuntimeError, and the t_alpha/t_beta3 schedulers were silently ignored. Subclass AdEMAMix with optim_bits=32, mirroring AdEMAMix8bit, and add a test comparing both classes against AdEMAMix(optim_bits=32). --- bitsandbytes/optim/ademamix.py | 12 +++++----- tests/test_optim.py | 40 ++++++++++++++++++++++++++++++++++ 2 files changed, 45 insertions(+), 7 deletions(-) diff --git a/bitsandbytes/optim/ademamix.py b/bitsandbytes/optim/ademamix.py index 7b79c8e88..92d4dc0c3 100644 --- a/bitsandbytes/optim/ademamix.py +++ b/bitsandbytes/optim/ademamix.py @@ -352,7 +352,7 @@ def __init__( ) -class AdEMAMix32bit(Optimizer2State): +class AdEMAMix32bit(AdEMAMix): def __init__( self, params: Iterable[torch.nn.Parameter], @@ -367,19 +367,17 @@ def __init__( is_paged: bool = False, ): super().__init__( - "ademamix", - params=params, + params, lr=lr, betas=betas, + alpha=alpha, + t_alpha=t_alpha, + t_beta3=t_beta3, eps=eps, weight_decay=weight_decay, optim_bits=32, - args=None, min_8bit_size=min_8bit_size, is_paged=is_paged, - alpha=alpha, - t_alpha=t_alpha, - t_beta3=t_beta3, ) diff --git a/tests/test_optim.py b/tests/test_optim.py index 29736311d..8d0326d44 100644 --- a/tests/test_optim.py +++ b/tests/test_optim.py @@ -575,6 +575,46 @@ def test_benchmark_blockwise(dim1, dim2, gtype, optim_name, device): # assert s < 3.9 +ademamix_32bit_classes = [ + ("AdEMAMix32bit", bnb.optim.AdEMAMix32bit), + ("PagedAdEMAMix32bit", bnb.optim.PagedAdEMAMix32bit), +] + + +@pytest.mark.parametrize("scheduled", [False, True], ids=["unscheduled", "scheduled"]) +@pytest.mark.parametrize( + "optim_name,optim_cls", + ademamix_32bit_classes, + ids=[x[0] for x in ademamix_32bit_classes], +) +@pytest.mark.parametrize("device", get_available_devices()) +@pytest.mark.skipif(not get_available_devices(), reason="No device") +def test_ademamix32bit_matches_ademamix(optim_name, optim_cls, scheduled, device): + """AdEMAMix32bit must allocate the (2, *p.shape) m1/m2 buffer and step exactly like AdEMAMix(optim_bits=32).""" + if device == "cpu" and optim_name.startswith("Paged"): + pytest.skip("Paged optimizers are not meaningful on CPU") + + sched = dict(t_alpha=100, t_beta3=100) if scheduled else {} + + torch.manual_seed(0) + p_ref = torch.nn.Parameter(torch.randn(4096, device=device)) + p_test = torch.nn.Parameter(p_ref.detach().clone()) + opt_ref = bnb.optim.AdEMAMix([p_ref], lr=1e-3, optim_bits=32, **sched) + opt_test = optim_cls([p_test], lr=1e-3, **sched) + + for _ in range(5): + g = torch.randn(4096, device=device) + p_ref.grad = g.clone() + p_test.grad = g.clone() + opt_ref.step() + opt_test.step() + + assert opt_test.state[p_test]["state1"].shape == (2, 4096) + torch.testing.assert_close(opt_test.state[p_test]["state1"], opt_ref.state[p_ref]["state1"]) + torch.testing.assert_close(opt_test.state[p_test]["state2"], opt_ref.state[p_ref]["state2"]) + torch.testing.assert_close(p_test, p_ref) + + ademamix_state_dict_opts = [ ("AdEMAMix8bit", lambda p: bnb.optim.AdEMAMix8bit(p, lr=1e-3)), ("AdEMAMix32bit", lambda p: bnb.optim.AdEMAMix(p, lr=1e-3)),