RoPE Positional Encoding
Rotary Position Embedding (RoPE).
Implements rotary position encoding by rotating query and key vectors by position-dependent angles. Each consecutive pair of dimensions is treated as a 2D plane and rotated by an angle proportional to the token position and inversely proportional to a geometric progression of frequencies. This encodes relative position information directly into the attention dot product without requiring learned or absolute position embeddings.
- Reference:
Su et al. (2024), “RoFormer: Enhanced Transformer with Rotary Position Embedding”, arXiv:2104.09864.
- class src.model.embeddings.rope.RoPE(*args: Any, **kwargs: Any)[source]
Bases:
ModuleRotary Position Embedding over consecutive dimension pairs.
Applies a 2D rotation to each pair of adjacent dimensions in the input tensor. The rotation angle for pair
iat positionpis:theta_i(p) = p * scaling * base^{-i / (pair_dim - 1)}
This enables the attention dot product
q^T kto depend only on the relative position between tokens, as the rotation satisfies:R(p_q)^T R(p_k) = R(p_k - p_q)
- Reference:
Su et al. (2024), “RoFormer: Enhanced Transformer with Rotary Position Embedding”, arXiv:2104.09864.
- Parameters:
head_dim – Dimensionality of each attention head. Must be even for proper pairing.
base – Base frequency for the geometric progression of rotation frequencies. Defaults to
10000.0.scaling – Position scaling factor applied to token indices before computing angles. Defaults to
1.0.
- head_dim
Total head dimensionality.
- pair_dim
Number of dimension pairs (
head_dim // 2).
- base
Base frequency for inverse frequency computation.
- scaling
Position scaling factor.
- forward(x: torch.Tensor, logical_layer_idx: int = 0) torch.Tensor[source]
Apply rotary position encoding.
- Parameters:
x – Input tensor of shape
(batch, heads, seq_len, head_dim).logical_layer_idx – Logical layer index (unused; accepted for interface compatibility with other positional encodings).
- Returns:
Tensor of same shape as
xwith rotary position encoding applied. Ifpair_dim == 0, returnsxunchanged.