Attention requires two things from positional encodings:
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 ( ei ) from positional location ( pi ). Worse, if a model was trained on sequence lengths up to N=2048 , position 2049 presented an unseen vector p2049 , causing immediate generation collapse.
Subsequent approaches tried adding relative bias terms bi,j directly into the attention matrix QKT+B . While functionally effective, modifying the N×N 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.
To make attention relative, we want an encoding function R(x,m) applied to a vector x at position m such that the dot product between Query at position m and Key at position n depends only on the relative offset ( m−n ):
How do you preserve vector norms while encoding an angle shift? Complex space rotations.
In a 2D plane, rotating a vector x=[x1,x2]T by an angle mθ is represented by the orthogonal rotation matrix:
When you compute the dot product between a rotated Query at position m and a rotated Key at position n :
Because RT(m)R(n)=R(n−m) , the absolute positions m and n cancel out completely. The resulting attention weight is strictly a function of the distance ( m−n ).
For a d -dimensional vector, RoPE splits the channels into d/2 pairs of 2D planes, applying a different rotation frequency θi to each pair:
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 x~=[−x2,x1,−x4,x3,…] .
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
inv_freq = 1.0 / (self.base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
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)
freqs = torch.outer(t, self.inv_freq)
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]
return (x * cos) + (self._rotate_half(x) * sin)