From 776f9dda084323c4339bcb7fe1369cf16c9f7fcb Mon Sep 17 00:00:00 2001 From: Vineeth Sai Date: Mon, 21 Sep 2026 12:18:24 -0700 Subject: [PATCH 1/3] Add the missing b_hn bias to nnx.GRUCell The class docstring gives the cell as n = tanh(W_in x + b_in + r * (W_hn h + b_hn)), and flax.linen.GRUCell builds b_hn as the bias of its `hn` layer. nnx fuses r, z and n into one `dense_h` with use_bias=False, so no hidden bias exists and __call__ computes n = activation_fn(xi_n + r * hh_n). A bias on the fused layer would also reach r and z, so carry it as its own parameter, which is what linen does with a separate `hn` bias. bias_init defaults to zeros, so a freshly initialised cell is unchanged; the divergence only appears once that bias is trained or when weights are ported from linen or torch.nn.GRUCell. --- flax/nnx/nn/recurrent.py | 11 +++++- tests/nnx/nn/recurrent_test.py | 66 ++++++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 1 deletion(-) diff --git a/flax/nnx/nn/recurrent.py b/flax/nnx/nn/recurrent.py index 9df3f065e..9f92bac9b 100644 --- a/flax/nnx/nn/recurrent.py +++ b/flax/nnx/nn/recurrent.py @@ -693,6 +693,15 @@ def __init__( kernel_metadata=recurrent_kernel_metadata, ) + # `b_hn` from the docstring. Only the reset-gated term carries a hidden + # bias, which the fused `dense_h` above cannot express: a bias there would + # reach r and z as well. `flax.linen.GRUCell` builds it as a separate `hn` + # bias for the same reason. + self.hn_bias = nnx.Param( + bias_init(rngs.params(), (hidden_features,), self.param_dtype), + **bias_metadata, + ) + if carry_init: warnings.warn( "carry_init is provided in __init__. " @@ -730,7 +739,7 @@ def __call__(self, carry: Array, inputs: Array) -> tuple[Array, Array]: # type: z = self.gate_fn(xi_z + hh_z) # Compute n with an additional linear transformation on h - n = self.activation_fn(xi_n + r * hh_n) + n = self.activation_fn(xi_n + r * (hh_n + self.hn_bias[...])) # Update hidden state new_h = (1.0 - z) * n + z * h diff --git a/tests/nnx/nn/recurrent_test.py b/tests/nnx/nn/recurrent_test.py index ca8d4cbd1..03526c627 100644 --- a/tests/nnx/nn/recurrent_test.py +++ b/tests/nnx/nn/recurrent_test.py @@ -204,6 +204,72 @@ def test_lstm_equivalence_with_flax_linen(self): np.testing.assert_allclose(c_nnx, c_linen, atol=1e-5) + def test_gru_equivalence_with_flax_linen(self): + """nnx.GRUCell must match flax.linen.GRUCell, including `b_hn`. + + The class docstring specifies + `n = tanh(W_in x + b_in + r * (W_hn h + b_hn))`, and linen builds `b_hn` + as the bias of its `hn` layer. nnx fuses r, z and n into one biasless + `dense_h`, so `b_hn` has to be carried separately; without it the two + implementations disagree as soon as that bias is non-zero. + """ + in_features = 3 + hidden_features = 4 + x = random.normal(random.PRNGKey(42), (1, in_features)) + + rngs_nnx = nnx.Rngs(0) + module_nnx = nnx.GRUCell( + in_features=in_features, + hidden_features=hidden_features, + rngs=rngs_nnx, + ) + carry_nnx = module_nnx.initialize_carry(x.shape, rngs_nnx) + + # A non-zero bias init: the default is zeros, which hides a missing bias. + module_linen = linen.GRUCell( + features=hidden_features, + bias_init=initializers.normal(stddev=1.0), + ) + carry_linen = module_linen.initialize_carry(random.PRNGKey(0), x.shape) + variables_linen = module_linen.init(random.PRNGKey(1), carry_linen, x) + params_linen = variables_linen['params'] + + # nnx splits the fused projections as r, z, n. + module_nnx.dense_i.kernel[...] = jnp.concatenate( + [params_linen[g]['kernel'] for g in ('ir', 'iz', 'in')], axis=-1 + ) + module_nnx.dense_i.bias[...] = jnp.concatenate( + [params_linen[g]['bias'] for g in ('ir', 'iz', 'in')] + ) + module_nnx.dense_h.kernel[...] = jnp.concatenate( + [params_linen[g]['kernel'] for g in ('hr', 'hz', 'hn')], axis=-1 + ) + module_nnx.hn_bias[...] = params_linen['hn']['bias'] + + new_carry_nnx, y_nnx = module_nnx(carry_nnx, x) + new_carry_linen, y_linen = module_linen.apply( + variables_linen, carry_linen, x + ) + + np.testing.assert_allclose(y_nnx, y_linen, atol=1e-5) + np.testing.assert_allclose(new_carry_nnx, new_carry_linen, atol=1e-5) + + def test_gru_has_the_hidden_bias_from_its_docstring(self): + """`b_hn` must exist as a parameter, not only in the documented formula.""" + module = nnx.GRUCell(in_features=3, hidden_features=4, rngs=nnx.Rngs(0)) + self.assertEqual(module.hn_bias.shape, (4,)) + + # It is the only hidden bias: r and z take none, matching linen. + module.hn_bias[...] = jnp.ones((4,)) + x = random.normal(random.PRNGKey(0), (1, 3)) + carry = module.initialize_carry(x.shape, nnx.Rngs(0)) + biased, _ = module(carry, x) + + module.hn_bias[...] = jnp.zeros((4,)) + unbiased, _ = module(carry, x) + self.assertFalse(np.allclose(biased, unbiased)) + + class TestRNN(absltest.TestCase): def test_rnn_with_lstm_cell(self): """Test RNN module using LSTMCell.""" From 50c4db85e78706a003c1d31ada624cda8d073abf Mon Sep 17 00:00:00 2001 From: Vineeth Sai Date: Mon, 28 Sep 2026 15:17:37 -0700 Subject: [PATCH 2/3] Drop the GRUCell equivalence test Per review, the change is covered by the existing tests. --- tests/nnx/nn/recurrent_test.py | 66 ---------------------------------- 1 file changed, 66 deletions(-) diff --git a/tests/nnx/nn/recurrent_test.py b/tests/nnx/nn/recurrent_test.py index 03526c627..ca8d4cbd1 100644 --- a/tests/nnx/nn/recurrent_test.py +++ b/tests/nnx/nn/recurrent_test.py @@ -204,72 +204,6 @@ def test_lstm_equivalence_with_flax_linen(self): np.testing.assert_allclose(c_nnx, c_linen, atol=1e-5) - def test_gru_equivalence_with_flax_linen(self): - """nnx.GRUCell must match flax.linen.GRUCell, including `b_hn`. - - The class docstring specifies - `n = tanh(W_in x + b_in + r * (W_hn h + b_hn))`, and linen builds `b_hn` - as the bias of its `hn` layer. nnx fuses r, z and n into one biasless - `dense_h`, so `b_hn` has to be carried separately; without it the two - implementations disagree as soon as that bias is non-zero. - """ - in_features = 3 - hidden_features = 4 - x = random.normal(random.PRNGKey(42), (1, in_features)) - - rngs_nnx = nnx.Rngs(0) - module_nnx = nnx.GRUCell( - in_features=in_features, - hidden_features=hidden_features, - rngs=rngs_nnx, - ) - carry_nnx = module_nnx.initialize_carry(x.shape, rngs_nnx) - - # A non-zero bias init: the default is zeros, which hides a missing bias. - module_linen = linen.GRUCell( - features=hidden_features, - bias_init=initializers.normal(stddev=1.0), - ) - carry_linen = module_linen.initialize_carry(random.PRNGKey(0), x.shape) - variables_linen = module_linen.init(random.PRNGKey(1), carry_linen, x) - params_linen = variables_linen['params'] - - # nnx splits the fused projections as r, z, n. - module_nnx.dense_i.kernel[...] = jnp.concatenate( - [params_linen[g]['kernel'] for g in ('ir', 'iz', 'in')], axis=-1 - ) - module_nnx.dense_i.bias[...] = jnp.concatenate( - [params_linen[g]['bias'] for g in ('ir', 'iz', 'in')] - ) - module_nnx.dense_h.kernel[...] = jnp.concatenate( - [params_linen[g]['kernel'] for g in ('hr', 'hz', 'hn')], axis=-1 - ) - module_nnx.hn_bias[...] = params_linen['hn']['bias'] - - new_carry_nnx, y_nnx = module_nnx(carry_nnx, x) - new_carry_linen, y_linen = module_linen.apply( - variables_linen, carry_linen, x - ) - - np.testing.assert_allclose(y_nnx, y_linen, atol=1e-5) - np.testing.assert_allclose(new_carry_nnx, new_carry_linen, atol=1e-5) - - def test_gru_has_the_hidden_bias_from_its_docstring(self): - """`b_hn` must exist as a parameter, not only in the documented formula.""" - module = nnx.GRUCell(in_features=3, hidden_features=4, rngs=nnx.Rngs(0)) - self.assertEqual(module.hn_bias.shape, (4,)) - - # It is the only hidden bias: r and z take none, matching linen. - module.hn_bias[...] = jnp.ones((4,)) - x = random.normal(random.PRNGKey(0), (1, 3)) - carry = module.initialize_carry(x.shape, nnx.Rngs(0)) - biased, _ = module(carry, x) - - module.hn_bias[...] = jnp.zeros((4,)) - unbiased, _ = module(carry, x) - self.assertFalse(np.allclose(biased, unbiased)) - - class TestRNN(absltest.TestCase): def test_rnn_with_lstm_cell(self): """Test RNN module using LSTMCell.""" From 998b0bf5095ac082eb6d71806cc1e5400e4078ae Mon Sep 17 00:00:00 2001 From: Vineeth Sai Date: Mon, 5 Oct 2026 23:43:10 -0700 Subject: [PATCH 3/3] Broadcast and promote the GRUCell hn bias like Linear Adding the (hidden,) bias to a batched hh_n relied on implicit rank promotion, which the test suite runs with set to raise, and with dtype set it promoted the carry to param_dtype. Promote it with promote_dtype and reshape it to hh_n's rank, as Linear does with its own bias. --- flax/nnx/nn/recurrent.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/flax/nnx/nn/recurrent.py b/flax/nnx/nn/recurrent.py index 9f92bac9b..a10438b2b 100644 --- a/flax/nnx/nn/recurrent.py +++ b/flax/nnx/nn/recurrent.py @@ -739,7 +739,11 @@ def __call__(self, carry: Array, inputs: Array) -> tuple[Array, Array]: # type: z = self.gate_fn(xi_z + hh_z) # Compute n with an additional linear transformation on h - n = self.activation_fn(xi_n + r * (hh_n + self.hn_bias[...])) + hh_n, hn_bias = self.promote_dtype( + (hh_n, self.hn_bias[...]), dtype=self.dtype + ) + hh_n += jnp.reshape(hn_bias, (1,) * (hh_n.ndim - 1) + (-1,)) + n = self.activation_fn(xi_n + r * hh_n) # Update hidden state new_h = (1.0 - z) * n + z * h