Spaces:
Running on Zero
Running on Zero
Commit ·
aa7bfed
1
Parent(s): 1d852fb
feat: add rope visualization only with some minor cmparison with sinusoidial positional embedding
Browse files- README.md +50 -13
- app.py +353 -20
- requirements.txt +7 -5
- src/absolute_pe.py +35 -0
- src/extract.py +260 -0
- src/plots.py +305 -0
- src/rope.py +124 -0
- tests/test_core.py +115 -0
README.md
CHANGED
|
@@ -1,13 +1,50 @@
|
|
| 1 |
-
---
|
| 2 |
-
title:
|
| 3 |
-
emoji: 📚
|
| 4 |
-
colorFrom: gray
|
| 5 |
-
colorTo: green
|
| 6 |
-
sdk: gradio
|
| 7 |
-
sdk_version: 6.26.0
|
| 8 |
-
python_version: 3.11
|
| 9 |
-
app_file: app.py
|
| 10 |
-
pinned: false
|
| 11 |
-
short_description:
|
| 12 |
-
---
|
| 13 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
title: RoPE Embedding Visualization
|
| 3 |
+
emoji: 📚
|
| 4 |
+
colorFrom: gray
|
| 5 |
+
colorTo: green
|
| 6 |
+
sdk: gradio
|
| 7 |
+
sdk_version: 6.26.0
|
| 8 |
+
python_version: 3.11
|
| 9 |
+
app_file: app.py
|
| 10 |
+
pinned: false
|
| 11 |
+
short_description: Visualize how RoPE rotates query and key vectors
|
| 12 |
+
---
|
| 13 |
+
|
| 14 |
+
# RoPE Explorer
|
| 15 |
+
|
| 16 |
+
Interactive Gradio app for **Rotary Position Embedding**. The Hugging Face Space
|
| 17 |
+
runs `app.py` (`sdk: gradio`). Local Docker (`Dockerfile`, `compose.yaml`) is for
|
| 18 |
+
running `python app.py` on port **7860**.
|
| 19 |
+
|
| 20 |
+
Production math lives in [`src/`](src/) (`rope.py`, `absolute_pe.py`, `extract.py`,
|
| 21 |
+
`plots.py`). Scratch scripts under `rope_implementation/`,
|
| 22 |
+
`absolute_sinusoidal_position_embedding/`, and `relative_pos_embedding/` are
|
| 23 |
+
learning notes only and are **not** imported by the app.
|
| 24 |
+
|
| 25 |
+
## Modes
|
| 26 |
+
|
| 27 |
+
1. **Random matrix** — sample even-width Q (and K) tensors, apply numpy RoPE,
|
| 28 |
+
inspect heatmaps, pairwise 2D rotation, `QK^T`, and additive sinusoidal PE.
|
| 29 |
+
2. **Real model** — lazy-load an ungated Llama-like checkpoint (default
|
| 30 |
+
`HuggingFaceTB/SmolLM2-135M`), take `embed_tokens`, first-layer `q_proj` /
|
| 31 |
+
`k_proj` (GQA-aware), and compare educational numpy RoPE (`llama` pairing)
|
| 32 |
+
to the model's `rotary_emb`. No Hugging Face token is required. Gated models
|
| 33 |
+
are not used.
|
| 34 |
+
|
| 35 |
+
First load of a model downloads weights into the cache; later runs reuse the
|
| 36 |
+
last loaded model in memory.
|
| 37 |
+
|
| 38 |
+
CPU is enough for SmolLM2 and Qwen2.5-0.5B. TinyLlama is included for a larger
|
| 39 |
+
example and may be slow on CPU. This Space does **not** require ZeroGPU
|
| 40 |
+
(`@spaces.GPU` is unused).
|
| 41 |
+
|
| 42 |
+
## Local run
|
| 43 |
+
|
| 44 |
+
```bash
|
| 45 |
+
pip install -r requirements.txt
|
| 46 |
+
python app.py
|
| 47 |
+
```
|
| 48 |
+
|
| 49 |
+
Or `docker compose up`. Hugging Face Cloud uses the README YAML (`sdk: gradio`),
|
| 50 |
+
not the Docker image, unless the Space SDK is switched to Docker.
|
app.py
CHANGED
|
@@ -1,20 +1,353 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""RoPE Explorer Gradio app. Imports only from ``src/``."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import plotly.graph_objects as go
|
| 7 |
+
import gradio as gr
|
| 8 |
+
|
| 9 |
+
from src.absolute_pe import add_positional_encoding
|
| 10 |
+
from src.extract import (
|
| 11 |
+
DEFAULT_MODEL,
|
| 12 |
+
MAX_SEQ_LEN,
|
| 13 |
+
MODEL_CHOICES,
|
| 14 |
+
expand_kv_heads,
|
| 15 |
+
extract_from_model,
|
| 16 |
+
random_qk,
|
| 17 |
+
select_head,
|
| 18 |
+
)
|
| 19 |
+
from src.plots import (
|
| 20 |
+
additive_pe_heatmaps,
|
| 21 |
+
attention_bars,
|
| 22 |
+
attention_heatmaps,
|
| 23 |
+
bulk_before_after_delta,
|
| 24 |
+
frequency_strip,
|
| 25 |
+
norm_compare_add_vs_rope,
|
| 26 |
+
norms_and_cosine,
|
| 27 |
+
position_sweep,
|
| 28 |
+
rotation_2d,
|
| 29 |
+
theta_heatmap,
|
| 30 |
+
)
|
| 31 |
+
from src.rope import (
|
| 32 |
+
attention_scores,
|
| 33 |
+
pair_dim_labels,
|
| 34 |
+
pair_frequencies,
|
| 35 |
+
pair_xy,
|
| 36 |
+
rotate_pair,
|
| 37 |
+
theta_grid,
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
HOWTO_MD = """
|
| 41 |
+
# How RoPE works
|
| 42 |
+
|
| 43 |
+
After each token has a vector from the **embedding table**, attention builds two extra
|
| 44 |
+
vectors per token with linear layers (`q_proj`, `k_proj`):
|
| 45 |
+
|
| 46 |
+
- **Q (query)** — what this token is looking for
|
| 47 |
+
- **K (key)** — what this token offers as a match
|
| 48 |
+
|
| 49 |
+
Attention scores are (scaled) **dot products** `Q · K`. **RoPE rotates those Q and K
|
| 50 |
+
vectors in 2D planes before the dot product.** It does **not** add a position vector
|
| 51 |
+
onto the raw token embeddings. This app shows embeddings only as context, then focuses
|
| 52 |
+
on Q and K before vs after RoPE.
|
| 53 |
+
|
| 54 |
+
## Pairwise rotation
|
| 55 |
+
|
| 56 |
+
For even dimension `d`, pair `i` uses frequency
|
| 57 |
+
|
| 58 |
+
$$\\omega_i = 10000^{-2i/d},\\qquad \\theta(k,i) = k\\,\\omega_i$$
|
| 59 |
+
|
| 60 |
+
**Interleaved (paper-style)** pairing `(2i, 2i+1)`:
|
| 61 |
+
|
| 62 |
+
$$
|
| 63 |
+
x'_{2i} = x_{2i}\\cos\\theta - x_{2i+1}\\sin\\theta,\\qquad
|
| 64 |
+
x'_{2i+1} = x_{2i}\\sin\\theta + x_{2i+1}\\cos\\theta
|
| 65 |
+
$$
|
| 66 |
+
|
| 67 |
+
Hugging Face Llama-like models use the same frequencies but pair `(i, i + d/2)`
|
| 68 |
+
(`rotate_half`). Random-matrix mode uses interleaved pairing; real models use the
|
| 69 |
+
Llama layout so the numpy implementation can be checksummed against `rotary_emb`.
|
| 70 |
+
|
| 71 |
+
Relative positions fall out of the algebra: `R(m)^T R(n) = R(n-m)`.
|
| 72 |
+
|
| 73 |
+
Shaw relative attention (learned bias `b_{m-n}` on scores) is a **different**
|
| 74 |
+
mechanism and is not computed here.
|
| 75 |
+
"""
|
| 76 |
+
|
| 77 |
+
PLACEHOLDER = go.Figure().update_layout(
|
| 78 |
+
title="Run **Compute** on the Setup tab first",
|
| 79 |
+
template="plotly_white",
|
| 80 |
+
height=320,
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def _safe_slider_max(n: int) -> int:
|
| 85 |
+
"""Gradio sliders need max > min; keep a one-step range even at edge cases."""
|
| 86 |
+
return max(int(n), 1)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def compute(
|
| 90 |
+
source: str,
|
| 91 |
+
sentence: str,
|
| 92 |
+
model_name: str,
|
| 93 |
+
seq_len: int,
|
| 94 |
+
dim: int,
|
| 95 |
+
seed: int,
|
| 96 |
+
base: float,
|
| 97 |
+
progress=gr.Progress(track_tqdm=False),
|
| 98 |
+
):
|
| 99 |
+
try:
|
| 100 |
+
if source.startswith("Random"):
|
| 101 |
+
progress(0.4, desc="Sampling random Q/K")
|
| 102 |
+
data = random_qk(int(seq_len), int(dim), seed=int(seed), base=float(base))
|
| 103 |
+
else:
|
| 104 |
+
progress(0.2, desc=f"Loading {model_name} (first time downloads weights)")
|
| 105 |
+
data = extract_from_model(model_name, sentence)
|
| 106 |
+
seq = int(select_head(data["q_before"], 0).shape[0])
|
| 107 |
+
head_dim = int(select_head(data["q_before"], 0).shape[1])
|
| 108 |
+
n_pairs = head_dim // 2
|
| 109 |
+
n_heads = max(int(data["n_q_heads"]) - 1, 0)
|
| 110 |
+
checksum = data["checksum"]
|
| 111 |
+
if checksum is None:
|
| 112 |
+
status = (
|
| 113 |
+
f"Random Q/K · seq={seq} · dim={head_dim} · base={data['base']:g} · "
|
| 114 |
+
f"style={data['style']}"
|
| 115 |
+
)
|
| 116 |
+
else:
|
| 117 |
+
status = (
|
| 118 |
+
f"Model `{data['model_name']}` · {seq} tokens · head_dim={head_dim} · "
|
| 119 |
+
f"Q heads={data['n_q_heads']} · KV heads={data['n_kv_heads']} · "
|
| 120 |
+
f"rope_theta={data['base']:g} · "
|
| 121 |
+
f"max |numpy RoPE − model rotary| on Q = **{checksum:.3e}**"
|
| 122 |
+
)
|
| 123 |
+
token_labels = ", ".join(data["tokens"][:seq])
|
| 124 |
+
status = status + f"\n\nTokens: `{token_labels}`"
|
| 125 |
+
return (
|
| 126 |
+
data,
|
| 127 |
+
status,
|
| 128 |
+
gr.update(maximum=_safe_slider_max(n_heads), value=0),
|
| 129 |
+
gr.update(maximum=_safe_slider_max(seq - 1), value=0),
|
| 130 |
+
gr.update(maximum=_safe_slider_max(n_pairs - 1), value=0),
|
| 131 |
+
gr.update(maximum=_safe_slider_max(seq - 1), value=0),
|
| 132 |
+
)
|
| 133 |
+
except Exception as exc:
|
| 134 |
+
return (
|
| 135 |
+
None,
|
| 136 |
+
f"**Error:** {exc}",
|
| 137 |
+
gr.update(),
|
| 138 |
+
gr.update(),
|
| 139 |
+
gr.update(),
|
| 140 |
+
gr.update(),
|
| 141 |
+
)
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def _qk_slice(data: dict, which: str, head: int):
|
| 145 |
+
before = data["q_before"] if which == "Q" else data["k_before"]
|
| 146 |
+
after = data["q_after"] if which == "Q" else data["k_after"]
|
| 147 |
+
return select_head(before, head), select_head(after, head)
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def update_bulk(data, which, head, mod_2pi):
|
| 151 |
+
if not data:
|
| 152 |
+
fig = PLACEHOLDER
|
| 153 |
+
return fig, fig, fig, fig
|
| 154 |
+
before, after = _qk_slice(data, which, int(head))
|
| 155 |
+
dim = before.shape[-1]
|
| 156 |
+
seq = before.shape[0]
|
| 157 |
+
return (
|
| 158 |
+
bulk_before_after_delta(before, after, tokens=data["tokens"]),
|
| 159 |
+
norms_and_cosine(before, after, tokens=data["tokens"]),
|
| 160 |
+
theta_heatmap(seq, dim, data["base"], mod_2pi=bool(mod_2pi)),
|
| 161 |
+
frequency_strip(dim, data["base"]),
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def update_individual(data, which, head, token, pair, sweep):
|
| 166 |
+
if not data:
|
| 167 |
+
return "Compute on the Setup tab first.", PLACEHOLDER, PLACEHOLDER
|
| 168 |
+
before, after = _qk_slice(data, which, int(head))
|
| 169 |
+
token = int(np.clip(token, 0, before.shape[0] - 1))
|
| 170 |
+
n_pairs = before.shape[1] // 2
|
| 171 |
+
pair = int(np.clip(pair, 0, n_pairs - 1))
|
| 172 |
+
style = data["style"]
|
| 173 |
+
xb, yb = pair_xy(before, token, pair, style=style)
|
| 174 |
+
xa, ya = pair_xy(after, token, pair, style=style)
|
| 175 |
+
theta = float(theta_grid(before.shape[0], before.shape[1], data["base"])[token, pair])
|
| 176 |
+
cos_t, sin_t = float(np.cos(theta)), float(np.sin(theta))
|
| 177 |
+
xe_chk, xo_chk = rotate_pair(np.array([xb]), np.array([yb]), np.array([theta]))
|
| 178 |
+
d0, d1 = pair_dim_labels(pair, before.shape[1], style=style)
|
| 179 |
+
table = f"""
|
| 180 |
+
### Token `{token}` · pair `{pair}` (`{d0}`, `{d1}`)
|
| 181 |
+
|
| 182 |
+
| | {d0} | {d1} |
|
| 183 |
+
|---|---:|---:|
|
| 184 |
+
| before | {xb:.6f} | {yb:.6f} |
|
| 185 |
+
| after | {xa:.6f} | {ya:.6f} |
|
| 186 |
+
| check (`rotate_pair`) | {float(xe_chk):.6f} | {float(xo_chk):.6f} |
|
| 187 |
+
|
| 188 |
+
**θ(k,i) = {theta:.6f} rad** · cos = {cos_t:.6f} · sin = {sin_t:.6f}
|
| 189 |
+
|
| 190 |
+
`x'_even = x_even cos θ − x_odd sin θ`
|
| 191 |
+
`x'_odd = x_even sin θ + x_odd cos θ`
|
| 192 |
+
"""
|
| 193 |
+
neighbors = [n for n in (token - 1, token + 1, token + 2) if 0 <= n < before.shape[0]]
|
| 194 |
+
rot = rotation_2d(before, after, token, pair, style, theta, neighbor_tokens=neighbors)
|
| 195 |
+
if sweep:
|
| 196 |
+
omega = float(pair_frequencies(before.shape[1], data["base"])[pair])
|
| 197 |
+
sweep_fig = position_sweep(xb, yb, omega, before.shape[0], token)
|
| 198 |
+
else:
|
| 199 |
+
sweep_fig = PLACEHOLDER
|
| 200 |
+
sweep_fig.update_layout(title="Enable “replay same pair at every k” to see position-only spin")
|
| 201 |
+
return table, rot, sweep_fig
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def update_attention(data, head, query_token):
|
| 205 |
+
if not data:
|
| 206 |
+
return PLACEHOLDER, PLACEHOLDER, ""
|
| 207 |
+
q_b = select_head(data["q_before"], int(head))
|
| 208 |
+
q_a = select_head(data["q_after"], int(head))
|
| 209 |
+
k_b_all = expand_kv_heads(data["k_before"], data["n_q_heads"])
|
| 210 |
+
k_a_all = expand_kv_heads(data["k_after"], data["n_q_heads"])
|
| 211 |
+
k_b = select_head(k_b_all, int(head))
|
| 212 |
+
k_a = select_head(k_a_all, int(head))
|
| 213 |
+
sb = attention_scores(q_b, k_b)
|
| 214 |
+
sa = attention_scores(q_a, k_a)
|
| 215 |
+
qt = int(np.clip(query_token, 0, sb.shape[0] - 1))
|
| 216 |
+
note = (
|
| 217 |
+
"Additive PE changes values by **addition**. RoPE encodes **relative** offset "
|
| 218 |
+
"because `R(m)^T R(n) = R(n−m)`: the score depends on the position difference, "
|
| 219 |
+
"not on absolute indices alone."
|
| 220 |
+
)
|
| 221 |
+
return attention_heatmaps(sb, sa), attention_bars(sb[qt], sa[qt], qt), note
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
def update_compare(data):
|
| 225 |
+
if not data:
|
| 226 |
+
return PLACEHOLDER, PLACEHOLDER, ""
|
| 227 |
+
emb = np.asarray(data["embeddings"], dtype=np.float64)
|
| 228 |
+
# Compare tab always uses additive PE on the embedding matrix (may be wider than a head).
|
| 229 |
+
pe, combined = add_positional_encoding(emb, base=data["base"])
|
| 230 |
+
q_b = select_head(data["q_before"], 0)
|
| 231 |
+
q_a = select_head(data["q_after"], 0)
|
| 232 |
+
heat = additive_pe_heatmaps(emb, pe, combined)
|
| 233 |
+
norms = norm_compare_add_vs_rope(emb, combined, q_b, q_a)
|
| 234 |
+
copy = """
|
| 235 |
+
**Absolute sinusoidal PE** *adds* a position-shaped vector, so both **norm and direction** change.
|
| 236 |
+
|
| 237 |
+
**RoPE** *rotates* query/key pairs: **norm stays**, and the relative angle depends on `m − n`.
|
| 238 |
+
|
| 239 |
+
Shaw-style relative attention (`q_m^T k_n + b_{m-n}`) is a third, learned-bias mechanism — not shown as a plot.
|
| 240 |
+
"""
|
| 241 |
+
return heat, norms, copy
|
| 242 |
+
|
| 243 |
+
|
| 244 |
+
def toggle_source(source: str):
|
| 245 |
+
is_random = source.startswith("Random")
|
| 246 |
+
return (
|
| 247 |
+
gr.update(visible=is_random),
|
| 248 |
+
gr.update(visible=is_random),
|
| 249 |
+
gr.update(visible=is_random),
|
| 250 |
+
gr.update(visible=not is_random),
|
| 251 |
+
gr.update(visible=not is_random),
|
| 252 |
+
)
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
with gr.Blocks(title="RoPE Explorer") as demo:
|
| 256 |
+
state = gr.State(None)
|
| 257 |
+
gr.Markdown("# RoPE Explorer")
|
| 258 |
+
gr.Markdown(
|
| 259 |
+
"Interactive view of **Rotary Position Embedding**: random Q/K matrices or "
|
| 260 |
+
"query/key vectors from a small ungated Hugging Face model."
|
| 261 |
+
)
|
| 262 |
+
with gr.Tabs():
|
| 263 |
+
with gr.Tab("How RoPE works"):
|
| 264 |
+
gr.Markdown(HOWTO_MD)
|
| 265 |
+
|
| 266 |
+
with gr.Tab("Setup"):
|
| 267 |
+
source = gr.Radio(
|
| 268 |
+
["Random matrix", "Real model"],
|
| 269 |
+
value="Random matrix",
|
| 270 |
+
label="Source",
|
| 271 |
+
)
|
| 272 |
+
with gr.Row():
|
| 273 |
+
sentence = gr.Textbox(
|
| 274 |
+
value="RoPE rotates query and key vectors.",
|
| 275 |
+
label="Sentence (real model)",
|
| 276 |
+
visible=False,
|
| 277 |
+
)
|
| 278 |
+
model_name = gr.Dropdown(
|
| 279 |
+
MODEL_CHOICES,
|
| 280 |
+
value=DEFAULT_MODEL,
|
| 281 |
+
label="Model (ungated, Llama-like)",
|
| 282 |
+
visible=False,
|
| 283 |
+
)
|
| 284 |
+
with gr.Row():
|
| 285 |
+
seq_len = gr.Slider(2, MAX_SEQ_LEN, value=16, step=1, label="Sequence length")
|
| 286 |
+
dim = gr.Slider(4, 128, value=32, step=2, label="Dimension (even)")
|
| 287 |
+
seed = gr.Number(value=0, label="Seed", precision=0)
|
| 288 |
+
base = gr.Number(
|
| 289 |
+
value=10000,
|
| 290 |
+
label="RoPE base (overridden by config.rope_theta for real models)",
|
| 291 |
+
)
|
| 292 |
+
compute_btn = gr.Button("Compute", variant="primary")
|
| 293 |
+
status = gr.Markdown("Choose a source and click Compute.")
|
| 294 |
+
head = gr.Slider(minimum=0, maximum=2, step=1, value=0, label="Head index (real models)")
|
| 295 |
+
|
| 296 |
+
with gr.Tab("Bulk changes"):
|
| 297 |
+
which = gr.Radio(["Q", "K"], value="Q", label="Tensor")
|
| 298 |
+
mod_2pi = gr.Checkbox(False, label="θ heatmap: wrap mod 2π")
|
| 299 |
+
bulk_main = gr.Plot(label="Before / after / delta")
|
| 300 |
+
bulk_norm = gr.Plot(label="Norms and cosine")
|
| 301 |
+
bulk_theta = gr.Plot(label="θ(k, i)")
|
| 302 |
+
bulk_freq = gr.Plot(label="ω_i")
|
| 303 |
+
|
| 304 |
+
with gr.Tab("Individual changes"):
|
| 305 |
+
with gr.Row():
|
| 306 |
+
token_k = gr.Slider(0, 15, step=1, value=0, label="Token index k")
|
| 307 |
+
pair_i = gr.Slider(0, 15, step=1, value=0, label="Pair index i")
|
| 308 |
+
sweep = gr.Checkbox(True, label="Replay the same content pair at every position k")
|
| 309 |
+
pair_table = gr.Markdown()
|
| 310 |
+
pair_plot = gr.Plot()
|
| 311 |
+
sweep_plot = gr.Plot()
|
| 312 |
+
|
| 313 |
+
with gr.Tab("Attention effect"):
|
| 314 |
+
query_token = gr.Slider(0, 15, step=1, value=0, label="Query token")
|
| 315 |
+
attn_heat = gr.Plot()
|
| 316 |
+
attn_bar = gr.Plot()
|
| 317 |
+
attn_note = gr.Markdown()
|
| 318 |
+
|
| 319 |
+
with gr.Tab("Compare to additive PE"):
|
| 320 |
+
pe_heat = gr.Plot()
|
| 321 |
+
pe_norm = gr.Plot()
|
| 322 |
+
pe_note = gr.Markdown()
|
| 323 |
+
|
| 324 |
+
compute_btn.click(
|
| 325 |
+
compute,
|
| 326 |
+
inputs=[source, sentence, model_name, seq_len, dim, seed, base],
|
| 327 |
+
outputs=[state, status, head, token_k, pair_i, query_token],
|
| 328 |
+
)
|
| 329 |
+
source.change(
|
| 330 |
+
toggle_source,
|
| 331 |
+
inputs=[source],
|
| 332 |
+
outputs=[seq_len, dim, seed, sentence, model_name],
|
| 333 |
+
)
|
| 334 |
+
|
| 335 |
+
bulk_inputs = [state, which, head, mod_2pi]
|
| 336 |
+
bulk_outputs = [bulk_main, bulk_norm, bulk_theta, bulk_freq]
|
| 337 |
+
for ctrl in bulk_inputs:
|
| 338 |
+
ctrl.change(update_bulk, inputs=bulk_inputs, outputs=bulk_outputs)
|
| 339 |
+
|
| 340 |
+
ind_inputs = [state, which, head, token_k, pair_i, sweep]
|
| 341 |
+
ind_outputs = [pair_table, pair_plot, sweep_plot]
|
| 342 |
+
for ctrl in ind_inputs:
|
| 343 |
+
ctrl.change(update_individual, inputs=ind_inputs, outputs=ind_outputs)
|
| 344 |
+
|
| 345 |
+
attn_inputs = [state, head, query_token]
|
| 346 |
+
attn_outputs = [attn_heat, attn_bar, attn_note]
|
| 347 |
+
for ctrl in attn_inputs:
|
| 348 |
+
ctrl.change(update_attention, inputs=attn_inputs, outputs=attn_outputs)
|
| 349 |
+
|
| 350 |
+
state.change(update_compare, inputs=[state], outputs=[pe_heat, pe_norm, pe_note])
|
| 351 |
+
|
| 352 |
+
if __name__ == "__main__":
|
| 353 |
+
demo.launch(server_name="0.0.0.0", server_port=7860)
|
requirements.txt
CHANGED
|
@@ -1,5 +1,7 @@
|
|
| 1 |
-
gradio==6.26.0
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
spaces==0.51.3
|
| 5 |
-
|
|
|
|
|
|
|
|
|
| 1 |
+
gradio==6.26.0
|
| 2 |
+
numpy>=2.0,<2.3
|
| 3 |
+
plotly==7.0.0
|
| 4 |
+
spaces==0.51.3
|
| 5 |
+
transformers>=4.44.0
|
| 6 |
+
accelerate>=0.34.0
|
| 7 |
+
torch>=2.2.0
|
src/absolute_pe.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Vectorized additive sinusoidal positional encoding (Vaswani et al.).
|
| 2 |
+
|
| 3 |
+
``p(k, i) = sin(k / base^{i/d})`` for even ``i``, and
|
| 4 |
+
``p(k, i) = cos(k / base^{(i-1)/d})`` for odd ``i``.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def sinusoidal_pe(seq_len: int, dim: int, base: float = 10000.0) -> np.ndarray:
|
| 13 |
+
"""Return a ``(seq_len, dim)`` additive PE matrix."""
|
| 14 |
+
if seq_len < 0 or dim < 1:
|
| 15 |
+
raise ValueError("seq_len must be >= 0 and dim >= 1")
|
| 16 |
+
positions = np.arange(seq_len, dtype=np.float64)[:, None]
|
| 17 |
+
dims = np.arange(dim, dtype=np.float64)[None, :]
|
| 18 |
+
exponent = np.where(dims % 2 == 0, dims / dim, (dims - 1.0) / dim)
|
| 19 |
+
angles = positions / (base ** exponent)
|
| 20 |
+
pe = np.empty((seq_len, dim), dtype=np.float64)
|
| 21 |
+
pe[:, 0::2] = np.sin(angles[:, 0::2])
|
| 22 |
+
pe[:, 1::2] = np.cos(angles[:, 1::2])
|
| 23 |
+
return pe
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def add_positional_encoding(
|
| 27 |
+
embeddings: np.ndarray,
|
| 28 |
+
base: float = 10000.0,
|
| 29 |
+
) -> tuple[np.ndarray, np.ndarray]:
|
| 30 |
+
"""Return ``(pe, embeddings + pe)`` for a ``(seq, dim)`` matrix."""
|
| 31 |
+
embeddings = np.asarray(embeddings, dtype=np.float64)
|
| 32 |
+
if embeddings.ndim != 2:
|
| 33 |
+
raise ValueError("embeddings must be 2D (seq, dim)")
|
| 34 |
+
pe = sinusoidal_pe(embeddings.shape[0], embeddings.shape[1], base=base)
|
| 35 |
+
return pe, embeddings + pe
|
src/extract.py
ADDED
|
@@ -0,0 +1,260 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Random tensors and lazy Hugging Face Q/K extraction.
|
| 2 |
+
|
| 3 |
+
Models are not loaded at import time. The last loaded model is cached.
|
| 4 |
+
"""
|
| 5 |
+
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
from typing import Any
|
| 9 |
+
|
| 10 |
+
import numpy as np
|
| 11 |
+
|
| 12 |
+
from src.rope import apply_rope
|
| 13 |
+
|
| 14 |
+
MAX_SEQ_LEN = 64
|
| 15 |
+
|
| 16 |
+
# Ungated Llama-like checkpoints (q_proj / k_proj + rotary). No HF token required.
|
| 17 |
+
MODEL_CHOICES = [
|
| 18 |
+
"HuggingFaceTB/SmolLM2-135M",
|
| 19 |
+
"Qwen/Qwen2.5-0.5B-Instruct",
|
| 20 |
+
"TinyLlama/TinyLlama-1.1B-Chat-v1.0",
|
| 21 |
+
]
|
| 22 |
+
DEFAULT_MODEL = MODEL_CHOICES[0]
|
| 23 |
+
|
| 24 |
+
_cache: dict[str, Any] = {"name": None, "model": None, "tokenizer": None}
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def random_matrix(seq_len: int, dim: int, seed: int = 0) -> np.ndarray:
|
| 28 |
+
if dim % 2 != 0:
|
| 29 |
+
raise ValueError(f"dim must be even for RoPE, got {dim}")
|
| 30 |
+
seq_len = int(np.clip(seq_len, 1, MAX_SEQ_LEN))
|
| 31 |
+
rng = np.random.default_rng(int(seed))
|
| 32 |
+
return rng.standard_normal((seq_len, dim)).astype(np.float64)
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def random_qk(
|
| 36 |
+
seq_len: int,
|
| 37 |
+
dim: int,
|
| 38 |
+
seed: int = 0,
|
| 39 |
+
base: float = 10000.0,
|
| 40 |
+
) -> dict[str, Any]:
|
| 41 |
+
q = random_matrix(seq_len, dim, seed=seed)
|
| 42 |
+
k = random_matrix(seq_len, dim, seed=seed + 1)
|
| 43 |
+
return _pack_tensors(
|
| 44 |
+
q_before=q,
|
| 45 |
+
k_before=k,
|
| 46 |
+
embeddings=q.copy(),
|
| 47 |
+
tokens=[f"t{i}" for i in range(q.shape[0])],
|
| 48 |
+
base=float(base),
|
| 49 |
+
style="interleaved",
|
| 50 |
+
n_q_heads=1,
|
| 51 |
+
n_kv_heads=1,
|
| 52 |
+
checksum=None,
|
| 53 |
+
source="random",
|
| 54 |
+
model_name=None,
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def _pack_tensors(
|
| 59 |
+
*,
|
| 60 |
+
q_before: np.ndarray,
|
| 61 |
+
k_before: np.ndarray,
|
| 62 |
+
embeddings: np.ndarray,
|
| 63 |
+
tokens: list[str],
|
| 64 |
+
base: float,
|
| 65 |
+
style: str,
|
| 66 |
+
n_q_heads: int,
|
| 67 |
+
n_kv_heads: int,
|
| 68 |
+
checksum: float | None,
|
| 69 |
+
source: str,
|
| 70 |
+
model_name: str | None,
|
| 71 |
+
q_after_model: np.ndarray | None = None,
|
| 72 |
+
k_after_model: np.ndarray | None = None,
|
| 73 |
+
) -> dict[str, Any]:
|
| 74 |
+
q_after = apply_rope(q_before, base=base, style=style)
|
| 75 |
+
k_after = apply_rope(k_before, base=base, style=style)
|
| 76 |
+
return {
|
| 77 |
+
"q_before": q_before,
|
| 78 |
+
"k_before": k_before,
|
| 79 |
+
"q_after": q_after,
|
| 80 |
+
"k_after": k_after,
|
| 81 |
+
"q_after_model": q_after_model,
|
| 82 |
+
"k_after_model": k_after_model,
|
| 83 |
+
"embeddings": embeddings,
|
| 84 |
+
"tokens": tokens,
|
| 85 |
+
"base": float(base),
|
| 86 |
+
"style": style,
|
| 87 |
+
"n_q_heads": int(n_q_heads),
|
| 88 |
+
"n_kv_heads": int(n_kv_heads),
|
| 89 |
+
"checksum": checksum,
|
| 90 |
+
"source": source,
|
| 91 |
+
"model_name": model_name,
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def _require_hf():
|
| 96 |
+
try:
|
| 97 |
+
import torch
|
| 98 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 99 |
+
except ImportError as exc:
|
| 100 |
+
raise RuntimeError(
|
| 101 |
+
"Real-model mode needs `torch` and `transformers`. "
|
| 102 |
+
"Install them or use Random matrix mode."
|
| 103 |
+
) from exc
|
| 104 |
+
return torch, AutoModelForCausalLM, AutoTokenizer
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def get_model(model_name: str):
|
| 108 |
+
"""Load tokenizer + causal LM on CPU; cache the last selection."""
|
| 109 |
+
torch, AutoModelForCausalLM, AutoTokenizer = _require_hf()
|
| 110 |
+
if _cache["name"] == model_name and _cache["model"] is not None:
|
| 111 |
+
return _cache["model"], _cache["tokenizer"]
|
| 112 |
+
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 113 |
+
model = AutoModelForCausalLM.from_pretrained(
|
| 114 |
+
model_name,
|
| 115 |
+
torch_dtype=torch.float32,
|
| 116 |
+
low_cpu_mem_usage=True,
|
| 117 |
+
)
|
| 118 |
+
model.eval()
|
| 119 |
+
model.to("cpu")
|
| 120 |
+
_cache["name"] = model_name
|
| 121 |
+
_cache["model"] = model
|
| 122 |
+
_cache["tokenizer"] = tokenizer
|
| 123 |
+
return model, tokenizer
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
def _backbone(model):
|
| 127 |
+
if hasattr(model, "model") and hasattr(model.model, "embed_tokens"):
|
| 128 |
+
return model.model
|
| 129 |
+
raise RuntimeError(
|
| 130 |
+
"This checkpoint is not Llama-like (expected model.model.embed_tokens)."
|
| 131 |
+
)
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def _rotary_module(backbone, attn):
|
| 135 |
+
if hasattr(backbone, "rotary_emb"):
|
| 136 |
+
return backbone.rotary_emb
|
| 137 |
+
if hasattr(attn, "rotary_emb"):
|
| 138 |
+
return attn.rotary_emb
|
| 139 |
+
return None
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def _rotate_half_torch(x, torch):
|
| 143 |
+
x1 = x[..., : x.shape[-1] // 2]
|
| 144 |
+
x2 = x[..., x.shape[-1] // 2 :]
|
| 145 |
+
return torch.cat((-x2, x1), dim=-1)
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
def _broadcast_cos_sin(cos, sin, q):
|
| 149 |
+
"""Make cos/sin broadcast with q of shape (batch, heads, seq, dim)."""
|
| 150 |
+
if cos.dim() == 2:
|
| 151 |
+
cos, sin = cos.unsqueeze(0).unsqueeze(0), sin.unsqueeze(0).unsqueeze(0)
|
| 152 |
+
elif cos.dim() == 3:
|
| 153 |
+
cos, sin = cos.unsqueeze(1), sin.unsqueeze(1)
|
| 154 |
+
return cos, sin
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def _hf_rotary(q, k, rotary, position_ids, torch):
|
| 158 |
+
try:
|
| 159 |
+
cos, sin = rotary(q, position_ids=position_ids)
|
| 160 |
+
except TypeError:
|
| 161 |
+
try:
|
| 162 |
+
cos, sin = rotary(q, seq_len=q.shape[-2])
|
| 163 |
+
except TypeError:
|
| 164 |
+
cos, sin = rotary(q)
|
| 165 |
+
if cos.shape[-1] == q.shape[-1] // 2:
|
| 166 |
+
cos = torch.cat((cos, cos), dim=-1)
|
| 167 |
+
sin = torch.cat((sin, sin), dim=-1)
|
| 168 |
+
cos, sin = _broadcast_cos_sin(cos, sin, q)
|
| 169 |
+
q_rot = (q * cos) + (_rotate_half_torch(q, torch) * sin)
|
| 170 |
+
k_rot = (k * cos) + (_rotate_half_torch(k, torch) * sin)
|
| 171 |
+
return q_rot, k_rot
|
| 172 |
+
|
| 173 |
+
|
| 174 |
+
def extract_from_model(model_name: str, sentence: str) -> dict[str, Any]:
|
| 175 |
+
torch, _, _ = _require_hf()
|
| 176 |
+
model, tokenizer = get_model(model_name)
|
| 177 |
+
text = sentence.strip() or "RoPE rotates query and key vectors."
|
| 178 |
+
encoded = tokenizer(
|
| 179 |
+
text,
|
| 180 |
+
return_tensors="pt",
|
| 181 |
+
truncation=True,
|
| 182 |
+
max_length=MAX_SEQ_LEN,
|
| 183 |
+
add_special_tokens=True,
|
| 184 |
+
)
|
| 185 |
+
input_ids = encoded["input_ids"]
|
| 186 |
+
tokens = tokenizer.convert_ids_to_tokens(input_ids[0].tolist())
|
| 187 |
+
backbone = _backbone(model)
|
| 188 |
+
cfg = model.config
|
| 189 |
+
n_heads = int(cfg.num_attention_heads)
|
| 190 |
+
n_kv = int(getattr(cfg, "num_key_value_heads", n_heads))
|
| 191 |
+
hidden_size = int(cfg.hidden_size)
|
| 192 |
+
head_dim = int(getattr(cfg, "head_dim", hidden_size // n_heads))
|
| 193 |
+
base = float(getattr(cfg, "rope_theta", 10000.0))
|
| 194 |
+
|
| 195 |
+
with torch.no_grad():
|
| 196 |
+
embeds = backbone.embed_tokens(input_ids)
|
| 197 |
+
layer = backbone.layers[0]
|
| 198 |
+
attn = layer.self_attn
|
| 199 |
+
hidden = layer.input_layernorm(embeds)
|
| 200 |
+
q = attn.q_proj(hidden)
|
| 201 |
+
k = attn.k_proj(hidden)
|
| 202 |
+
seq = q.shape[1]
|
| 203 |
+
q = q.view(1, seq, n_heads, head_dim).transpose(1, 2).contiguous()
|
| 204 |
+
k = k.view(1, seq, n_kv, head_dim).transpose(1, 2).contiguous()
|
| 205 |
+
position_ids = torch.arange(seq).unsqueeze(0)
|
| 206 |
+
rotary = _rotary_module(backbone, attn)
|
| 207 |
+
q_model = k_model = None
|
| 208 |
+
if rotary is not None:
|
| 209 |
+
q_model, k_model = _hf_rotary(q, k, rotary, position_ids, torch)
|
| 210 |
+
|
| 211 |
+
q_np = q[0].cpu().numpy().astype(np.float64)
|
| 212 |
+
k_np = k[0].cpu().numpy().astype(np.float64)
|
| 213 |
+
emb_np = embeds[0].cpu().numpy().astype(np.float64)
|
| 214 |
+
q_model_np = k_model_np = None
|
| 215 |
+
checksum = None
|
| 216 |
+
if q_model is not None:
|
| 217 |
+
q_model_np = q_model[0].cpu().numpy().astype(np.float64)
|
| 218 |
+
k_model_np = k_model[0].cpu().numpy().astype(np.float64)
|
| 219 |
+
q_edu = apply_rope(q_np, base=base, style="llama")
|
| 220 |
+
checksum = float(np.max(np.abs(q_edu - q_model_np)))
|
| 221 |
+
|
| 222 |
+
return _pack_tensors(
|
| 223 |
+
q_before=q_np,
|
| 224 |
+
k_before=k_np,
|
| 225 |
+
embeddings=emb_np,
|
| 226 |
+
tokens=tokens,
|
| 227 |
+
base=base,
|
| 228 |
+
style="llama",
|
| 229 |
+
n_q_heads=n_heads,
|
| 230 |
+
n_kv_heads=n_kv,
|
| 231 |
+
checksum=checksum,
|
| 232 |
+
source="model",
|
| 233 |
+
model_name=model_name,
|
| 234 |
+
q_after_model=q_model_np,
|
| 235 |
+
k_after_model=k_model_np,
|
| 236 |
+
)
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def select_head(tensor: np.ndarray, head: int) -> np.ndarray:
|
| 240 |
+
"""Return a ``(seq, dim)`` slice; ``tensor`` is 2D or ``(heads, seq, dim)``."""
|
| 241 |
+
t = np.asarray(tensor)
|
| 242 |
+
if t.ndim == 2:
|
| 243 |
+
return t
|
| 244 |
+
if t.ndim != 3:
|
| 245 |
+
raise ValueError(f"expected 2D or 3D tensor, got {t.ndim}D")
|
| 246 |
+
h = int(np.clip(head, 0, t.shape[0] - 1))
|
| 247 |
+
return t[h]
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
def expand_kv_heads(k: np.ndarray, n_q_heads: int) -> np.ndarray:
|
| 251 |
+
"""Repeat GQA key heads so they align with query heads."""
|
| 252 |
+
t = np.asarray(k)
|
| 253 |
+
if t.ndim == 2:
|
| 254 |
+
return t
|
| 255 |
+
n_kv = t.shape[0]
|
| 256 |
+
if n_kv == n_q_heads:
|
| 257 |
+
return t
|
| 258 |
+
if n_q_heads % n_kv != 0:
|
| 259 |
+
return t
|
| 260 |
+
return np.repeat(t, n_q_heads // n_kv, axis=0)
|
src/plots.py
ADDED
|
@@ -0,0 +1,305 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Plotly figures for bulk heatmaps, pair rotation, and attention."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import plotly.graph_objects as go
|
| 7 |
+
from plotly.subplots import make_subplots
|
| 8 |
+
|
| 9 |
+
from src.rope import l2_norms, pair_frequencies, pair_xy, row_cosine, theta_grid
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def _empty(title: str) -> go.Figure:
|
| 13 |
+
fig = go.Figure()
|
| 14 |
+
fig.update_layout(title=title, template="plotly_white")
|
| 15 |
+
return fig
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def downsample(mat: np.ndarray, max_cols: int = 64, max_rows: int = 64) -> np.ndarray:
|
| 19 |
+
m = np.asarray(mat)
|
| 20 |
+
row_step = max(1, int(np.ceil(m.shape[0] / max_rows)))
|
| 21 |
+
if m.ndim == 1:
|
| 22 |
+
return m[::row_step]
|
| 23 |
+
col_step = max(1, int(np.ceil(m.shape[1] / max_cols)))
|
| 24 |
+
return m[::row_step, ::col_step]
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def heatmap(z: np.ndarray, title: str, xtitle: str, ytitle: str, tokens=None) -> go.Figure:
|
| 28 |
+
z = downsample(np.asarray(z, dtype=np.float64))
|
| 29 |
+
fig = go.Figure(
|
| 30 |
+
data=go.Heatmap(
|
| 31 |
+
z=z,
|
| 32 |
+
colorbar=dict(title="value"),
|
| 33 |
+
hovertemplate="y=%{y}<br>x=%{x}<br>z=%{z:.4f}<extra></extra>",
|
| 34 |
+
)
|
| 35 |
+
)
|
| 36 |
+
fig.update_layout(
|
| 37 |
+
title=title,
|
| 38 |
+
xaxis_title=xtitle,
|
| 39 |
+
yaxis_title=ytitle,
|
| 40 |
+
yaxis=dict(autorange="reversed"),
|
| 41 |
+
template="plotly_white",
|
| 42 |
+
height=380,
|
| 43 |
+
margin=dict(l=60, r=40, t=50, b=50),
|
| 44 |
+
)
|
| 45 |
+
if tokens is not None and len(tokens) == z.shape[0]:
|
| 46 |
+
fig.update_yaxes(tickmode="array", tickvals=list(range(len(tokens))), ticktext=tokens)
|
| 47 |
+
return fig
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def bulk_before_after_delta(before: np.ndarray, after: np.ndarray, tokens=None) -> go.Figure:
|
| 51 |
+
before = np.asarray(before, dtype=np.float64)
|
| 52 |
+
after = np.asarray(after, dtype=np.float64)
|
| 53 |
+
delta = after - before
|
| 54 |
+
mats = [downsample(before), downsample(after), downsample(delta)]
|
| 55 |
+
titles = ["Before RoPE", "After RoPE", "Delta (after − before)"]
|
| 56 |
+
fig = make_subplots(rows=1, cols=3, subplot_titles=titles)
|
| 57 |
+
for i, mat in enumerate(mats, start=1):
|
| 58 |
+
fig.add_trace(
|
| 59 |
+
go.Heatmap(
|
| 60 |
+
z=mat,
|
| 61 |
+
showscale=(i == 3),
|
| 62 |
+
hovertemplate="token=%{y}<br>dim=%{x}<br>z=%{z:.4f}<extra></extra>",
|
| 63 |
+
),
|
| 64 |
+
row=1,
|
| 65 |
+
col=i,
|
| 66 |
+
)
|
| 67 |
+
fig.update_yaxes(autorange="reversed", row=1, col=i)
|
| 68 |
+
fig.update_layout(template="plotly_white", height=400, margin=dict(t=60))
|
| 69 |
+
return fig
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def norms_and_cosine(before: np.ndarray, after: np.ndarray, tokens=None) -> go.Figure:
|
| 73 |
+
nb = l2_norms(before)
|
| 74 |
+
na = l2_norms(after)
|
| 75 |
+
cos = row_cosine(before, after)
|
| 76 |
+
xs = list(range(len(nb)))
|
| 77 |
+
fig = make_subplots(rows=1, cols=2, subplot_titles=["Per-token L2 norm", "Cosine (original vs rotated)"])
|
| 78 |
+
fig.add_trace(go.Scatter(x=xs, y=nb, name="before", mode="lines+markers"), row=1, col=1)
|
| 79 |
+
fig.add_trace(go.Scatter(x=xs, y=na, name="after", mode="lines+markers"), row=1, col=1)
|
| 80 |
+
fig.add_trace(go.Bar(x=xs, y=cos, name="cosine", showlegend=False), row=1, col=2)
|
| 81 |
+
fig.update_xaxes(title_text="token", row=1, col=1)
|
| 82 |
+
fig.update_xaxes(title_text="token", row=1, col=2)
|
| 83 |
+
fig.update_yaxes(title_text="L2", row=1, col=1)
|
| 84 |
+
fig.update_yaxes(title_text="cosine", range=[min(0.0, float(np.min(cos)) - 0.05), 1.02], row=1, col=2)
|
| 85 |
+
fig.update_layout(template="plotly_white", height=380, barmode="group")
|
| 86 |
+
return fig
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def theta_heatmap(seq_len: int, dim: int, base: float, mod_2pi: bool = False) -> go.Figure:
|
| 90 |
+
grid = theta_grid(seq_len, dim, base=base)
|
| 91 |
+
if mod_2pi:
|
| 92 |
+
grid = np.mod(grid, 2 * np.pi)
|
| 93 |
+
title = "θ(k, i) mod 2π"
|
| 94 |
+
else:
|
| 95 |
+
title = "θ(k, i) = k · ω_i"
|
| 96 |
+
z = downsample(grid, max_cols=64, max_rows=64)
|
| 97 |
+
fig = go.Figure(
|
| 98 |
+
data=go.Heatmap(
|
| 99 |
+
z=z,
|
| 100 |
+
colorbar=dict(title="radians"),
|
| 101 |
+
hovertemplate="token k=%{y}<br>pair i=%{x}<br>θ=%{z:.4f}<extra></extra>",
|
| 102 |
+
)
|
| 103 |
+
)
|
| 104 |
+
fig.update_layout(
|
| 105 |
+
title=title,
|
| 106 |
+
xaxis_title="pair index i",
|
| 107 |
+
yaxis_title="token position k",
|
| 108 |
+
yaxis=dict(autorange="reversed"),
|
| 109 |
+
template="plotly_white",
|
| 110 |
+
height=380,
|
| 111 |
+
)
|
| 112 |
+
return fig
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def frequency_strip(dim: int, base: float) -> go.Figure:
|
| 116 |
+
omega = pair_frequencies(dim, base=base)
|
| 117 |
+
fig = go.Figure(data=go.Bar(x=list(range(len(omega))), y=omega, name="ω_i"))
|
| 118 |
+
fig.update_layout(
|
| 119 |
+
title="Pair frequencies ω_i (pair 0 is fastest)",
|
| 120 |
+
xaxis_title="pair index i",
|
| 121 |
+
yaxis_title="ω",
|
| 122 |
+
yaxis_type="log",
|
| 123 |
+
template="plotly_white",
|
| 124 |
+
height=280,
|
| 125 |
+
)
|
| 126 |
+
return fig
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def _arc_points(x0, y0, x1, y1, n=40):
|
| 130 |
+
r0 = float(np.hypot(x0, y0))
|
| 131 |
+
r1 = float(np.hypot(x1, y1))
|
| 132 |
+
r = 0.5 * (r0 + r1)
|
| 133 |
+
if r < 1e-9:
|
| 134 |
+
return [], []
|
| 135 |
+
a0 = float(np.arctan2(y0, x0))
|
| 136 |
+
a1 = float(np.arctan2(y1, x1))
|
| 137 |
+
delta = (a1 - a0 + np.pi) % (2 * np.pi) - np.pi
|
| 138 |
+
ts = np.linspace(0.0, 1.0, n)
|
| 139 |
+
angs = a0 + ts * delta
|
| 140 |
+
rr = r * 0.55
|
| 141 |
+
return list(rr * np.cos(angs)), list(rr * np.sin(angs))
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def rotation_2d(
|
| 145 |
+
matrix_before: np.ndarray,
|
| 146 |
+
matrix_after: np.ndarray,
|
| 147 |
+
token: int,
|
| 148 |
+
pair: int,
|
| 149 |
+
style: str,
|
| 150 |
+
theta: float,
|
| 151 |
+
neighbor_tokens: list[int] | None = None,
|
| 152 |
+
) -> go.Figure:
|
| 153 |
+
xb, yb = pair_xy(matrix_before, token, pair, style=style)
|
| 154 |
+
xa, ya = pair_xy(matrix_after, token, pair, style=style)
|
| 155 |
+
fig = go.Figure()
|
| 156 |
+
fig.add_trace(
|
| 157 |
+
go.Scatter(
|
| 158 |
+
x=[0, xb],
|
| 159 |
+
y=[0, yb],
|
| 160 |
+
mode="lines+markers",
|
| 161 |
+
name="before",
|
| 162 |
+
line=dict(width=3),
|
| 163 |
+
)
|
| 164 |
+
)
|
| 165 |
+
fig.add_trace(
|
| 166 |
+
go.Scatter(
|
| 167 |
+
x=[0, xa],
|
| 168 |
+
y=[0, ya],
|
| 169 |
+
mode="lines+markers",
|
| 170 |
+
name="after",
|
| 171 |
+
line=dict(width=3),
|
| 172 |
+
)
|
| 173 |
+
)
|
| 174 |
+
xs, ys = _arc_points(xb, yb, xa, ya)
|
| 175 |
+
if xs:
|
| 176 |
+
fig.add_trace(
|
| 177 |
+
go.Scatter(x=xs, y=ys, mode="lines", name=f"θ={theta:.3f} rad", line=dict(dash="dot"))
|
| 178 |
+
)
|
| 179 |
+
if neighbor_tokens:
|
| 180 |
+
for n in neighbor_tokens:
|
| 181 |
+
if n == token:
|
| 182 |
+
continue
|
| 183 |
+
nx, ny = pair_xy(matrix_after, n, pair, style=style)
|
| 184 |
+
fig.add_trace(
|
| 185 |
+
go.Scatter(
|
| 186 |
+
x=[0, nx],
|
| 187 |
+
y=[0, ny],
|
| 188 |
+
mode="lines+markers",
|
| 189 |
+
name=f"token {n} after",
|
| 190 |
+
opacity=0.45,
|
| 191 |
+
)
|
| 192 |
+
)
|
| 193 |
+
lim = max(abs(xb), abs(yb), abs(xa), abs(ya), 1e-3) * 1.3
|
| 194 |
+
fig.update_layout(
|
| 195 |
+
title=f"Pair {pair} at token {token}: 2D rotation",
|
| 196 |
+
xaxis=dict(scaleanchor="y", scaleratio=1, range=[-lim, lim], zeroline=True),
|
| 197 |
+
yaxis=dict(range=[-lim, lim], zeroline=True),
|
| 198 |
+
template="plotly_white",
|
| 199 |
+
height=460,
|
| 200 |
+
legend=dict(orientation="h"),
|
| 201 |
+
)
|
| 202 |
+
return fig
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
def position_sweep(
|
| 206 |
+
x_even: float,
|
| 207 |
+
x_odd: float,
|
| 208 |
+
omega: float,
|
| 209 |
+
seq_len: int,
|
| 210 |
+
highlight: int,
|
| 211 |
+
) -> go.Figure:
|
| 212 |
+
"""Same content pair spun by θ(k)=k ω — position is the only change."""
|
| 213 |
+
ks = np.arange(seq_len)
|
| 214 |
+
thetas = ks * omega
|
| 215 |
+
xs = x_even * np.cos(thetas) - x_odd * np.sin(thetas)
|
| 216 |
+
ys = x_even * np.sin(thetas) + x_odd * np.cos(thetas)
|
| 217 |
+
fig = go.Figure()
|
| 218 |
+
fig.add_trace(go.Scatter(x=xs, y=ys, mode="markers+lines", name="path over k"))
|
| 219 |
+
fig.add_trace(
|
| 220 |
+
go.Scatter(
|
| 221 |
+
x=[0, xs[highlight]],
|
| 222 |
+
y=[0, ys[highlight]],
|
| 223 |
+
mode="lines+markers",
|
| 224 |
+
name=f"k={highlight}",
|
| 225 |
+
line=dict(width=3),
|
| 226 |
+
)
|
| 227 |
+
)
|
| 228 |
+
lim = max(float(np.max(np.abs(xs))), float(np.max(np.abs(ys))), 1e-3) * 1.3
|
| 229 |
+
fig.update_layout(
|
| 230 |
+
title="Same (x_even, x_odd) rotated at every position k",
|
| 231 |
+
xaxis=dict(scaleanchor="y", scaleratio=1, range=[-lim, lim]),
|
| 232 |
+
yaxis=dict(range=[-lim, lim]),
|
| 233 |
+
template="plotly_white",
|
| 234 |
+
height=420,
|
| 235 |
+
)
|
| 236 |
+
return fig
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def attention_heatmaps(scores_before: np.ndarray, scores_after: np.ndarray) -> go.Figure:
|
| 240 |
+
fig = make_subplots(rows=1, cols=2, subplot_titles=["QKᵀ without RoPE", "QKᵀ with RoPE"])
|
| 241 |
+
for i, mat in enumerate([scores_before, scores_after], start=1):
|
| 242 |
+
fig.add_trace(
|
| 243 |
+
go.Heatmap(
|
| 244 |
+
z=downsample(mat),
|
| 245 |
+
showscale=(i == 2),
|
| 246 |
+
hovertemplate="query=%{y}<br>key=%{x}<br>score=%{z:.4f}<extra></extra>",
|
| 247 |
+
),
|
| 248 |
+
row=1,
|
| 249 |
+
col=i,
|
| 250 |
+
)
|
| 251 |
+
fig.update_yaxes(autorange="reversed", title_text="query token", row=1, col=i)
|
| 252 |
+
fig.update_xaxes(title_text="key token", row=1, col=i)
|
| 253 |
+
fig.update_layout(template="plotly_white", height=420)
|
| 254 |
+
return fig
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
def attention_bars(logits_before: np.ndarray, logits_after: np.ndarray, query_token: int) -> go.Figure:
|
| 258 |
+
xs = list(range(len(logits_before)))
|
| 259 |
+
fig = go.Figure()
|
| 260 |
+
fig.add_trace(go.Bar(x=xs, y=logits_before, name="without RoPE"))
|
| 261 |
+
fig.add_trace(go.Bar(x=xs, y=logits_after, name="with RoPE"))
|
| 262 |
+
fig.update_layout(
|
| 263 |
+
title=f"Attention logits from query token {query_token}",
|
| 264 |
+
xaxis_title="key token",
|
| 265 |
+
yaxis_title="Q·K",
|
| 266 |
+
barmode="group",
|
| 267 |
+
template="plotly_white",
|
| 268 |
+
height=360,
|
| 269 |
+
)
|
| 270 |
+
return fig
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
def additive_pe_heatmaps(emb: np.ndarray, pe: np.ndarray, combined: np.ndarray) -> go.Figure:
|
| 274 |
+
titles = ["Token embeddings", "Additive sinusoidal PE", "Embeddings + PE"]
|
| 275 |
+
fig = make_subplots(rows=1, cols=3, subplot_titles=titles)
|
| 276 |
+
for i, mat in enumerate([emb, pe, combined], start=1):
|
| 277 |
+
fig.add_trace(
|
| 278 |
+
go.Heatmap(z=downsample(mat), showscale=(i == 3)),
|
| 279 |
+
row=1,
|
| 280 |
+
col=i,
|
| 281 |
+
)
|
| 282 |
+
fig.update_yaxes(autorange="reversed", row=1, col=i)
|
| 283 |
+
fig.update_layout(template="plotly_white", height=400)
|
| 284 |
+
return fig
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
def norm_compare_add_vs_rope(
|
| 288 |
+
emb: np.ndarray,
|
| 289 |
+
emb_plus_pe: np.ndarray,
|
| 290 |
+
q_before: np.ndarray,
|
| 291 |
+
q_after: np.ndarray,
|
| 292 |
+
) -> go.Figure:
|
| 293 |
+
fig = go.Figure()
|
| 294 |
+
fig.add_trace(go.Scatter(y=l2_norms(emb), mode="lines+markers", name="embeddings"))
|
| 295 |
+
fig.add_trace(go.Scatter(y=l2_norms(emb_plus_pe), mode="lines+markers", name="embeddings + PE"))
|
| 296 |
+
fig.add_trace(go.Scatter(y=l2_norms(q_before), mode="lines+markers", name="Q before RoPE"))
|
| 297 |
+
fig.add_trace(go.Scatter(y=l2_norms(q_after), mode="lines+markers", name="Q after RoPE"))
|
| 298 |
+
fig.update_layout(
|
| 299 |
+
title="L2 norms: additive PE changes magnitude; RoPE does not",
|
| 300 |
+
xaxis_title="token",
|
| 301 |
+
yaxis_title="L2",
|
| 302 |
+
template="plotly_white",
|
| 303 |
+
height=360,
|
| 304 |
+
)
|
| 305 |
+
return fig
|
src/rope.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Vectorized RoPE: pairwise 2D rotations of query/key slices.
|
| 2 |
+
|
| 3 |
+
Two pairing conventions:
|
| 4 |
+
|
| 5 |
+
- ``interleaved`` — paper / GPT-J style: rotate dims ``(2i, 2i+1)``.
|
| 6 |
+
- ``llama`` — Hugging Face Llama/Qwen/SmolLM style: rotate
|
| 7 |
+
``(i, i + dim/2)`` (the ``rotate_half`` layout).
|
| 8 |
+
|
| 9 |
+
Frequencies match the usual schedule ``ω_i = base^{-2i/d}``.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
from __future__ import annotations
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
|
| 16 |
+
VALID_STYLES = ("interleaved", "llama")
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def pair_frequencies(dim: int, base: float = 10000.0) -> np.ndarray:
|
| 20 |
+
"""Return ``ω`` of shape ``(dim/2,)`` with ``ω_i = base ** (-2i / dim)``."""
|
| 21 |
+
if dim % 2 != 0:
|
| 22 |
+
raise ValueError(f"RoPE requires an even dim, got {dim}")
|
| 23 |
+
n_pairs = dim // 2
|
| 24 |
+
pair_idx = np.arange(n_pairs, dtype=np.float64) * 2.0
|
| 25 |
+
return 1.0 / (base ** (pair_idx / dim))
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def theta_grid(seq_len: int, dim: int, base: float = 10000.0) -> np.ndarray:
|
| 29 |
+
"""Rotation angles ``θ[k, i] = k * ω_i``, shape ``(seq_len, dim/2)``."""
|
| 30 |
+
omega = pair_frequencies(dim, base=base)
|
| 31 |
+
positions = np.arange(seq_len, dtype=np.float64)[:, None]
|
| 32 |
+
return positions * omega[None, :]
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def _broadcast_angles(angles: np.ndarray, like: np.ndarray) -> np.ndarray:
|
| 36 |
+
"""Broadcast ``(seq, n_pairs)`` onto ``like`` with shape ``(..., seq, n_pairs)``."""
|
| 37 |
+
lead = like.ndim - 2
|
| 38 |
+
return angles.reshape((1,) * lead + angles.shape)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def rotate_pair(x_even: np.ndarray, x_odd: np.ndarray, theta: np.ndarray):
|
| 42 |
+
"""Apply a 2D rotation by ``theta`` (radians) to coordinate pairs."""
|
| 43 |
+
cos = np.cos(theta)
|
| 44 |
+
sin = np.sin(theta)
|
| 45 |
+
return x_even * cos - x_odd * sin, x_even * sin + x_odd * cos
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def apply_rope(
|
| 49 |
+
x: np.ndarray,
|
| 50 |
+
base: float = 10000.0,
|
| 51 |
+
style: str = "interleaved",
|
| 52 |
+
) -> np.ndarray:
|
| 53 |
+
"""Rotate last-dimension pairs of ``x`` (``(..., seq_len, dim)``)."""
|
| 54 |
+
if style not in VALID_STYLES:
|
| 55 |
+
raise ValueError(f"style must be one of {VALID_STYLES}, got {style!r}")
|
| 56 |
+
x = np.asarray(x, dtype=np.float64)
|
| 57 |
+
if x.ndim < 2:
|
| 58 |
+
raise ValueError("x must be at least 2D (seq, dim)")
|
| 59 |
+
dim = x.shape[-1]
|
| 60 |
+
seq_len = x.shape[-2]
|
| 61 |
+
angles = theta_grid(seq_len, dim, base=base)
|
| 62 |
+
cos = np.cos(angles)
|
| 63 |
+
sin = np.sin(angles)
|
| 64 |
+
n_pairs = dim // 2
|
| 65 |
+
out = np.empty_like(x)
|
| 66 |
+
|
| 67 |
+
if style == "interleaved":
|
| 68 |
+
even = x[..., 0::2]
|
| 69 |
+
odd = x[..., 1::2]
|
| 70 |
+
cos_b = _broadcast_angles(cos, even)
|
| 71 |
+
sin_b = _broadcast_angles(sin, odd)
|
| 72 |
+
out[..., 0::2] = even * cos_b - odd * sin_b
|
| 73 |
+
out[..., 1::2] = even * sin_b + odd * cos_b
|
| 74 |
+
else:
|
| 75 |
+
first = x[..., :n_pairs]
|
| 76 |
+
second = x[..., n_pairs:]
|
| 77 |
+
cos_b = _broadcast_angles(cos, first)
|
| 78 |
+
sin_b = _broadcast_angles(sin, second)
|
| 79 |
+
out[..., :n_pairs] = first * cos_b - second * sin_b
|
| 80 |
+
out[..., n_pairs:] = first * sin_b + second * cos_b
|
| 81 |
+
return out
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def pair_xy(x: np.ndarray, token: int, pair: int, style: str = "interleaved") -> tuple[float, float]:
|
| 85 |
+
"""Return the 2D coordinates of one RoPE pair at one token."""
|
| 86 |
+
vec = np.asarray(x)
|
| 87 |
+
if vec.ndim != 2:
|
| 88 |
+
raise ValueError("pair_xy expects (seq, dim)")
|
| 89 |
+
dim = vec.shape[-1]
|
| 90 |
+
n_pairs = dim // 2
|
| 91 |
+
if not (0 <= pair < n_pairs):
|
| 92 |
+
raise IndexError(f"pair {pair} out of range 0..{n_pairs - 1}")
|
| 93 |
+
if style == "interleaved":
|
| 94 |
+
return float(vec[token, 2 * pair]), float(vec[token, 2 * pair + 1])
|
| 95 |
+
if style == "llama":
|
| 96 |
+
return float(vec[token, pair]), float(vec[token, pair + n_pairs])
|
| 97 |
+
raise ValueError(f"unknown style {style!r}")
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def pair_dim_labels(pair: int, dim: int, style: str = "interleaved") -> tuple[str, str]:
|
| 101 |
+
n_pairs = dim // 2
|
| 102 |
+
if style == "interleaved":
|
| 103 |
+
return f"dim {2 * pair}", f"dim {2 * pair + 1}"
|
| 104 |
+
return f"dim {pair}", f"dim {pair + n_pairs}"
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def l2_norms(x: np.ndarray) -> np.ndarray:
|
| 108 |
+
return np.linalg.norm(np.asarray(x, dtype=np.float64), axis=-1)
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def row_cosine(a: np.ndarray, b: np.ndarray) -> np.ndarray:
|
| 112 |
+
a = np.asarray(a, dtype=np.float64)
|
| 113 |
+
b = np.asarray(b, dtype=np.float64)
|
| 114 |
+
an = l2_norms(a)
|
| 115 |
+
bn = l2_norms(b)
|
| 116 |
+
dots = np.sum(a * b, axis=-1)
|
| 117 |
+
return dots / (an * bn + 1e-12)
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def attention_scores(q: np.ndarray, k: np.ndarray) -> np.ndarray:
|
| 121 |
+
"""Unnormalized ``Q K^T`` for ``(seq, dim)`` slices."""
|
| 122 |
+
q = np.asarray(q, dtype=np.float64)
|
| 123 |
+
k = np.asarray(k, dtype=np.float64)
|
| 124 |
+
return q @ k.T
|
tests/test_core.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import unittest
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
|
| 5 |
+
from src.absolute_pe import add_positional_encoding, sinusoidal_pe
|
| 6 |
+
from src.extract import random_qk, select_head
|
| 7 |
+
from src.rope import (
|
| 8 |
+
apply_rope,
|
| 9 |
+
attention_scores,
|
| 10 |
+
l2_norms,
|
| 11 |
+
pair_frequencies,
|
| 12 |
+
pair_xy,
|
| 13 |
+
rotate_pair,
|
| 14 |
+
row_cosine,
|
| 15 |
+
)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class RopeTests(unittest.TestCase):
|
| 19 |
+
def test_even_dim_required(self):
|
| 20 |
+
with self.assertRaises(ValueError):
|
| 21 |
+
apply_rope(np.ones((4, 5)))
|
| 22 |
+
|
| 23 |
+
def test_position_zero_is_identity(self):
|
| 24 |
+
x = np.random.default_rng(0).standard_normal((8, 16))
|
| 25 |
+
y = apply_rope(x, base=10000.0, style="interleaved")
|
| 26 |
+
np.testing.assert_allclose(y[0], x[0], atol=1e-12)
|
| 27 |
+
|
| 28 |
+
def test_preserves_l2(self):
|
| 29 |
+
x = np.random.default_rng(1).standard_normal((12, 32))
|
| 30 |
+
for style in ("interleaved", "llama"):
|
| 31 |
+
y = apply_rope(x, style=style)
|
| 32 |
+
np.testing.assert_allclose(l2_norms(x), l2_norms(y), atol=1e-10)
|
| 33 |
+
|
| 34 |
+
def test_matches_scratch_even_odd(self):
|
| 35 |
+
x = np.random.default_rng(2).standard_normal((6, 10))
|
| 36 |
+
seq, dim = x.shape
|
| 37 |
+
positions = np.arange(seq)[:, None]
|
| 38 |
+
pair_indices = np.arange(0, dim, 2)
|
| 39 |
+
inv_freq = 1 / (10000 ** (pair_indices / dim))
|
| 40 |
+
angles = positions * inv_freq
|
| 41 |
+
cos, sin = np.cos(angles), np.sin(angles)
|
| 42 |
+
expected = np.empty_like(x)
|
| 43 |
+
expected[:, 0::2] = x[:, 0::2] * cos - x[:, 1::2] * sin
|
| 44 |
+
expected[:, 1::2] = x[:, 0::2] * sin + x[:, 1::2] * cos
|
| 45 |
+
np.testing.assert_allclose(apply_rope(x, style="interleaved"), expected)
|
| 46 |
+
|
| 47 |
+
def test_relative_angle_depends_on_offset(self):
|
| 48 |
+
omega = pair_frequencies(8, base=10000.0)[0]
|
| 49 |
+
x = np.zeros((5, 8))
|
| 50 |
+
x[:, 0] = 1.0
|
| 51 |
+
y = apply_rope(x, style="interleaved")
|
| 52 |
+
xm, ym = pair_xy(y, 3, 0)
|
| 53 |
+
xn, yn = pair_xy(y, 1, 0)
|
| 54 |
+
a_m = np.arctan2(ym, xm)
|
| 55 |
+
a_n = np.arctan2(yn, xn)
|
| 56 |
+
self.assertAlmostEqual(a_m - a_n, (3 - 1) * omega, places=10)
|
| 57 |
+
|
| 58 |
+
def test_rotate_pair_and_batched_heads(self):
|
| 59 |
+
even, odd = np.array([1.0]), np.array([0.0])
|
| 60 |
+
re, ro = rotate_pair(even, odd, np.array([np.pi / 2]))
|
| 61 |
+
np.testing.assert_allclose(re, 0.0, atol=1e-12)
|
| 62 |
+
np.testing.assert_allclose(ro, 1.0, atol=1e-12)
|
| 63 |
+
x = np.random.default_rng(3).standard_normal((4, 7, 16))
|
| 64 |
+
y = apply_rope(x, style="llama")
|
| 65 |
+
self.assertEqual(y.shape, x.shape)
|
| 66 |
+
np.testing.assert_allclose(l2_norms(x), l2_norms(y), atol=1e-10)
|
| 67 |
+
|
| 68 |
+
def test_cosine_identity_at_zero(self):
|
| 69 |
+
x = np.random.default_rng(4).standard_normal((5, 12))
|
| 70 |
+
y = apply_rope(x)
|
| 71 |
+
np.testing.assert_allclose(row_cosine(x, y)[0], 1.0, atol=1e-10)
|
| 72 |
+
|
| 73 |
+
def test_attention_scores_shape(self):
|
| 74 |
+
q = np.ones((3, 4))
|
| 75 |
+
k = np.ones((3, 4))
|
| 76 |
+
s = attention_scores(q, k)
|
| 77 |
+
self.assertEqual(s.shape, (3, 3))
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
class AbsolutePeTests(unittest.TestCase):
|
| 81 |
+
def test_matches_loop_formula(self):
|
| 82 |
+
seq, dim, base = 4, 5, 10000.0
|
| 83 |
+
pe = sinusoidal_pe(seq, dim, base=base)
|
| 84 |
+
for k in range(seq):
|
| 85 |
+
for i in range(dim):
|
| 86 |
+
if i % 2 == 0:
|
| 87 |
+
expected = np.sin(k / base ** (i / dim))
|
| 88 |
+
else:
|
| 89 |
+
expected = np.cos(k / base ** ((i - 1) / dim))
|
| 90 |
+
self.assertAlmostEqual(pe[k, i], expected, places=12)
|
| 91 |
+
|
| 92 |
+
def test_add_changes_norm(self):
|
| 93 |
+
emb = np.random.default_rng(0).standard_normal((6, 8))
|
| 94 |
+
pe, combined = add_positional_encoding(emb)
|
| 95 |
+
self.assertEqual(pe.shape, emb.shape)
|
| 96 |
+
self.assertFalse(np.allclose(l2_norms(emb), l2_norms(combined)))
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
class ExtractHelpersTests(unittest.TestCase):
|
| 100 |
+
def test_random_qk(self):
|
| 101 |
+
data = random_qk(seq_len=8, dim=16, seed=0, base=10000.0)
|
| 102 |
+
self.assertEqual(data["q_before"].shape, (8, 16))
|
| 103 |
+
np.testing.assert_allclose(l2_norms(data["q_before"]), l2_norms(data["q_after"]), atol=1e-10)
|
| 104 |
+
slice_q = select_head(data["q_before"], 99)
|
| 105 |
+
self.assertEqual(slice_q.shape, (8, 16))
|
| 106 |
+
from src.plots import bulk_before_after_delta, attention_heatmaps
|
| 107 |
+
|
| 108 |
+
fig = bulk_before_after_delta(data["q_before"], data["q_after"])
|
| 109 |
+
self.assertTrue(len(fig.data) >= 1)
|
| 110 |
+
s = attention_scores(data["q_before"], data["k_before"])
|
| 111 |
+
attention_heatmaps(s, s)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
if __name__ == "__main__":
|
| 115 |
+
unittest.main()
|