Skip to content

Add the missing b_hn bias to nnx.GRUCell - #5595

Open
vineethsaivs wants to merge 3 commits into
google:mainfrom
vineethsaivs:fix-nnx-grucell-hidden-bias-20260921
Open

vineethsaivs wants to merge 3 commits into
google:mainfrom
vineethsaivs:fix-nnx-grucell-hidden-bias-20260921

Conversation

@vineethsaivs

Copy link
Copy Markdown

Closes #5594.

nnx.GRUCell's docstring gives the cell as n = tanh(W_in x + b_in + r * (W_hn h + b_hn)), and 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).

Parameter trees for the same cell:

linen:  hn/bias (4,)  hn/kernel  hr/kernel  hz/kernel  in/bias  in/kernel  ir/bias  ir/kernel  iz/bias  iz/kernel
nnx:    dense_h/kernel (4, 12)   dense_i/bias (12,)   dense_i/kernel (3, 12)

A bias on the fused dense_h would also reach r and z, which neither linen nor torch.nn.GRUCell biases, so it is carried as its own parameter instead. That keeps the single fused matmul.

Two new tests in tests/nnx/nn/recurrent_test.py: a linen-parity case built like the existing test_lstm_equivalence_with_flax_linen but with a non-zero bias_init, since the default zeros hide a missing bias, and one asserting the parameter exists and moves the output. Both fail on main with AttributeError: 'GRUCell' object has no attribute 'hn_bias'; the file goes from 20 to 22 passing. ruff check clean. Run on CPU.

Compatibility: this adds one array to the parameter tree, so a checkpoint saved from the current nnx.GRUCell will not restore without a fallback. bias_init defaults to zeros, so a freshly initialised cell is numerically unchanged and the divergence only appears once that bias is trained or weights are ported in. If you would rather not move the tree, say so on #5594 and I will send the docstring change instead.

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.
@samanklesaria

Copy link
Copy Markdown
Collaborator

@vineethsaivs Thanks for catching this! The code looks good, but the tests seem unnecessary here. Can you drop them?

Per review, the change is covered by the existing tests.
@vineethsaivs

Copy link
Copy Markdown
Author

Dropped the test in 50c4db8, so the PR is now just the recurrent.py change.

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.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

nnx.GRUCell is missing the b_hn bias that its docstring and linen.GRUCell both specify

2 participants