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
12 changes: 5 additions & 7 deletions bitsandbytes/optim/ademamix.py
Original file line number Diff line number Diff line change
Expand Up @@ -352,7 +352,7 @@ def __init__(
)


class AdEMAMix32bit(Optimizer2State):
class AdEMAMix32bit(AdEMAMix):
def __init__(
self,
params: Iterable[torch.nn.Parameter],
Expand All @@ -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,
)


Expand Down
40 changes: 40 additions & 0 deletions tests/test_optim.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),
Expand Down