diff --git a/bitsandbytes/optim/optimizer.py b/bitsandbytes/optim/optimizer.py index dfc6e5d65..7147bd592 100644 --- a/bitsandbytes/optim/optimizer.py +++ b/bitsandbytes/optim/optimizer.py @@ -108,7 +108,10 @@ def override_config(self, parameters, key=None, value=None, key_value_dict=None) if id(p) in self.pid2config: self.pid2config[id(p)].update(key_value_dict) else: - self.pid2config[id(p)] = key_value_dict + # Copy per parameter. Storing the same dict for several parameters means a + # later single-parameter override lands on all of them, and storing a + # caller-supplied `key_value_dict` by reference mutates their dict too. + self.pid2config[id(p)] = dict(key_value_dict) def register_module_override(self, module, param_name, config): self.module_weight_config_triple.append((module, param_name, config)) diff --git a/tests/test_optim.py b/tests/test_optim.py index 29736311d..f6a6423be 100644 --- a/tests/test_optim.py +++ b/tests/test_optim.py @@ -369,6 +369,43 @@ def test_override_config_after_register(device): assert adam.state[p2]["state1"].dtype == torch.uint8 +def test_override_config_does_not_share_one_dict_across_parameters(): + """A per-parameter override must not land on the other parameters of the same call. + + Parameters that had no config yet were all assigned the same dict object, so a later + `override_config(p1, ...)` took the update branch and mutated the config every one of them + was pointing at. + """ + mng = bnb.optim.GlobalOptimManager.get_instance() + mng.initialize() + + p1 = torch.nn.Parameter(torch.zeros(2)) + p2 = torch.nn.Parameter(torch.zeros(2)) + + mng.override_config([p1, p2], "optim_bits", 32) + assert mng.pid2config[id(p1)] is not mng.pid2config[id(p2)] + + mng.override_config(p1, "lr", 0.01) + + assert mng.pid2config[id(p1)] == {"optim_bits": 32, "lr": 0.01} + assert mng.pid2config[id(p2)] == {"optim_bits": 32} + + +def test_override_config_does_not_mutate_the_callers_dict(): + """A `key_value_dict` passed in by the caller is theirs, not the manager's to extend.""" + mng = bnb.optim.GlobalOptimManager.get_instance() + mng.initialize() + + p = torch.nn.Parameter(torch.zeros(2)) + config = {"optim_bits": 8} + + mng.override_config(p, key_value_dict=config) + mng.override_config(p, "lr", 0.5) + + assert config == {"optim_bits": 8} + assert mng.pid2config[id(p)] == {"optim_bits": 8, "lr": 0.5} + + optimizer_names_8bit = [ "adam8bit_blockwise", "lion8bit_blockwise",