From f226eeec9ec4bd9cbfa1566ef4a7ce56ab9098f3 Mon Sep 17 00:00:00 2001 From: mm65x Date: Mon, 16 Mar 2026 20:28:14 +0000 Subject: [PATCH] Fix nn.GRU skipping bhn bias when hidden is None (#3252) Co-authored-by: mm65x --- python/mlx/nn/layers/recurrent.py | 2 ++ python/tests/test_nn.py | 8 ++++++++ 2 files changed, 10 insertions(+) diff --git a/python/mlx/nn/layers/recurrent.py b/python/mlx/nn/layers/recurrent.py index 3ffa7654..a5d31fd8 100644 --- a/python/mlx/nn/layers/recurrent.py +++ b/python/mlx/nn/layers/recurrent.py @@ -184,6 +184,8 @@ class GRU(Module): if hidden is not None: n = n + r * h_proj_n + elif self.bhn is not None: + n = n + r * self.bhn n = mx.tanh(n) if hidden is not None: diff --git a/python/tests/test_nn.py b/python/tests/test_nn.py index 174823f1..5957184b 100644 --- a/python/tests/test_nn.py +++ b/python/tests/test_nn.py @@ -1941,6 +1941,14 @@ class TestLayers(mlx_tests.MLXTestCase): h_out = layer(inp, h_out[-1, :]) self.assertEqual(h_out.shape, (44, 12)) + # hidden=None should be equivalent to hidden=zeros (issue #3249) + for bias in [True, False]: + layer = nn.GRU(5, 12, bias=bias) + inp = mx.random.normal((2, 25, 5)) + h_none = layer(inp) + h_zeros = layer(inp, hidden=mx.zeros((2, 12))) + self.assertTrue(mx.allclose(h_none, h_zeros).item()) + def test_lstm(self): layer = nn.LSTM(5, 12) inp = mx.random.normal((2, 25, 5))