Why 2D Rotations Solved Transformer Long-Context: A Mechanics-First Look at RoPE
Attention requires two things from positional encodings:
- Relative Distance Sensitivity: Token attending to Token should care primarily about how far apart they are ( ), not where they sit globally in the context window.
- Feature Preservation: Injecting position information must not destroy or mangle the semantic embedding features learned by the model.
Original absolute positional encodings added static sine/cosine waves directly to input embeddings:
This forced the model to burn parameter capacity un-mixing semantic meaning ( ) from positional location ( ). Worse, if a model was trained on sequence lengths up to , position presented an unseen vector , causing immediate generation collapse.
Subsequent approaches tried adding relative bias terms directly into the attention matrix . While functionally effective, modifying the matrix broke kernel-level GPU fusions like FlashAttention and introduced huge memory overheads.
Enter Rotary Position Embedding (RoPE) (Su et al., 2021). Instead of adding positional vectors or patching the attention matrix, RoPE rotates the Query and Key vectors in 2D sub-planes before computing attention.
The Geometry: Inner Products Under Rotation
To make attention relative, we want an encoding function applied to a vector at position such that the dot product between Query at position and Key at position depends only on the relative offset ( ):
How do you preserve vector norms while encoding an angle shift? Complex space rotations.
In a 2D plane, rotating a vector by an angle is represented by the orthogonal rotation matrix:
When you compute the dot product between a rotated Query at position and a rotated Key at position :
Because , the absolute positions and cancel out completely. The resulting attention weight is strictly a function of the distance ( ).
For a -dimensional vector, RoPE splits the channels into pairs of 2D planes, applying a different rotation frequency to each pair:
High-Performance Vectorized PyTorch Implementation
In practice, explicitly constructing full block-diagonal rotation matrices for every head and token is slow. We can implement RoPE efficiently using the complex number representation or the real-valued slice trick:
where .
Here is a clean, production-ready PyTorch implementation:
python
import torch
import torch.nn as nn
class RotaryPositionalEmbedding(nn.Module):
"""
Rotary Position Embedding (RoPE) for multi-head attention.
Applies 2D plane rotations across d_model / 2 feature pairs.
"""
def __init__(self, dim: int, max_seq_len: int = 8192, base: float = 10000.0):
super().__init__()
self.dim = dim
self.base = base
# Compute frequency bands theta_i = 10000^(-2(i-1)/d)
inv_freq = 1.0 / (self.base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
# Precompute cosine and sine cache for max_seq_len
self._build_cache(max_seq_len)
def _build_cache(self, max_seq_len: int):
t = torch.arange(max_seq_len, dtype=self.inv_freq.dtype)
# Outer product: [seq_len, dim / 2]
freqs = torch.outer(t, self.inv_freq)
# Duplicate frequencies to match feature dimensions: [seq_len, dim]
emb = torch.cat((freqs, freqs), dim=-1)
self.register_buffer("cos_cached", emb.cos(), persistent=False)
self.register_buffer("sin_cached", emb.sin(), persistent=False)
def _rotate_half(self, x: torch.Tensor) -> torch.Tensor:
"""Splits vector in half and negates/swaps halves: [-x2, x1]"""
x1 = x[..., : self.dim // 2]
x2 = x[..., self.dim // 2 :]
return torch.cat((-x2, x1), dim=-1)
def forward(self, x: torch.Tensor, seq_len: int) -> torch.Tensor:
"""
Args:
x: Tensor of shape [batch_size, num_heads, seq_len, head_dim]
seq_len: Current sequence length
Returns:
Tensor with RoPE applied to Query or Key
"""
cos = self.cos_cached[:seq_len, :].unsqueeze(0).unsqueeze(1) # [1, 1, seq_len, dim]
sin = self.sin_cached[:seq_len, :].unsqueeze(0).unsqueeze(1) # [1, 1, seq_len, dim]
# Apply elementwise rotation formula
return (x * cos) + (self._rotate_half(x) * sin)
Top comments (0)