DebasishDhal99 commited on
Commit
aa7bfed
·
1 Parent(s): 1d852fb

feat: add rope visualization only with some minor cmparison with sinusoidial positional embedding

Browse files
Files changed (8) hide show
  1. README.md +50 -13
  2. app.py +353 -20
  3. requirements.txt +7 -5
  4. src/absolute_pe.py +35 -0
  5. src/extract.py +260 -0
  6. src/plots.py +305 -0
  7. src/rope.py +124 -0
  8. tests/test_core.py +115 -0
README.md CHANGED
@@ -1,13 +1,50 @@
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: CVisualize how ROPE changes the embeddings
12
- ---
13
- # rope-implementation
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- import gradio as gr
2
- import spaces
3
-
4
- @spaces.GPU
5
- def gpu_test():
6
- return "GPU available"
7
-
8
-
9
- with gr.Blocks(title="RoPE Explorer") as demo:
10
- with gr.Tabs():
11
- with gr.Tab("Block 01"):
12
- gr.Markdown("## Block 01")
13
- gr.Markdown("Placeholder for the first visualization.")
14
-
15
- with gr.Tab("Block 02"):
16
- gr.Markdown("## Block 02")
17
- gr.Markdown("Placeholder for the second visualization.")
18
-
19
- if __name__ == "__main__":
20
- demo.launch( server_name="0.0.0.0", server_port=7860, )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- numpy>=2.0,<2.3
4
- spaces==0.51.3
5
- plotly==7.0.0
 
 
 
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()