""" Rotary Position Embeddings (RoPE) ────────────────────────────────── RoPE replaces absolute sinusoidal PE with a per-head rotation applied directly inside each attention layer. WHY ROPE OVER ABSOLUTE PE: - Encodes *relative* distance between tokens (token i attends to token j via their angular difference), not absolute position. - Generalises better to longer sequences than were seen during training. - Is the foundation for context-extension techniques (YaRN, NTK-scaling) that scale Llama from 4k → 128k context with minimal quality loss. - Used in: Llama, Mistral, Falcon, Gemma, Qwen. HOW IT WORKS: For a head dimension `dim`, RoPE treats each consecutive pair of values (x_{2i}, x_{2i+1}) as a 2D vector and rotates it by angle (pos × θ_i), where θ_i = 1 / base^(2i/dim). The dot-product of two rotated vectors depends only on their positional difference — giving attention a clean relative-distance signal for free. """ import torch import torch.nn as nn class RotaryEmbedding(nn.Module): """ Pre-computes and caches cos/sin tables for RoPE. Args: dim: Head dimension (d_model // n_heads). RoPE operates per-head. max_seq_len: Pre-compute tables up to this length. Safe to set high (8192). base: Frequency base (θ). Default 10000. Use 500000 for Llama 3.1-style extended context. """ def __init__(self, dim: int, max_seq_len: int = 8192, base: int = 10000): super().__init__() # Inverse frequencies: θ_i = 1 / base^(2i / dim), shape [dim/2] inv_freq = 1.0 / (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, seq_len: int) -> None: """Pre-compute cos and sin tables for positions [0, seq_len).""" t = torch.arange(seq_len, device=self.inv_freq.device).float() freqs = torch.outer(t, self.inv_freq) # [seq, dim/2] emb = torch.cat([freqs, freqs], dim=-1) # [seq, dim] (duplicate for rotate_half) # Shape [1, 1, seq, dim] so it broadcasts with [B, heads, T, head_dim] self.register_buffer("cos_cached", emb.cos()[None, None], persistent=False) self.register_buffer("sin_cached", emb.sin()[None, None], persistent=False) def forward( self, seq_len: int, start_pos: int = 0, ) -> tuple[torch.Tensor, torch.Tensor]: """ Return (cos, sin) sliced to [start_pos : start_pos + seq_len]. Args: seq_len: Number of tokens to generate embeddings for. start_pos: Starting position index in the sequence. For training or prompt prefill, start_pos = 0. For single-token incremental decoding with KV-cache, start_pos equals the number of tokens already stored in the cache. Returns: cos: shape [1, 1, seq_len, dim] sin: shape [1, 1, seq_len, dim] """ return ( self.cos_cached[:, :, start_pos : start_pos + seq_len, :], self.sin_cached[:, :, start_pos : start_pos + seq_len, :], ) def rotate_half(x: torch.Tensor) -> torch.Tensor: """ Rotate the second half of the last dimension to the front with a sign flip. For x = [x1, x2] (split at midpoint): rotate_half(x) = [-x2, x1] This implements the 2D rotation without computing sin/cos per element. """ half = x.shape[-1] // 2 x1, x2 = x[..., :half], x[..., half:] return torch.cat([-x2, x1], dim=-1) def apply_rotary_emb( q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: """ Apply RoPE rotation to query and key tensors. Args: q: Query tensor, shape [B, n_heads, T, head_dim] k: Key tensor, shape [B, n_kv_heads, T, head_dim] cos: Cosine table, shape [1, 1, T, head_dim] sin: Sine table, shape [1, 1, T, head_dim] Returns: Rotated (q, k) tensors with the same shapes. """ q_rot = (q * cos) + (rotate_half(q) * sin) k_rot = (k * cos) + (rotate_half(k) * sin) return q_rot, k_rot