Add an RMSNorm layer and fix float16 normalization returning NaN - #159
Merged
Merged
Conversation
cetagostini
added a commit
to cetagostini/pytensor_ml
that referenced
this pull request
Sep 12, 2026
Retain upstream RMSNorm from pymc-devs#159, replace the positional-only base with VariadicLayer, and keep frequency arithmetic in the input dtype. Simplify RoPE documentation and calls, audit numerical and serialization tests, and exercise general positional tests through the JAX/MLX CI matrix. Record the existing dynamic-head-dimension backend limitations as strict expected failures.
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.
Added
RMSNorm, a layer that rescales activations by the root mean square of their features and applies a learned scale. It does not subtract the mean and has no learned shift, which is what separates it fromLayerNorm. This is F-13 on #53._standardize, the helper behindLayerNorm,GroupNormandBatchNorm, now takes its mean and variance at float32 when the input is float16 or bfloat16, then casts back. Squaring a float16 activation overflows once it reaches about 256 in absolute value, so all three returned NaN for every element on inputs a diffusion decoder produces.Matches
torch.nn.RMSNorm,transformers.T5LayerNormanddiffusers.RMSNormto 5.6e-07. None of the three is a test dependency, so that check is not in the diff.📚 Documentation preview 📚: https://pytensor-ml--159.org.readthedocs.build/en/159/