Skip to content

Add an RMSNorm layer and fix float16 normalization returning NaN - #159

Merged
jessegrabowski merged 6 commits into
pymc-devs:mainfrom
jessegrabowski:rms-norm
Sep 9, 2026
Merged

jessegrabowski merged 6 commits into
pymc-devs:mainfrom
jessegrabowski:rms-norm

Conversation

@jessegrabowski

@jessegrabowski jessegrabowski commented Sep 9, 2026 •

Copy link
Copy Markdown
Member

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 from LayerNorm. This is F-13 on #53.

_standardize, the helper behind LayerNorm, GroupNorm and BatchNorm, 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.T5LayerNorm and diffusers.RMSNorm to 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/

@jessegrabowski
jessegrabowski merged commit bf1fbc7 into pymc-devs:main Sep 9, 2026
12 checks passed
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.
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.

1 participant