Add rotary position embeddings - #55
cetagostini wants to merge 4 commits into
Conversation
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.
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 ☂️ |
| return (X - mu) / pt.sqrt(sigma_sq + epsilon), mu, sigma_sq | ||
|
|
||
|
|
||
| def _rms_normalize(X, epsilon): |
There was a problem hiding this comment.
This is only called once, just inline it
| return inferred | ||
|
|
||
|
|
||
| def _scale_parameter(name: str, n_in: int) -> TrainableParameter: |
There was a problem hiding this comment.
I don't like this, it makes _affine_parameters asymmetric. At this point just inline the variable creation everywhere.
| if not self.affine: | ||
| return [X_normalized] | ||
|
|
||
| # Scale-only, so the affine transform contributes one input rather than the ``(loc, scale)`` |
|
|
||
| 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 |
There was a problem hiding this comment.
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.
| 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, |
There was a problem hiding this comment.
simplify, don't need the huge essay
There was a problem hiding this comment.
also i'd rather we just pick something for our library
| 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 |
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
This all needs to be dialed in, each parameter has the book of genesis under it
| scaling_factor : float, optional | ||
| Extension factor for the scaled variants. Default is 1.0. | ||
|
|
||
| See Also |
There was a problem hiding this comment.
I don't personally care for the see also section, ymmv
| self.scaling = scaling | ||
| self.scaling_factor = scaling_factor | ||
|
|
||
| # Positions are a required second input, which the one-tensor Layer.__call__ signature does not |
| # 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 |
There was a problem hiding this comment.
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.
|
hey jesse, addressed your comments.
spelling/refs/permalinks/doc trimming done. 365 passed, mypy and pre-commit clean. |
jessegrabowski
left a comment
There was a problem hiding this comment.
very nice, one more round and we're done i think
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.
|
Addressed the remaining review in 3d54ea0 and merged current main. Since #159 now provides RMSNorm, its implementation and tests are retained unchanged.
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
left a comment
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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.
| expected = rope_np( | ||
| x_np, positions_np, pairing=pairing, scaling=scaling, scaling_factor=scaling_factor | ||
| ) | ||
| np.testing.assert_allclose(result, expected, atol=1e-6) |
There was a problem hiding this comment.
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.
| out.eval({x: x_np, positions: positions_np}), | ||
| rope_np(x_np, positions_np, pairing=pairing), | ||
| atol=1e-6, | ||
| ) |
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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] |
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
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.
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 configuredRotaryEmbeddinglayer: explicit positions, half/adjacent pairing, linear interpolation, and static NTK scaling.VariadicLayerbase with__call__(self, *inputs: pt.TensorLike) -> pt.TensorVariable, replacingPositionalLayer. Existing unaryLayerremains unchanged.UnaryLayerOpalready accepts multiple inputs and names its single-output contract.zip_concatenateprimitive.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
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/