MiniTransformer-91M / model /positional_encoding.py
Vivid86's picture
Upload folder using huggingface_hub
a2ec932 verified
Raw History Blame Contribute Delete
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