diff --git a/pytensor_ml/layers/__init__.py b/pytensor_ml/layers/__init__.py index 3f6822a..c8ccbb6 100644 --- a/pytensor_ml/layers/__init__.py +++ b/pytensor_ml/layers/__init__.py @@ -41,6 +41,8 @@ LayerNormLayer, NoRunningStatsBatchNormLayer, PredictionBatchNormLayer, + RMSNorm, + RMSNormLayer, ) from pytensor_ml.layers.padding import ( ConstantPad1D, @@ -99,6 +101,7 @@ "MultiheadAttention", "PoolLayer", "PoolLayerGrad", + "RMSNorm", "Recurrent", "RecurrentCell", "ReflectionPad1D", diff --git a/pytensor_ml/layers/norm.py b/pytensor_ml/layers/norm.py index 85774eb..db872c4 100644 --- a/pytensor_ml/layers/norm.py +++ b/pytensor_ml/layers/norm.py @@ -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. @@ -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: @@ -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, @@ -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): @@ -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") diff --git a/tests/dispatch/mlx/test_norm.py b/tests/dispatch/mlx/test_norm.py new file mode 100644 index 0000000..45ccf79 --- /dev/null +++ b/tests/dispatch/mlx/test_norm.py @@ -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) diff --git a/tests/test_layers.py b/tests/test_layers.py index 59d278f..d991bd4 100644 --- a/tests/test_layers.py +++ b/tests/test_layers.py @@ -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 @@ -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 ( @@ -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) diff --git a/tests/test_serialize.py b/tests/test_serialize.py index 2d2fcb0..4ade0fa 100644 --- a/tests/test_serialize.py +++ b/tests/test_serialize.py @@ -39,6 +39,7 @@ GroupNorm, LayerNorm, Linear, + RMSNorm, Sequential, Squeeze, ) @@ -138,6 +139,14 @@ def test_layernorm_roundtrips(affine): assert_outputs_roundtrip([X], output, [np.random.default_rng(1).normal(size=(8, 4))]) +@pytest.mark.parametrize("affine", [True, False], ids=["affine", "no_affine"]) +def test_rmsnorm_roundtrips(affine): + X, output = initialized_network( + Linear("fc", n_in=4, n_out=6), RMSNorm("rms", n_in=6, affine=affine) + ) + assert_outputs_roundtrip([X], output, [np.random.default_rng(1).normal(size=(8, 4))]) + + @pytest.mark.parametrize("affine", [True, False], ids=["affine", "no_affine"]) def test_groupnorm_roundtrips(affine): X, output = initialized_network(