diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 829e952..e9b33fc 100644 --- a/.github/workflows/run_tests.yml +++ b/.github/workflows/run_tests.yml @@ -20,6 +20,7 @@ jobs: matrix: os: [ubuntu-latest] python-version: ["3.12"] + pytensor-mode: [FAST_RUN] # Off by default, flipped to 1 by the per-backend jobs under `include`. A core job installs no # backend at all, which is what proves the library runs without one. install-jax: [0] @@ -35,6 +36,7 @@ jobs: paths: | tests/test_transformer.py tests/test_layers.py + tests/test_positional.py tests/test_recurrent.py tests/test_conv.py tests/test_conv_transpose.py @@ -64,8 +66,8 @@ jobs: tests/test_workflow_groups.py # Windows runners are roughly twice as slow, and almost everything here is platform independent. # It runs one job over the parts that are not: file IO, paths, and compiling a graph to train. - # One job per backend, each installing only that backend. They live here rather than in - # `test-subset` because a matrix entry cannot change `os`, and mlx needs macOS. + # Backend jobs each install only their backend. Dispatch jobs keep the default mode; positional + # jobs run the same general tests in the backend's mode. These live here because mlx needs macOS. include: - os: ubuntu-latest python-version: "3.12" @@ -79,6 +81,21 @@ jobs: test-subset: name: dispatch-mlx paths: tests/dispatch/mlx/ + - os: ubuntu-latest + python-version: "3.12" + install-jax: 1 + pytensor-mode: JAX + test-subset: + name: positional-jax + paths: tests/test_positional.py + - os: macos-15 + python-version: "3.12" + install-mlx: 1 + pytensor-mode: MLX + floatX: float32 + test-subset: + name: positional-mlx + paths: tests/test_positional.py - os: ubuntu-latest python-version: "3.12" install-torch: 1 @@ -98,6 +115,7 @@ jobs: env: TEST_SUBSET: ${{ matrix.test-subset.paths }} + PYTENSOR_FLAGS: mode=${{ matrix.pytensor-mode || 'FAST_RUN' }},floatX=${{ matrix.floatX || 'float64' }} defaults: run: diff --git a/docs/source/api/layers.rst b/docs/source/api/layers.rst index 391f17c..4208d7c 100644 --- a/docs/source/api/layers.rst +++ b/docs/source/api/layers.rst @@ -10,6 +10,7 @@ Base :toctree: generated/ Layer + pytensor_ml.base.VariadicLayer Combinators ----------- @@ -72,6 +73,7 @@ Normalization and regularization BatchNorm LayerNorm + RMSNorm GroupNorm Dropout @@ -102,3 +104,5 @@ Attention and transformers FeedForward TransformerBlock scaled_dot_product_attention + RotaryEmbedding + rotary_embedding diff --git a/pytensor_ml/base.py b/pytensor_ml/base.py index bb0c277..e773a90 100644 --- a/pytensor_ml/base.py +++ b/pytensor_ml/base.py @@ -85,6 +85,32 @@ def __call__(self, X): def __call__(self, x: pt.TensorLike) -> pt.TensorVariable: ... +class VariadicLayer(ABC): + """ + Base class for graph-building layers that consume multiple tensors. + + Examples + -------- + Declare a layer whose inputs are supplied together: + + .. code-block:: python + + import pytensor.tensor as pt + + from pytensor_ml.base import VariadicLayer + + class Add(VariadicLayer): + def __call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable: + left, right = inputs + return pt.as_tensor(left) + pt.as_tensor(right) + + total = Add()(pt.vector("left"), pt.vector("right")) + """ + + @abstractmethod + def __call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable: ... + + class LayerOp(SymbolicOp): """Base class for the library's neural-network ops. @@ -157,4 +183,4 @@ def update_chain_root(variable: Variable) -> tuple[Variable, int] | None: depth += 1 -__all__ = ["Layer", "LayerOp", "StatefulOp", "UnaryLayerOp", "update_chain_root"] +__all__ = ["Layer", "LayerOp", "StatefulOp", "UnaryLayerOp", "VariadicLayer", "update_chain_root"] diff --git a/pytensor_ml/layers/__init__.py b/pytensor_ml/layers/__init__.py index c8ccbb6..829b806 100644 --- a/pytensor_ml/layers/__init__.py +++ b/pytensor_ml/layers/__init__.py @@ -54,6 +54,7 @@ ZeroPad1D, ZeroPad2D, ) +from pytensor_ml.layers.positional import RotaryEmbedding, RotaryEmbeddingLayer, rotary_embedding from pytensor_ml.layers.recurrent import ( GRU, LSTM, @@ -108,6 +109,7 @@ "ReflectionPad2D", "ReplicationPad1D", "ReplicationPad2D", + "RotaryEmbedding", "Sequential", "Squeeze", "TransformerBlock", @@ -115,5 +117,6 @@ "Upsample2D", "ZeroPad1D", "ZeroPad2D", + "rotary_embedding", "scaled_dot_product_attention", ] diff --git a/pytensor_ml/layers/positional.py b/pytensor_ml/layers/positional.py new file mode 100644 index 0000000..94972d7 --- /dev/null +++ b/pytensor_ml/layers/positional.py @@ -0,0 +1,339 @@ +from typing import Literal, get_args + +import pytensor.tensor as pt + +from pytensor.tensor.type import float_dtypes +from pytensor.tensor.variable import TensorVariable + +from pytensor_ml.base import UnaryLayerOp, VariadicLayer, _resolve_layer_name + +Pairing = Literal["half", "adjacent"] +Scaling = Literal["none", "linear", "ntk"] + +_PAIRINGS = get_args(Pairing) +_SCALINGS = get_args(Scaling) + + +def _validate_options(pairing: str, scaling: str, scaling_factor: float) -> None: + """Reject unknown pairing and scaling options.""" + if pairing not in _PAIRINGS: + raise ValueError(f"pairing must be one of {_PAIRINGS}, got {pairing!r}") + if scaling not in _SCALINGS: + raise ValueError(f"scaling must be one of {_SCALINGS}, got {scaling!r}") + if scaling != "none" and scaling_factor <= 0: + raise ValueError(f"scaling_factor must be positive, got {scaling_factor}") + + +def _head_dim(x: TensorVariable) -> int | TensorVariable: + """Return the feature size, symbolic when it is not statically known.""" + if x.type.dtype not in float_dtypes: + raise ValueError(f"RotaryEmbedding needs a floating-point input, got dtype {x.type.dtype}.") + + static_size = x.type.shape[-1] + if static_size is None: + return x.shape[-1] + if static_size % 2: + raise ValueError( + f"RotaryEmbedding rotates channel pairs, so head_dim must be even, got {static_size}." + ) + + return static_size + + +def _add_head_axes( + angles: TensorVariable, x: TensorVariable, position_ids: TensorVariable +) -> TensorVariable: + """Unsqueeze ``angles`` so its sequence axis lines up with ``x``'s. + + ``position_ids`` indexes tokens, so its last axis is the sequence and any leading axes are batch + axes. ``x`` carries head axes between the two, and the number of them follows from the two ranks -- + which is why the caller does not pass an axis to unsqueeze. + """ + n_head_axes = (x.type.ndim - 1) - position_ids.type.ndim + if n_head_axes < 0: + raise ValueError( + f"position_ids has {position_ids.type.ndim} dimensions, more than the " + f"{x.type.ndim - 1} non-feature dimensions of x." + ) + if n_head_axes == 0: + return angles + + aligned_angles = angles[(Ellipsis, *(None,) * n_head_axes, slice(None), slice(None))] + return aligned_angles + + +def _inverse_frequencies( + head_dim: int | TensorVariable, dtype: str, base: float, scaling: str, scaling_factor: float +) -> TensorVariable: + r""" + Angular frequencies :math:`\theta_i = \mathrm{base}^{-2i/d}`, one per rotated pair. + + Returns + ------- + inverse_frequencies : TensorVariable + Shape ``(head_dim // 2,)``, dtype ``dtype``. Folds to a constant when ``head_dim`` is static. + """ + dimension = pt.cast(head_dim, dtype) + frequency_base = pt.constant(base, dtype=dtype) + if scaling == "ntk": + if isinstance(head_dim, int) and head_dim <= 2: + raise ValueError( + f"NTK scaling rescales the base by scaling_factor ** (d / (d - 2)), which is " + f"undefined for head_dim <= 2; got {head_dim}." + ) + # Static NTK scaling changes the base independently of the sequence length. + factor = pt.constant(scaling_factor, dtype=dtype) + frequency_base = frequency_base * factor ** (dimension / (dimension - 2)) + + exponent = pt.arange(0, head_dim, 2, dtype=dtype) / dimension + inverse_frequencies = frequency_base**-exponent + + if scaling == "linear": + inverse_frequencies = inverse_frequencies / pt.constant(scaling_factor, dtype=dtype) + + inverse_frequencies.name = "inverse_frequencies" + + return inverse_frequencies + + +def _split_pairs( + x: TensorVariable, pairing: str, head_dim: int | TensorVariable +) -> tuple[TensorVariable, TensorVariable]: + """Split the feature axis into the two members of every rotated pair. + + The conventions are a permutation of the feature axis apart and are not interchangeable: weights + trained under one produce nonsense under the other. + + ``"half"`` pairs channel ``i`` with ``i + d/2``; ``"adjacent"`` pairs ``2i`` with ``2i + 1``. + """ + if pairing == "half": + half = head_dim // 2 + first, second = x[..., :half], x[..., half:] + + else: + first, second = x[..., 0::2], x[..., 1::2] + + return first, second + + +def _join_pairs( + x: TensorVariable, first: TensorVariable, second: TensorVariable, pairing: str +) -> TensorVariable: + """Reassemble the feature axis, inverting :func:`_split_pairs` for the same ``pairing``.""" + if pairing == "half": + rotated = pt.concatenate([first, second], axis=-1) + else: + rotated = x[..., 0::2].set(first) + rotated = rotated[..., 1::2].set(second) + return rotated + + +class RotaryEmbeddingLayer(UnaryLayerOp): + __props__ = ("base", "pairing", "scaling", "scaling_factor") + + def build_inner_graph(self, x, position_ids): + _validate_options( + pairing=self.pairing, scaling=self.scaling, scaling_factor=self.scaling_factor + ) + head_dim = _head_dim(x) + dtype = x.type.dtype + + inverse_frequencies = _inverse_frequencies( + head_dim=head_dim, + dtype=dtype, + base=self.base, + scaling=self.scaling, + scaling_factor=self.scaling_factor, + ) + angles = position_ids[..., None].astype(dtype) * inverse_frequencies + angles = _add_head_axes(angles=angles, x=x, position_ids=position_ids) + cos = pt.cos(angles) + sin = pt.sin(angles) + + first, second = _split_pairs(x=x, pairing=self.pairing, head_dim=head_dim) + rotated = _join_pairs( + x=x, + first=first * cos - second * sin, + second=second * cos + first * sin, + pairing=self.pairing, + ) + rotated.name = "rotary_embedding" + + return [rotated] + + +def rotary_embedding( + x: pt.TensorLike, + position_ids: pt.TensorLike, + *, + base: float = 10_000.0, + pairing: Pairing = "half", + scaling: Scaling = "none", + scaling_factor: float = 1.0, +) -> TensorVariable: + r""" + Rotary position embedding (RoPE) applied to the trailing feature axis. + + Rotate each two-dimensional subspace of the feature axis by an angle proportional to the token's + position: + + .. math:: + + \begin{pmatrix} x'_a \\ x'_b \end{pmatrix} = + \begin{pmatrix} \cos m\theta_i & -\sin m\theta_i \\ + \sin m\theta_i & \cos m\theta_i \end{pmatrix} + \begin{pmatrix} x_a \\ x_b \end{pmatrix}, + \qquad \theta_i = \mathrm{base}^{-2i/d}, + + where :math:`m` is the position and :math:`(a, b)` is the :math:`i`-th channel pair [1]_. + Query-key dot products depend only on relative position. + + Apply to queries and keys before + :func:`~pytensor_ml.layers.attention.scaled_dot_product_attention`, never to values. + + Parameters + ---------- + x : TensorLike + Tensor whose last axis is rotated, typically queries or keys of shape + ``(..., n_head, seq, head_dim)``. ``head_dim`` must be even. + For JAX with ``pairing="half"`` and for MLX, declare ``head_dim`` in the input's static shape. + position_ids : TensorLike + Token positions, shape ``(..., seq)``. Head axes are inserted so ``(seq,)`` and + ``(batch, seq)`` both broadcast over ``(batch, n_head, seq, head_dim)``. + base : float, optional + Geometric base of the frequency ladder. Default 10000.0. + pairing : str, optional + Pair channels across halves (``"half"`` [3]_) or consecutively (``"adjacent"`` [4]_). + Must match the trained weights. Default ``"half"``. + scaling : str, optional + Context-extension scheme: ``"none"`` (default), ``"linear"`` for position interpolation [2]_, + or ``"ntk"`` for static NTK-aware scaling. + scaling_factor : float, optional + Context-extension factor, ignored when ``scaling="none"``. Default 1.0. + + Returns + ------- + rotated : TensorVariable + ``x`` with its last axis rotated, same shape and dtype. + + Examples + -------- + Rotate a token at its absolute position in a decoded sequence: + + .. code-block:: python + + import numpy as np + + from pytensor_ml.layers import rotary_embedding + + token = np.ones((1, 8), dtype="float32") + rotated = rotary_embedding(token, position_ids=np.array([12])).eval() + + References + ---------- + .. [1] Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B., & Liu, Y. (2021). RoFormer: Enhanced + Transformer with Rotary Position Embedding. arXiv:2104.09864. https://arxiv.org/abs/2104.09864. + .. [2] Chen, S., Wong, S., Chen, L., & Tian, Y. (2023). Extending Context Window of Large Language + Models via Positional Interpolation. arXiv:2306.15595. https://arxiv.org/abs/2306.15595. + .. [3] https://github.com/huggingface/transformers/blob/v4.57.6/src/transformers/models/llama/modeling_llama.py#L109-L113 + .. [4] https://github.com/meta-pytorch/torchtune/blob/v0.6.1/torchtune/modules/position_embeddings.py#L99-L113 + """ + _validate_options(pairing=pairing, scaling=scaling, scaling_factor=scaling_factor) + + x = pt.as_tensor(x) + position_ids = pt.as_tensor(position_ids) + + rotated = RotaryEmbeddingLayer( + name="RotaryEmbedding", + base=base, + pairing=pairing, + scaling=scaling, + scaling_factor=scaling_factor, + )(x, position_ids) + rotated.name = "rotary_embedding_output" + + return rotated + + +class RotaryEmbedding(VariadicLayer): + r""" + Rotary position embeddings as a configured layer. + + Share one frequency configuration between queries and keys [1]_. Call with ``(x, position_ids)``. + The layer has no learned parameters. + + Parameters + ---------- + name : str or None, optional + Name prefix for the layer's output. Defaults to "RotaryEmbedding" when None. + base : float, optional + Geometric base of the frequency ladder. Default 10000.0. + pairing : str, optional + ``"half"`` (default) or ``"adjacent"``. + scaling : str, optional + ``"none"`` (default), ``"linear"``, or ``"ntk"``. + scaling_factor : float, optional + Extension factor for the scaled variants. Default 1.0. + + Examples + -------- + Apply the same rotation to queries and keys before causal attention: + + .. code-block:: python + + import pytensor.tensor as pt + + from pytensor_ml.layers import Input, RotaryEmbedding, scaled_dot_product_attention + + queries = Input("queries", shape=(None, 4, None, 8)) + keys = Input("keys", shape=(None, 4, None, 8)) + values = Input("values", shape=(None, 4, None, 8)) + positions = pt.lvector("positions") + rope = RotaryEmbedding("rope", pairing="half") + output = scaled_dot_product_attention( + rope(queries, positions), rope(keys, positions), values, is_causal=True + ) + + References + ---------- + .. [1] Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B., & Liu, Y. (2021). RoFormer: Enhanced + Transformer with Rotary Position Embedding. arXiv:2104.09864. https://arxiv.org/abs/2104.09864. + """ + + def __init__( + self, + name: str | None = None, + *, + base: float = 10_000.0, + pairing: Pairing = "half", + scaling: Scaling = "none", + scaling_factor: float = 1.0, + ): + _validate_options(pairing=pairing, scaling=scaling, scaling_factor=scaling_factor) + + self.name = _resolve_layer_name(name, "RotaryEmbedding", "base") + self.base = base + self.pairing = pairing + self.scaling = scaling + self.scaling_factor = scaling_factor + + def __call__(self, *inputs: pt.TensorLike) -> TensorVariable: + """Rotate ``(x, position_ids)`` using this layer's configuration.""" + x, position_ids = inputs + rotated = rotary_embedding( + x=x, + position_ids=position_ids, + base=self.base, + pairing=self.pairing, + scaling=self.scaling, + scaling_factor=self.scaling_factor, + ) + rotated.name = f"{self.name}_output" + + return rotated + + +__all__ = [ + "RotaryEmbedding", + "rotary_embedding", +] diff --git a/tests/dispatch/jax/test_basic.py b/tests/dispatch/jax/test_basic.py index 851aff5..6568729 100644 --- a/tests/dispatch/jax/test_basic.py +++ b/tests/dispatch/jax/test_basic.py @@ -7,13 +7,12 @@ from pytensor.compile.mode import JAX, Mode from pytensor.graph.basic import Variable -from pytensor.graph.rewriting.db import RewriteDatabaseQuery from pytensor.link.jax.linker import JAXLinker jax = pytest.importorskip("jax") -optimizer = RewriteDatabaseQuery(include=["jax"], exclude=JAX._optimizer.exclude) -jax_mode = Mode(linker=JAXLinker(), optimizer=optimizer) +# Include canonicalization for ops that lower to supported primitives before JAX dispatch. +jax_mode = Mode(linker=JAXLinker(), optimizer=JAX._optimizer) py_mode = Mode(linker="py", optimizer=None) diff --git a/tests/test_positional.py b/tests/test_positional.py new file mode 100644 index 0000000..285a6f1 --- /dev/null +++ b/tests/test_positional.py @@ -0,0 +1,288 @@ +import numpy as np +import pytensor +import pytensor.tensor as pt +import pytest + +from pytensor_ml.layers.attention import scaled_dot_product_attention +from pytensor_ml.layers.positional import RotaryEmbedding, rotary_embedding +from tests.test_attention import sdpa_np + +floatX = pytensor.config.floatX + + +@pytest.fixture(scope="module") +def rng(): + return np.random.default_rng(sum(map(ord, "Rotary Test"))) + + +def rope_np(x, positions, *, base=10_000.0, pairing="half", scaling="none", scaling_factor=1.0): + """ + Independent RoPE reference, built pair by pair from explicit 2x2 rotations. + + Deliberately written from the RoFormer definition rather than from the implementation under test: + it imports no pytensor and no pytensor_ml, indexes scalars instead of slicing halves, and applies + the rotation matrix literally. A shared mistake in the pairing, a sign, or the frequency ladder + therefore cannot cancel out between the two. + + Parameters + ---------- + x : ndarray + Shape ``(..., seq, head_dim)``. + positions : ndarray + Shape ``(seq,)``, shared by every leading axis of ``x``. + """ + x = np.asarray(x, dtype="float64") + positions = np.asarray(positions) + head_dim = x.shape[-1] + half = head_dim // 2 + + if scaling == "ntk": + base = base * scaling_factor ** (head_dim / (head_dim - 2)) + + theta = np.array([base ** (-2.0 * i / head_dim) for i in range(half)]) + if scaling == "linear": + theta = theta / scaling_factor + + out = np.empty_like(x) + for lead in np.ndindex(*x.shape[:-2]): + for step, position in enumerate(positions): + for i in range(half): + first, second = (i, i + half) if pairing == "half" else (2 * i, 2 * i + 1) + cos, sin = np.cos(position * theta[i]), np.sin(position * theta[i]) + x_first, x_second = x[(*lead, step, first)], x[(*lead, step, second)] + out[(*lead, step, first)] = x_first * cos - x_second * sin + out[(*lead, step, second)] = x_second * cos + x_first * sin + + return out + + +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +@pytest.mark.parametrize( + "scaling, scaling_factor", + [("none", 1.0), ("linear", 4.0), ("ntk", 4.0)], + ids=["unscaled", "linear", "ntk"], +) +def test_rope_matches_independent_reference(pairing, scaling, scaling_factor, rng): + x_np = rng.normal(size=(2, 3, 5, 8)).astype(floatX) + positions_np = np.arange(5) + + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + out = rotary_embedding( + x, positions, pairing=pairing, scaling=scaling, scaling_factor=scaling_factor + ) + + result = out.eval({x: x_np, positions: positions_np}) + expected = rope_np( + x_np, positions_np, pairing=pairing, scaling=scaling, scaling_factor=scaling_factor + ) + np.testing.assert_allclose(result, expected, atol=1e-6) + + +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +def test_rope_one_step_at_a_time_equals_whole_sequence(pairing, rng): + x_np = rng.normal(size=(2, 3, 6, 4)).astype(floatX) + positions_np = np.arange(6) + + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + whole = rotary_embedding(x, positions, pairing=pairing).eval({x: x_np, positions: positions_np}) + + step_x = pt.tensor("step_x", shape=(2, 3, 1, 4)) + step_out = rotary_embedding(step_x, positions, pairing=pairing) + for step, position in enumerate(positions_np): + one = step_out.eval({step_x: x_np[:, :, step : step + 1], positions: np.array([position])}) + np.testing.assert_allclose(one[:, :, 0], whole[:, :, step], atol=1e-6) + + +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +def test_rope_scores_depend_only_on_relative_position(pairing, rng): + q_np = rng.normal(size=(1, 1, 1, 8)).astype(floatX) + k_np = rng.normal(size=(1, 1, 1, 8)).astype(floatX) + + q, k = pt.tensor("q", shape=q_np.shape), pt.tensor("k", shape=k_np.shape) + positions = pt.lvector("positions") + q_out = rotary_embedding(q, positions, pairing=pairing) + k_out = rotary_embedding(k, positions, pairing=pairing) + + def score(query_position, key_position): + rotated_q = q_out.eval({q: q_np, positions: np.array([query_position])}) + rotated_k = k_out.eval({k: k_np, positions: np.array([key_position])}) + return float((rotated_q * rotated_k).sum()) + + np.testing.assert_allclose(score(3, 0), score(10, 7), rtol=1e-6) + np.testing.assert_allclose(score(3, 0), score(103, 100), rtol=1e-6) + assert not np.isclose(score(3, 0), score(4, 0)) + + +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +def test_rope_is_an_orthogonal_rotation(pairing, rng): + x_np = rng.normal(size=(2, 4, 6)).astype(floatX) + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + out = rotary_embedding(x, positions, pairing=pairing) + + rotated = out.eval({x: x_np, positions: np.arange(4)}) + np.testing.assert_allclose( + np.linalg.norm(rotated, axis=-1), np.linalg.norm(x_np, axis=-1), rtol=1e-6 + ) + + identity = out.eval({x: x_np, positions: np.zeros(4, dtype="int64")}) + np.testing.assert_allclose(identity, x_np, atol=1e-6) + + +def test_rope_positions_broadcast_shared_and_per_sequence(rng): + x_np = rng.normal(size=(2, 3, 5, 8)).astype(floatX) + x = pt.tensor("x", shape=x_np.shape) + + shared = pt.lvector("shared") + per_sequence = pt.lmatrix("per_sequence") + shared_out = rotary_embedding(x, shared) + batched_out = rotary_embedding(x, per_sequence) + + positions_np = np.arange(5) + shared_result = shared_out.eval({x: x_np, shared: positions_np}) + repeated = batched_out.eval({x: x_np, per_sequence: np.tile(positions_np, (2, 1))}) + np.testing.assert_allclose(shared_result, repeated, atol=1e-6) + + offsets = np.stack([positions_np, positions_np + 100]) + offset_result = batched_out.eval({x: x_np, per_sequence: offsets}) + np.testing.assert_allclose(offset_result[0], shared_result[0], atol=1e-6) + np.testing.assert_allclose(offset_result[1], rope_np(x_np[1], positions_np + 100), atol=1e-6) + + +@pytest.mark.parametrize("scaling", ["linear", "ntk"]) +def test_rope_scaling_factor_one_is_the_identity(scaling, rng): + x_np = rng.normal(size=(3, 6)).astype(floatX) + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + positions_np = np.arange(3) + + unscaled = rotary_embedding(x, positions).eval({x: x_np, positions: positions_np}) + scaled = rotary_embedding(x, positions, scaling=scaling, scaling_factor=1.0).eval( + {x: x_np, positions: positions_np} + ) + np.testing.assert_allclose(unscaled, scaled, atol=1e-6) + + +def test_linear_scaling_interpolates_positions(rng): + x_np = rng.normal(size=(1, 6)).astype(floatX) + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + + interpolated = rotary_embedding(x, positions, scaling="linear", scaling_factor=4.0).eval( + {x: x_np, positions: np.array([8])} + ) + plain = rotary_embedding(x, positions).eval({x: x_np, positions: np.array([2])}) + np.testing.assert_allclose(interpolated, plain, atol=1e-6) + + +def test_rope_composes_with_attention(rng): + q_np, k_np, v_np = (rng.normal(size=(2, 3, 5, 4)).astype(floatX) for _ in range(3)) + q = pt.tensor("q", shape=q_np.shape) + k = pt.tensor("k", shape=k_np.shape) + v = pt.tensor("v", shape=v_np.shape) + positions = pt.lvector("positions") + positions_np = np.arange(5) + options = dict(base=500.0, pairing="adjacent", scaling="linear", scaling_factor=2.5) + + rope = RotaryEmbedding("rope", **options) + output = scaled_dot_product_attention(rope(q, positions), rope(k, positions), v, is_causal=True) + result = output.eval({q: q_np, k: k_np, v: v_np, positions: positions_np}) + expected = sdpa_np( + rope_np(q_np, positions_np, **options), + rope_np(k_np, positions_np, **options), + v_np, + is_causal=True, + ) + np.testing.assert_allclose(result, expected, atol=1e-6) + + +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +def test_rope_input_gradient_matches_transpose_rotation(pairing, rng): + """Backpropagation applies the inverse rotation to the output cotangent.""" + x_np = rng.normal(size=(2, 3, 4)).astype(floatX) + weights = rng.normal(size=x_np.shape).astype(floatX) + positions_np = np.array([1, 7, 19]) + x = pt.tensor("x", shape=x_np.shape) + positions = pt.lvector("positions") + out = rotary_embedding(x, positions, pairing=pairing) + + grad = pt.grad(((out * weights) ** 2).sum(), x).eval({x: x_np, positions: positions_np}) + cotangent = 2 * weights**2 * rope_np(x_np, positions_np, pairing=pairing) + expected = rope_np(cotangent, -positions_np, pairing=pairing) + np.testing.assert_allclose(grad, expected, rtol=1e-6, atol=1e-6) + + +@pytest.mark.parametrize( + "pairing", + [ + pytest.param( + "half", + marks=pytest.mark.xfail( + condition=pytensor.config.mode == "JAX", + reason="PyTensor's JAX lowering requires static slice bounds for the half split", + raises=IndexError, + strict=True, + ), + ), + "adjacent", + ], +) +@pytest.mark.xfail( + condition=pytensor.config.mode == "MLX", + reason="PyTensor's MLX arange lowering requires a statically known head dimension", + raises=NotImplementedError, + strict=True, +) +def test_unknown_head_dimension_is_supported(pairing, rng): + x_np = rng.normal(size=(3, 8)).astype(floatX) + positions_np = np.arange(3) + + x = pt.tensor("x", shape=(3, None)) + positions = pt.lvector("positions") + out = rotary_embedding(x, positions, pairing=pairing) + + np.testing.assert_allclose( + out.eval({x: x_np, positions: positions_np}), + rope_np(x_np, positions_np, pairing=pairing), + atol=1e-6, + ) + + +def test_odd_head_dimension_raises(): + with pytest.raises(ValueError, match="head_dim must be even"): + rotary_embedding(pt.tensor("x", shape=(4, 7)), pt.lvector("positions")) + + +def test_integer_input_raises(): + with pytest.raises(ValueError, match="floating-point input"): + rotary_embedding(pt.tensor("x", shape=(4, 8), dtype="int64"), pt.lvector("positions")) + + +def test_positions_wider_than_input_raises(): + with pytest.raises(ValueError, match="more than the 1 non-feature dimensions"): + rotary_embedding(pt.tensor("x", shape=(4, 8)), pt.lmatrix("positions")) + + +@pytest.mark.parametrize( + "kwargs, match", + [ + ({"pairing": "interleaved"}, "pairing must be one of"), + ({"scaling": "yarn"}, "scaling must be one of"), + ({"scaling": "linear", "scaling_factor": 0.0}, "scaling_factor must be positive"), + ], + ids=["bad_pairing", "bad_scaling", "bad_factor"], +) +def test_invalid_options_raise(kwargs, match): + with pytest.raises(ValueError, match=match): + rotary_embedding(pt.tensor("x", shape=(4, 8)), pt.lvector("positions"), **kwargs) + with pytest.raises(ValueError, match=match): + RotaryEmbedding("rope", **kwargs) + + +def test_ntk_scaling_needs_more_than_two_channels(): + with pytest.raises(ValueError, match="undefined for head_dim <= 2"): + rotary_embedding( + pt.tensor("x", shape=(4, 2)), pt.lvector("positions"), scaling="ntk", scaling_factor=2.0 + ) diff --git a/tests/test_serialize.py b/tests/test_serialize.py index 4ade0fa..8c40638 100644 --- a/tests/test_serialize.py +++ b/tests/test_serialize.py @@ -44,6 +44,7 @@ Squeeze, ) from pytensor_ml.layers.attention import scaled_dot_product_attention +from pytensor_ml.layers.positional import rotary_embedding from pytensor_ml.pytensorf import collect_shared_variables, collect_trainable_params from pytensor_ml.serialize.base import _TYPE_FROM_JSON, _TYPE_TO_JSON @@ -155,6 +156,23 @@ def test_groupnorm_roundtrips(affine): assert_outputs_roundtrip([X], output, [np.random.default_rng(1).normal(size=(8, 4))]) +@pytest.mark.parametrize("pairing", ["half", "adjacent"]) +@pytest.mark.parametrize( + "scaling, scaling_factor", + [("none", 1.0), ("linear", 4.0), ("ntk", 4.0)], + ids=["unscaled", "linear", "ntk"], +) +def test_rotary_embedding_roundtrips(pairing, scaling, scaling_factor): + """A restored graph preserves non-default rotation frequencies and pairing.""" + x = pt.tensor("x", shape=(2, 3, 5, 8)) + positions = pt.lvector("positions") + output = rotary_embedding( + x, positions, base=500.0, pairing=pairing, scaling=scaling, scaling_factor=scaling_factor + ) + values = [np.random.default_rng(0).normal(size=(2, 3, 5, 8)), np.arange(5)] + assert_outputs_roundtrip([x, positions], output, values) + + def test_squeeze_roundtrips(): X = pt.matrix("X") assert_outputs_roundtrip(