Skip to content

Add rotary position embeddings - #55

Open
cetagostini wants to merge 4 commits into
pymc-devs:mainfrom
cetagostini:agent/rmsnorm-rope
Open

cetagostini wants to merge 4 commits into
pymc-devs:mainfrom
cetagostini:agent/rmsnorm-rope

Conversation

@cetagostini

@cetagostini cetagostini commented Aug 7, 2026 •

Copy link
Copy Markdown

Add rotary position embeddings for query/key tensors.

Merged current main (bf1fbc7, including #159). RMSNorm is now upstream, so this PR retains upstream's implementation and tests unchanged rather than introducing a competing version.

Scope

  • rotary_embedding(x, position_ids, ...) and configured RotaryEmbedding layer: explicit positions, half/adjacent pairing, linear interpolation, and static NTK scaling.
  • General VariadicLayer base with __call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable, replacing PositionalLayer. Existing unary Layer remains unchanged. UnaryLayerOp already accepts multiple inputs and names its single-output contract.
  • Frequency arithmetic stays in the input dtype, including symbolic head dimensions and NTK scaling.
  • API documentation and runnable examples. No custom backend dispatch or zip_concatenate primitive.

Review follow-up

Simplified docstrings, named returned values, split cosine/sine assignments, and used keyword arguments with one expression per line. Audited the added tests: removed parameter/name checks, prediction-rewrite tautologies, duplicate RMSNorm serialization coverage, and dedicated JAX positional tests. Retained independent numerical references, decode and relative-position invariants, broadcasting, scaling, errors, and serialization. Attention composition and gradients now use independent numerical expectations.

CI runs the same general RoPE tests on native CPU, JAX, and MLX, with float32 on MLX. The existing JAX harness uses the complete optimizer, including canonicalization. The earlier claim that SplitDims/JoinDims cannot compile on JAX was incorrect.

Verification

  • Native integration suite: 271 passed (positional, layers, serialization, base, workflow groups, and JAX dispatch tests).
  • JAX RoPE suite: 27 passed, 1 strict expected failure.
  • MLX RoPE suite: 26 passed, 2 strict expected failures.
  • mypy: no errors; repository pre-commit checks passed.
  • Three documentation examples executed. RMSNorm → RoPE → causal attention smoke exercised native CPU, JAX, and MLX; token-by-token rotations match prefill.

The expected failures record PyTensor 3.3.1 backend limitations: MLX requires a static head dimension for arange, and JAX's half pairing requires static slice bounds. Native CPU supports unknown dimensions for both pairings; JAX supports the tested unknown-dimension adjacent pairing. The limitations are documented and strict expected failures will flag an upstream fix.


📚 Documentation preview 📚: https://pytensor-ml--55.org.readthedocs.build/en/55/

Two additive decoder primitives, from the in-scope set agreed on pymc-devs#53. Both are
`UnaryLayerOp`s with complete `__props__`, so they serialize as leaves for free,
and neither touches the `__props__` of an already-serialized op -- no
`GRAPH_FORMAT_VERSION` bump.

RMSNorm (`RMSNorm` / `RMSNormLayer`, beside `LayerNorm`): scale-only affine with
no shift and no mean subtraction, matching `torch.nn.RMSNorm`,
`flax.linen/nnx.RMSNorm` and `tinygrad.nn.RMSNorm`. Default epsilon is 1e-6
(flax/tinygrad), deliberately not `LayerNorm`'s 1e-5 (torch).

RoPE (`rotary_embedding` / `RotaryEmbedding` / `RotaryEmbeddingLayer` in a new
`layers/positional.py`): positions are an explicit input rather than an implied
`0..seq-1`, so one graph serves a full sequence and a single decode step. Both
pairing conventions are props, because they are not interchangeable and the
choice is fixed by the weights: `"half"` pairs `i` with `i + d/2` (HuggingFace
`rotate_half`, GPT-NeoX, flax) and `"adjacent"` pairs `2i` with `2i + 1`
(RoFormer, GPT-J, torchtune). Linear position interpolation and static NTK-aware
scaling are also props. Composes with the existing `AttentionLayer` unchanged --
`_sdpa_graph` is already position-agnostic.

`_constant_like` moves from `activations.py` to `base.py` as `constant_like`:
both the activations and the new norm need dtype-pinned scalar constants, and
`base.py` is already the shared root neither can cycle through. Without it an
epsilon of 1e-6 widens a float32 input to float64, since pytensor's autocaster
types a bare Python float by value.

Correctness is pinned against an independent scalar reference in
`tests/test_positional.py` that imports neither pytensor nor pytensor_ml and
applies each 2x2 rotation literally: flipping the pairing on the reference side
fails the test. Also covered: one-step-at-a-time decoding equals the
whole-sequence result, scores depend only on relative position, rotation is
orthogonal, position 0 is the identity, and linear scaling by 4 at position 8
equals plain RoPE at position 2.

Both ops lower on JAX through pytensor's `OpFromGraph` fallback with no
`jax_funcify` of their own, asserted in `tests/dispatch/jax/test_positional.py`.
The interleaved pairing therefore reassembles its output with strided
`set_subtensor` writes rather than `split_dims`/`join_dims`, which have no JAX or
MLX conversion.
@codecov-commenter

Copy link
Copy Markdown

Welcome to Codecov 🎉

Once you merge this PR into your default branch, you're all set! Codecov will compare coverage reports and display results in all future pull requests.

Thanks for integrating Codecov - We've got you covered ☂️

Comment thread pytensor_ml/layers/norm.py Outdated
return (X - mu) / pt.sqrt(sigma_sq + epsilon), mu, sigma_sq


def _rms_normalize(X, epsilon):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is only called once, just inline it

Comment thread pytensor_ml/layers/norm.py Outdated
return inferred


def _scale_parameter(name: str, n_in: int) -> TrainableParameter:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't like this, it makes _affine_parameters asymmetric. At this point just inline the variable creation everywhere.

Comment thread pytensor_ml/layers/norm.py Outdated
if not self.affine:
return [X_normalized]

# Scale-only, so the affine transform contributes one input rather than the ``(loc, scale)``

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

remove

Comment thread pytensor_ml/layers/norm.py Outdated

y = \frac{x}{\sqrt{\frac{1}{n} \sum_i x_i^2 + \epsilon}} \cdot \gamma.

Unlike :class:`LayerNorm` there is no mean subtraction and no learned shift. That is not a

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Simplify this, it's weirdly obsessed with the fact that there's no location, so what. Nice to credit who we're copying and where it's used, though.

Comment thread pytensor_ml/layers/norm.py Outdated
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,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

simplify, don't need the huge essay

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

also i'd rather we just pick something for our library

Comment thread pytensor_ml/layers/positional.py Outdated
the feature axis apart, so this must match whatever the weights were trained with.
scaling : str, optional
Context-extension scheme. ``"none"`` (default) for plain RoPE, ``"linear"`` for position
interpolation (Chen et al. 2023), or ``"ntk"`` for static NTK-aware scaling, which rescales the

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

needs References section and correct cross-reference for academic paper

Apply this to queries and keys before :func:`~pytensor_ml.layers.attention.scaled_dot_product_attention`,
never to values.

Parameters

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This all needs to be dialed in, each parameter has the book of genesis under it

Comment thread pytensor_ml/layers/positional.py Outdated
scaling_factor : float, optional
Extension factor for the scaled variants. Default is 1.0.

See Also

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't personally care for the see also section, ymmv

Comment thread pytensor_ml/layers/positional.py Outdated
self.scaling = scaling
self.scaling_factor = scaling_factor

# Positions are a required second input, which the one-tensor Layer.__call__ signature does not

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

So create a new subclass

Comment thread pytensor_ml/activations.py Outdated
# np.finfo maps complex64 -> float32, keeping complex inputs at their own precision.
dtype = np.finfo(dtype).dtype if np.issubdtype(dtype, np.inexact) else np.dtype(config.floatX)
return pt.constant(np.asarray(value, dtype=dtype))
from pytensor_ml.base import Layer, constant_like

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

these changes look like they were mixed in from other PRs

Revert the unrelated churn: `activations.py` is untouched again and the
`constant_like` move is gone. RMSNorm pins its epsilon with the core one-liner
`pt.constant(self.epsilon, dtype=X.type.dtype)` instead.

`base.py` now only gains `PositionalLayer`, a sibling ABC for layers taking
positions. It is a sibling and not a subclass of `Layer` because narrowing
`Layer.__call__` from one tensor to two is not a valid override -- verified that
subclassing still errors under mypy, so a subclass would only move the
`type: ignore` up a level. The ignore is now gone entirely.

Simplifications: `_rms_normalize` inlined into `RMSNormLayer` (single caller),
`_scale_parameter` deleted and `_affine_parameters` restored to its symmetric
form, dead comments removed.

RoPE no longer requires a static `head_dim`. The frequency ladder is built with
`pt.arange` and folds to a single constant when the size is known, so the static
requirement bought nothing; a symbolic feature axis now works and is tested.
Everything is computed in the input's dtype via `x.type.dtype`, so no float64
leaks into single-precision graphs. Returned variables are named. The
interleaved pairing uses `x[..., 0::2].set(...)`.

Retract a false claim: `SplitDims`/`JoinDims` do lower on JAX -- they are
rewritten to `Reshape` during canonicalization. The real defect was in
`tests/dispatch/jax/test_basic.py`, whose `compare_jax_and_py` built its own
`RewriteDatabaseQuery(include=["jax"])` that omits canonicalize, so any op which
only reaches the backend after being rewritten raised "No JAX conversion" there
while compiling fine in real use. It now uses `JAX._optimizer`. `MultiheadAttention`
is and always was JAX-compilable.

Docs: American English, References sections with the RoFormer, positional
interpolation and RMSNorm papers, pinned permalinks to the HuggingFace and
torchtune implementations for the two pairing conventions, parameter docs cut
down, `See Also` removed.
@cetagostini

cetagostini commented Aug 8, 2026 •

Copy link
Copy Markdown
Author

hey jesse, addressed your comments.

  • activations.py/constant_like churn reverted, file is untouched again. epsilon is pt.constant(..., dtype=X.type.dtype) now
  • _rms_normalize and _scale_parameter inlined, _affine_parameters symmetric again, dead comments gone
  • static head_dim dropped, freqs built with pt.arange and fold to a constant, all in x.type.dtype so no float64
  • you were right about split/join. the actual bug was ours: compare_jax_and_py omits canonicalize, so it reported a gap that doesn't exist. fixed
  • subclassing Layer still fails mypy, can't narrow __call__, so PositionalLayer is a sibling ABC. no type: ignore left

spelling/refs/permalinks/doc trimming done. 365 passed, mypy and pre-commit clean.

@jessegrabowski jessegrabowski left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

very nice, one more round and we're done i think

Comment thread pytensor_ml/base.py Outdated
Comment thread pytensor_ml/layers/positional.py Outdated
Comment thread pytensor_ml/layers/positional.py Outdated
Comment thread pytensor_ml/layers/positional.py
Comment thread tests/dispatch/jax/test_positional.py Outdated
Comment thread tests/test_layers.py Outdated
Comment thread tests/test_positional.py Outdated
@jessegrabowski jessegrabowski added layer New or improvement to existing LayerOp enhancement New feature or request labels Aug 8, 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.
@cetagostini cetagostini changed the title Add RMSNorm and rotary position embeddings Add rotary position embeddings Sep 12, 2026
@cetagostini

Copy link
Copy Markdown
Author

Addressed the remaining review in 3d54ea0 and merged current main. Since #159 now provides RMSNorm, its implementation and tests are retained unchanged.

  • Replaced PositionalLayer with a small general VariadicLayer; existing Layer stays unary. UnaryLayerOp already accepts multiple inputs and refers to its single output.
  • Split cosine/sine assignments, used named arguments for the rotation expressions, shortened prose, and removed the remaining float64 frequency intermediates.
  • Audited the tests and removed tautologies/duplicates; attention and gradient checks now compare independent numerical expectations.
  • Removed the dedicated JAX positional tests. The CI matrix runs the same general RoPE tests on native CPU, JAX, and MLX.
  • Left zip_concatenate out, as requested.

Verification: 271 native integration tests passed; JAX 27 passed/1 strict expected failure; MLX 26 passed/2 strict expected failures. The expected failures are documented upstream limitations for unknown head dimensions (MLX arange and JAX half-split bounds). Mypy and pre-commit passed; API examples and the decoder smoke ran successfully.

@jessegrabowski jessegrabowski left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

lets gets this over the finish line, it's been way too long (my fault). A few final requests, hit them then feel free to merge it.

scaling=self.scaling,
scaling_factor=self.scaling_factor,
)
angles = position_ids[..., None].astype(dtype) * inverse_frequencies

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Widen the frequency arithmetic to float32 when x is float16 and cast cos and sin back to x.dtype, the way _accumulator_dtype does in norm.py: float16 in, float32 for the angles, float16 out. At float16 a position past 2048 is not representable and a base of 500000 casts to inf, which leaves every pair but the first at angle zero.

Comment thread tests/test_positional.py
expected = rope_np(
x_np, positions_np, pairing=pairing, scaling=scaling, scaling_factor=scaling_factor
)
np.testing.assert_allclose(result, expected, atol=1e-6)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Assert that out.dtype matches the input dtype here. A float32 input in a float64-default graph is the configuration #164 caught in SDPA and no case in this file covers it.

Comment thread tests/test_positional.py
out.eval({x: x_np, positions: positions_np}),
rope_np(x_np, positions_np, pairing=pairing),
atol=1e-6,
)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Add a float16 case at head_dim 128 or a position past 4096. The angles are built at the input dtype, so float16 cosines come out up to 0.5 off and nothing in the suite sees it.

:toctree: generated/

Layer
pytensor_ml.base.VariadicLayer

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Import VariadicLayer in pytensor_ml/layers/__init__.py and list it here as a bare name. Every other entry in this file resolves against the currentmodule on line 4.

matrix:
os: [ubuntu-latest]
python-version: ["3.12"]
pytensor-mode: [FAST_RUN]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Delete this axis. Line 118 already defaults the mode with || 'FAST_RUN', and floatX gets no axis of its own.

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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Move this harness change to its own PR. Twelve compare_jax_and_py call sites exist to prove an op lowers to JAX, and a canonicalizing database can rewrite that op away before the linker sees it.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request layer New or improvement to existing LayerOp

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants