Download model/positional_encoding.py from Vivid86/MiniTransformer-91M: direct link, hf CLI and curl.
- Browser
- Download file 4.4 kB
-
https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/positional_encoding.py
- Command line
-
hf download hf://Vivid86/MiniTransformer-91M/model/positional_encoding.py
-
curl -L -o positional_encoding.py https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/positional_encoding.py
4.4 kB
| """ | |
| 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 | |