Repository navigation
Add the missing b_hn bias to nnx.GRUCell - #5595
Open
vineethsaivs wants to merge 3 commits into
Open
vineethsaivs wants to merge 3 commits into
vineethsaivs wants to merge 3 commits into
Conversation
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.
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.
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Closes #5594.
nnx.GRUCell's docstring gives the cell asn = tanh(W_in x + b_in + r * (W_hn h + b_hn)), andlinen.GRUCellbuildsb_hnas the bias of itshnlayer. nnx fuses r, z and n into onedense_hwithuse_bias=False, so no hidden bias exists and__call__computesn = activation_fn(xi_n + r * hh_n).Parameter trees for the same cell:
A bias on the fused
dense_hwould also reach r and z, which neither linen nortorch.nn.GRUCellbiases, 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 existingtest_lstm_equivalence_with_flax_linenbut with a non-zerobias_init, since the default zeros hide a missing bias, and one asserting the parameter exists and moves the output. Both fail on main withAttributeError: 'GRUCell' object has no attribute 'hn_bias'; the file goes from 20 to 22 passing.ruff checkclean. Run on CPU.Compatibility: this adds one array to the parameter tree, so a checkpoint saved from the current
nnx.GRUCellwill not restore without a fallback.bias_initdefaults 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.