Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions pytensor_ml/layers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,8 @@
LayerNormLayer,
NoRunningStatsBatchNormLayer,
PredictionBatchNormLayer,
RMSNorm,
RMSNormLayer,
)
from pytensor_ml.layers.padding import (
ConstantPad1D,
Expand Down Expand Up @@ -99,6 +101,7 @@
"MultiheadAttention",
"PoolLayer",
"PoolLayerGrad",
"RMSNorm",
"Recurrent",
"RecurrentCell",
"ReflectionPad1D",
Expand Down
179 changes: 158 additions & 21 deletions pytensor_ml/layers/norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,23 @@ def _batch_axes(X: pt.TensorVariable) -> tuple[int, ...]:
return tuple(range(X.ndim - 1))


def _accumulator_dtype(X) -> str:
"""The width a normalization's statistics accumulate at.

Squaring a float16 activation overflows past :math:`|x|` of about 256, so the statistics are
taken wider than the input and cast back.
"""
return "float32" if X.dtype in ("float16", "bfloat16") else X.dtype


def _root_mean_square(X, epsilon):
"""Scale ``X`` by its root mean square over the last axis, leaving its mean where it is."""
wide = X.astype(_accumulator_dtype(X))
mean_square = pt.mean(pt.square(wide), axis=-1, keepdims=True)

return (wide / pt.sqrt(mean_square + epsilon)).astype(X.dtype)


def _standardize(X, epsilon, axis, keepdims=False):
"""
Center and scale ``X`` over ``axis``, returning the statistics used.
Expand All @@ -32,9 +49,13 @@ def _standardize(X, epsilon, axis, keepdims=False):
sigma_sq : TensorVariable
Biased variance of ``X`` over ``axis``.
"""
mu = X.mean(axis=axis, keepdims=keepdims)
sigma_sq = X.var(axis=axis, keepdims=keepdims)
return (X - mu) / pt.sqrt(sigma_sq + epsilon), mu, sigma_sq
wide = X.astype(_accumulator_dtype(X))

mu = wide.mean(axis=axis, keepdims=keepdims)
sigma_sq = wide.var(axis=axis, keepdims=keepdims)
standardized = (wide - mu) / pt.sqrt(sigma_sq + epsilon)

return standardized.astype(X.dtype), mu.astype(X.dtype), sigma_sq.astype(X.dtype)


def _affine_input_count(affine: bool) -> int:
Expand Down Expand Up @@ -81,6 +102,25 @@ def _resolve_n_in(name: str, n_in: int | None, X: pt.TensorVariable | None) -> i
return inferred


def _norm_parameter(
name: str, suffix: str, n_in: int, initializer: Initializer | None, default: Initializer
) -> TrainableParameter:
"""Build one of a norm layer's learned vectors, named ``{name}_{suffix}``.

It declares its initializer, so a redraw returns it to the identity transform -- normalizing and then
rescaling by a random factor defeats the point of the layer. A caller who wants something else says so,
and their choice becomes the declaration.
"""
resolved = default if initializer is None else initializer

return trainable(
resolved.initial_value((n_in,)),
f"{name}_{suffix}",
initializer=resolved,
layer_name=name,
)


def _affine_parameters(
name: str,
n_in: int,
Expand All @@ -89,26 +129,11 @@ def _affine_parameters(
) -> tuple[TrainableParameter, TrainableParameter]:
"""Build the learned shift and scale. Returns them in the ``(loc, scale)`` order that every norm
op unpacks its inputs in, so the two cannot drift apart.

Both declare their initializer, so a redraw returns them to the identity transform -- normalizing and
then rescaling by a random factor defeats the point of the layer. A caller who wants something else says
so, and their choice becomes the declaration.
"""
resolved_loc = ZeroInitializer() if loc_initializer is None else loc_initializer
resolved_scale = OneInitializer() if scale_initializer is None else scale_initializer
loc = trainable(
resolved_loc.initial_value((n_in,)),
f"{name}_loc",
initializer=resolved_loc,
layer_name=name,
)
scale = trainable(
resolved_scale.initial_value((n_in,)),
f"{name}_scale",
initializer=resolved_scale,
layer_name=name,
return (
_norm_parameter(name, "loc", n_in, loc_initializer, ZeroInitializer()),
_norm_parameter(name, "scale", n_in, scale_initializer, OneInitializer()),
)
return loc, scale


class BatchNormLayer(LayerOp):
Expand Down Expand Up @@ -456,6 +481,118 @@ def __call__(self, X: pt.TensorLike) -> pt.TensorVariable:
return X_transformed


class RMSNormLayer(UnaryLayerOp):
__props__ = ("n_in", "epsilon", "affine")

def build_inner_graph(self, X, *rest):
normalized = _root_mean_square(X, self.epsilon)
if not self.affine:
return [normalized]

(scale,) = rest

return [normalized * scale]


class RMSNorm(Layer):
r"""
Root-mean-square normalization over the last (feature) axis.

Rescale each sample by the root mean square of its own features, then optionally apply a learned
scale:

.. math::

y = \frac{x}{\sqrt{\mathrm{E}[x^2] + \epsilon}} \cdot \gamma,

where the mean of squares is taken over the last axis. Unlike :class:`LayerNorm` the mean is
left where it is, so the transform is a rescaling rather than a standardization, and there is no
learned shift to go with the scale.

Parameters
----------
name : str, optional
Name used as a prefix for the layer's parameters. Default is "RMSNorm".
n_in : int, optional
Size of the normalized feature axis. Inferred from the input's last dimension on the first
call when omitted.
epsilon : float, optional
Constant :math:`\epsilon` added to the mean square for numerical stability. Default is 1e-6.
affine : bool, optional
Apply the learned scale :math:`\gamma`, starting from :math:`\gamma = 1`, which a redraw
returns it to. Default is True.
scale_initializer : Initializer, optional
How :math:`\gamma` is drawn. Ones when omitted, which is the identity; drawing a random
factor to rescale a normalized activation by would defeat the layer.

Examples
--------
Normalize the queries of an attention head, which is where a transformer puts it to keep the
logits in range:

.. code-block:: python

from pytensor_ml.layers import Input, RMSNorm

queries = Input("queries", shape=(None, 8, 512, 64))
normalized = RMSNorm("q_norm", n_in=64)(queries)
"""

def __init__(
self,
name: str | None = None,
*,
n_in: int | None = None,
epsilon: float = 1e-6,
affine: bool = True,
scale_initializer: Initializer | None = None,
):
self.name = _resolve_layer_name(name, type(self).__name__, "n_in")
self.n_in = n_in
self.epsilon = epsilon
self.affine = affine
self._scale_initializer = scale_initializer

self.scale: TrainableParameter | None = None

self.initialized = False
self._initialize_params(None)

def _initialize_params(self, X: pt.TensorVariable | None):
if self.initialized:
return

n_in = _resolve_n_in(self.name, self.n_in, X)
if n_in is None:
return

if self.affine:
self.scale = _norm_parameter(
self.name, "scale", n_in, self._scale_initializer, OneInitializer()
)

self.initialized = True

def __call__(self, X: pt.TensorLike) -> pt.TensorVariable:
X = pt.as_tensor(X)
self._initialize_params(X)

inputs = [X]
if self.affine:
assert self.scale is not None
inputs.append(self.scale)

X_transformed = RMSNormLayer(
name=self.name,
n_in=self.n_in,
epsilon=self.epsilon,
affine=self.affine,
)(*inputs)
X_transformed.name = f"{self.name}_output"

return X_transformed


class GroupNormLayer(UnaryLayerOp):
__props__ = ("n_in", "n_groups", "epsilon", "affine")

Expand Down
33 changes: 33 additions & 0 deletions tests/dispatch/mlx/test_norm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
import numpy as np
import pytensor
import pytensor.tensor as pt
import pytest

pytest.importorskip("mlx.core")

from pytensor.compile.mode import MLX

from pytensor_ml.layers import GroupNorm, LayerNorm, RMSNorm


@pytest.mark.parametrize(
"build",
[
lambda: GroupNorm("group", n_groups=4, n_in=64, epsilon=1e-6),
lambda: LayerNorm("layer", n_in=64),
lambda: RMSNorm("rms", n_in=64),
],
ids=["group", "layer", "rms"],
)
@pytest.mark.parametrize("scale", [1.0, 3000.0], ids=["small", "past_the_square_root"])
def test_a_norm_survives_float16_activations_it_cannot_square(build, scale):
"""mlx executes float16 natively, so this is where the overflow actually happens: an activation
of 3000 squares to 9e6 against float16's 65504 ceiling. A diffusion decoder reaches those
magnitudes on its way out, so the statistics have to accumulate wider than the input."""
values = (np.random.default_rng(0).normal(size=(2, 16, 64)) * scale).astype("float16")

X = pt.tensor("X", shape=values.shape, dtype="float16")
computed = np.asarray(pytensor.function([X], build()(X), mode=MLX)(values))

assert not np.isnan(computed).any()
np.testing.assert_allclose(computed.astype("float64").std(), 1.0, rtol=0.05)
98 changes: 98 additions & 0 deletions tests/test_layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import pytensor.tensor as pt
import pytest

from pytensor.compile.mode import Mode
from pytensor.gradient import verify_grad
from pytensor.graph.replace import vectorize_graph

Expand All @@ -24,8 +25,10 @@
LayerNorm,
Linear,
MaxPool2D,
RMSNorm,
Sequential,
)
from pytensor_ml.layers.norm import _standardize
from pytensor_ml.layers.recurrent import RecurrentCell
from pytensor_ml.optim import adam
from pytensor_ml.pytensorf import (
Expand Down Expand Up @@ -1012,3 +1015,98 @@ def test_bidirectional_rejects_a_non_string_name():
backward = pytensor_ml.layers.RNN("backward", n_in=2, n_hidden=2)
with pytest.raises(TypeError, match=r"Bidirectional's `name` must be a string"):
pytensor_ml.layers.Bidirectional(forward, backward, name=5)


def test_standardizing_accumulates_wider_than_float16():
"""Squaring a float16 activation overflows past |x| of about 256, and a diffusion decoder's
activations reach the thousands -- the variance goes to infinity and every element standardizes
to NaN, which is a black image rather than a wrong one. Run on the python linker, since the
default backend cannot execute float16 at all."""
values = (np.random.default_rng(0).normal(size=(4, 8)) * 3000).astype("float16")
assert np.abs(values.astype("float64")).max() ** 2 > np.finfo(np.float16).max

X = pt.tensor("X", shape=values.shape, dtype="float16")
standardized, mu, sigma_sq = _standardize(X, 1e-6, axis=0)
computed = pytensor.function([X], standardized, mode=Mode(linker="py", optimizer=None))(values)

assert not np.isnan(computed).any()
assert (computed.dtype, mu.dtype, sigma_sq.dtype) == ("float16",) * 3
np.testing.assert_allclose(computed.astype("float64").std(), 1.0, rtol=1e-3)


def rms_norm_reference(X_np, scale_np, epsilon):
"""Walk the rows one at a time rather than broadcasting the way the layer does, so the two share
nothing but the definition and cannot agree about a wrong axis."""
flat = X_np.reshape(-1, X_np.shape[-1])
expected = np.empty_like(flat)
for row in range(flat.shape[0]):
features = flat[row].astype("float64")
root_mean_square = np.sqrt(sum(value**2 for value in features) / len(features) + epsilon)
expected[row] = features / root_mean_square * scale_np

return expected.reshape(X_np.shape)


@pytest.mark.parametrize("batch_shape", [(10,), (2, 4)], ids=["2d", "3d"])
@pytest.mark.parametrize("n_in", [6, None], ids=["specified", "lazy"])
def test_rms_norm_forward(n_in, batch_shape, rng):
X = pt.tensor("X", shape=(*(None,) * len(batch_shape), 6))
rms_norm = RMSNorm(name="RMSNorm_1", n_in=n_in)
out = rms_norm(X)
assert out.name == "RMSNorm_1_output"

X_np = rng.normal(size=(*batch_shape, 6)).astype(floatX)
scale_np = rng.normal(size=(6,)).astype(floatX)
rms_norm.scale.set_value(scale_np)

res = out.eval({X: X_np})

np.testing.assert_allclose(res, rms_norm_reference(X_np, scale_np, rms_norm.epsilon), rtol=1e-5)


def test_rms_norm_leaves_the_mean_where_it_found_it(rng):
"""The one thing separating it from LayerNorm is that it rescales without centering. An
implementation that subtracted the mean would pass every magnitude check and quietly discard
whatever offset the previous layer encoded."""
X = pt.tensor("X", shape=(None, 8))
out = RMSNorm(name="RMSNorm_1", n_in=8, affine=False)(X)

X_np = rng.normal(loc=3.0, scale=0.1, size=(10, 8)).astype(floatX)
res = out.eval({X: X_np})

# The rows are a tight cloud around 3, so dividing by their root mean square lands them near 1.
# Centering first would land them at 0 instead.
np.testing.assert_allclose(res.mean(axis=-1), 1.0, rtol=0.02)
np.testing.assert_allclose((res**2).mean(axis=-1), 1.0, rtol=1e-3)


def test_rms_norm_applies_epsilon_inside_the_square_root(rng):
# At the default epsilon, dividing by sqrt(mean_square + eps) and by sqrt(mean_square) + eps
# agree to well under the tolerance any other test runs at, so only a large one separates them.
X = pt.tensor("X", shape=(None, 6))
rms_norm = RMSNorm("rms", n_in=6, epsilon=0.5, affine=False)
out = rms_norm(X)

X_np = rng.normal(size=(3, 6)).astype(floatX)

np.testing.assert_allclose(
out.eval({X: X_np}),
rms_norm_reference(X_np, np.ones(6, dtype=floatX), rms_norm.epsilon),
rtol=1e-5,
)


def test_rms_norm_accumulates_wider_than_float16():
"""Squaring is the whole operation here, so a float16 activation of a few thousand overflows
before it is ever averaged. Run on the python linker, since the default backend cannot execute
float16 at all."""
values = (np.random.default_rng(0).normal(size=(4, 8)) * 3000).astype("float16")
assert np.abs(values.astype("float64")).max() ** 2 > np.finfo(np.float16).max

X = pt.tensor("X", shape=values.shape, dtype="float16")
out = RMSNorm(name="RMSNorm_1", n_in=8, affine=False)(X)
computed = pytensor.function([X], out, mode=Mode(linker="py", optimizer=None))(values)

assert not np.isnan(computed).any()
assert computed.dtype == np.dtype("float16")
np.testing.assert_allclose((computed.astype("float64") ** 2).mean(axis=-1), 1.0, rtol=1e-2)
Loading
Loading