Fix non-strict module update with extra weights (#3214)

Co-authored-by: Michelle DiMarco <[email protected]>
This commit is contained in:
Michelle DiMarco
2026-03-09 22:03:50 -07:00
committed by GitHub
co-authored by Michelle DiMarco
parent 6ac5280db4
commit d2702a4fc1
2 changed files with 24 additions and 0 deletions
+7
View File
@@ -340,6 +340,13 @@ class Module(dict):
raise ValueError(f'Module does not have parameter named "{k}".')
elif isinstance(parameters, list):
for i in range(len(parameters)):
if i >= len(dst):
if strict:
raise ValueError(
f"List index {i} is out of bounds for "
f"destination of length {len(dst)}."
)
continue
current_value = dst[i]
new_value = parameters[i]
if isinstance(current_value, mx.array):
+17
View File
@@ -170,6 +170,23 @@ class TestBase(mlx_tests.MLXTestCase):
# Empty weights is ok if strict is false
m.load_weights([], strict=False)
# Extra weights for non-existent layers are filtered when strict
# is false. Flat keys like "extra.weight" are silently dropped by
# Module.update, but nested indexed keys like "layers.1.weight"
# cause an IndexError in tree_unflatten/update without filtering.
m = nn.Sequential(nn.Linear(2, 2))
m.load_weights(
[
("layers.0.weight", mx.ones((2, 2))),
("layers.0.bias", mx.ones((2,))),
("layers.1.weight", mx.ones((2, 2))),
("layers.1.bias", mx.ones((2,))),
],
strict=False,
)
self.assertTrue(mx.array_equal(m.layers[0].weight, mx.ones((2, 2))))
self.assertEqual(len(m.layers), 1)
def test_module_state(self):
m = nn.Linear(10, 1)
m.state["hello"] = "world"