From 5480e6d9c1967225083df44c5f38a91a65d6e5f2 Mon Sep 17 00:00:00 2001 From: ErenAta16 Date: Wed, 29 Jul 2026 18:08:26 +0300 Subject: [PATCH] Give each parameter its own config dict in override_config Parameters that had no entry yet were all assigned the same dict object: self.pid2config[id(p)] = key_value_dict so overriding one parameter afterwards silently changed the others registered in the same call: mng.override_config([p1, p2], "optim_bits", 32) mng.pid2config[id(p1)] is mng.pid2config[id(p2)] -> True mng.override_config(p1, "lr", 0.01) p1: {'optim_bits': 32, 'lr': 0.01} p2: {'optim_bits': 32, 'lr': 0.01} <- never asked for The docstring's own example is the single-parameter override, so this is the documented usage rather than an edge case, and nothing raises: the parameter just trains with settings the caller did not choose. The same reference is also handed back out when the caller supplies `key_value_dict` themselves, so their dict grows behind their back: d = {"optim_bits": 8} mng.override_config(p3, key_value_dict=d) mng.override_config(p3, "lr", 0.5) d -> {'optim_bits': 8, 'lr': 0.5} Copy on insert. After: p1: {'optim_bits': 32, 'lr': 0.01} p2: {'optim_bits': 32} d : {'optim_bits': 8} --- bitsandbytes/optim/optimizer.py | 5 ++++- tests/test_optim.py | 37 +++++++++++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 1 deletion(-) 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",