File size: 4,401 Bytes
a2ec932
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
"""
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