File size: 4,343 Bytes
aa7bfed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
115
116
117
118
119
120
121
122
123
124
125
"""Vectorized RoPE: pairwise 2D rotations of query/key slices.

Two pairing conventions:

- ``interleaved`` — paper / GPT-J style: rotate dims ``(2i, 2i+1)``.
- ``llama`` — Hugging Face Llama/Qwen/SmolLM style: rotate
  ``(i, i + dim/2)`` (the ``rotate_half`` layout).

Frequencies match the usual schedule ``ω_i = base^{-2i/d}``.
"""

from __future__ import annotations

import numpy as np

VALID_STYLES = ("interleaved", "llama")


def pair_frequencies(dim: int, base: float = 10000.0) -> np.ndarray:
    """Return ``ω`` of shape ``(dim/2,)`` with ``ω_i = base ** (-2i / dim)``."""
    if dim % 2 != 0:
        raise ValueError(f"RoPE requires an even dim, got {dim}")
    n_pairs = dim // 2
    pair_idx = np.arange(n_pairs, dtype=np.float64) * 2.0
    return 1.0 / (base ** (pair_idx / dim))


def theta_grid(seq_len: int, dim: int, base: float = 10000.0) -> np.ndarray:
    """Rotation angles ``θ[k, i] = k * ω_i``, shape ``(seq_len, dim/2)``."""
    omega = pair_frequencies(dim, base=base)
    positions = np.arange(seq_len, dtype=np.float64)[:, None]
    return positions * omega[None, :]


def _broadcast_angles(angles: np.ndarray, like: np.ndarray) -> np.ndarray:
    """Broadcast ``(seq, n_pairs)`` onto ``like`` with shape ``(..., seq, n_pairs)``."""
    lead = like.ndim - 2
    return angles.reshape((1,) * lead + angles.shape)


def rotate_pair(x_even: np.ndarray, x_odd: np.ndarray, theta: np.ndarray):
    """Apply a 2D rotation by ``theta`` (radians) to coordinate pairs."""
    cos = np.cos(theta)
    sin = np.sin(theta)
    return x_even * cos - x_odd * sin, x_even * sin + x_odd * cos


def apply_rope(
    x: np.ndarray,
    base: float = 10000.0,
    style: str = "interleaved",
) -> np.ndarray:
    """Rotate last-dimension pairs of ``x`` (``(..., seq_len, dim)``)."""
    if style not in VALID_STYLES:
        raise ValueError(f"style must be one of {VALID_STYLES}, got {style!r}")
    x = np.asarray(x, dtype=np.float64)
    if x.ndim < 2:
        raise ValueError("x must be at least 2D (seq, dim)")
    dim = x.shape[-1]
    seq_len = x.shape[-2]
    angles = theta_grid(seq_len, dim, base=base)
    cos = np.cos(angles)
    sin = np.sin(angles)
    n_pairs = dim // 2
    out = np.empty_like(x)

    if style == "interleaved":
        even = x[..., 0::2]
        odd = x[..., 1::2]
        cos_b = _broadcast_angles(cos, even)
        sin_b = _broadcast_angles(sin, odd)
        out[..., 0::2] = even * cos_b - odd * sin_b
        out[..., 1::2] = even * sin_b + odd * cos_b
    else:
        first = x[..., :n_pairs]
        second = x[..., n_pairs:]
        cos_b = _broadcast_angles(cos, first)
        sin_b = _broadcast_angles(sin, second)
        out[..., :n_pairs] = first * cos_b - second * sin_b
        out[..., n_pairs:] = first * sin_b + second * cos_b
    return out


def pair_xy(x: np.ndarray, token: int, pair: int, style: str = "interleaved") -> tuple[float, float]:
    """Return the 2D coordinates of one RoPE pair at one token."""
    vec = np.asarray(x)
    if vec.ndim != 2:
        raise ValueError("pair_xy expects (seq, dim)")
    dim = vec.shape[-1]
    n_pairs = dim // 2
    if not (0 <= pair < n_pairs):
        raise IndexError(f"pair {pair} out of range 0..{n_pairs - 1}")
    if style == "interleaved":
        return float(vec[token, 2 * pair]), float(vec[token, 2 * pair + 1])
    if style == "llama":
        return float(vec[token, pair]), float(vec[token, pair + n_pairs])
    raise ValueError(f"unknown style {style!r}")


def pair_dim_labels(pair: int, dim: int, style: str = "interleaved") -> tuple[str, str]:
    n_pairs = dim // 2
    if style == "interleaved":
        return f"dim {2 * pair}", f"dim {2 * pair + 1}"
    return f"dim {pair}", f"dim {pair + n_pairs}"


def l2_norms(x: np.ndarray) -> np.ndarray:
    return np.linalg.norm(np.asarray(x, dtype=np.float64), axis=-1)


def row_cosine(a: np.ndarray, b: np.ndarray) -> np.ndarray:
    a = np.asarray(a, dtype=np.float64)
    b = np.asarray(b, dtype=np.float64)
    an = l2_norms(a)
    bn = l2_norms(b)
    dots = np.sum(a * b, axis=-1)
    return dots / (an * bn + 1e-12)


def attention_scores(q: np.ndarray, k: np.ndarray) -> np.ndarray:
    """Unnormalized ``Q K^T`` for ``(seq, dim)`` slices."""
    q = np.asarray(q, dtype=np.float64)
    k = np.asarray(k, dtype=np.float64)
    return q @ k.T