DebasishDhal99 commited on
Commit
36a72e5
·
1 Parent(s): 4f0e23c

feat: add interactive plot to view rope rotation angle as a function of dimension pair index (d) and token position index (k)

Browse files
Files changed (2) hide show
  1. app.py +80 -0
  2. src/plots.py +122 -0
app.py CHANGED
@@ -30,6 +30,8 @@ from src.plots import (
30
  norm_compare_add_vs_rope,
31
  norms_and_cosine,
32
  position_sweep,
 
 
33
  rotation_2d,
34
  theta_heatmap,
35
  )
@@ -282,6 +284,57 @@ def update_bulk(data, which, head, mod_2pi):
282
  )
283
 
284
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
285
  def update_individual(data, which, head, token, pair, sweep):
286
  if not data:
287
  return "Compute on the Setup tab first.", PLACEHOLDER, PLACEHOLDER
@@ -449,6 +502,28 @@ with gr.Blocks(title="RoPE Explorer") as demo:
449
  bulk_theta = gr.Plot(label="θ(k, i)")
450
  bulk_freq = gr.Plot(label="ω_i")
451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
452
  with gr.Tab("Individual changes"):
453
  with gr.Row():
454
  token_k = gr.Slider(0, 15, step=1, value=0, label="Token index k")
@@ -552,6 +627,11 @@ score does not represent a universal threshold of importance.
552
  for ctrl in bulk_inputs:
553
  ctrl.change(update_bulk, inputs=bulk_inputs, outputs=bulk_outputs)
554
 
 
 
 
 
 
555
  ind_inputs = [state, which, head, token_k, pair_i, sweep]
556
  ind_outputs = [pair_table, pair_plot, sweep_plot]
557
  for ctrl in ind_inputs:
 
30
  norm_compare_add_vs_rope,
31
  norms_and_cosine,
32
  position_sweep,
33
+ rope_angle_heatmap,
34
+ rope_angle_slices,
35
  rotation_2d,
36
  theta_heatmap,
37
  )
 
284
  )
285
 
286
 
287
+ def update_angle_explorer(data, head, token, pair, display_mode):
288
+ if not data:
289
+ return (
290
+ PLACEHOLDER,
291
+ PLACEHOLDER,
292
+ "Compute on the Setup tab first.",
293
+ gr.update(maximum=1, value=0),
294
+ gr.update(maximum=1, value=0),
295
+ )
296
+ before = select_head(data["q_before"], int(head))
297
+ seq, dim = before.shape
298
+ n_pairs = dim // 2
299
+ token = int(np.clip(token, 0, seq - 1))
300
+ pair = int(np.clip(pair, 0, n_pairs - 1))
301
+ base = float(data["base"])
302
+ theta = float(theta_grid(seq, dim, base)[token, pair])
303
+ if display_mode == "turns":
304
+ shown_theta = theta / (2 * np.pi)
305
+ display_label = "θ / 2π (turns)"
306
+ elif display_mode == "wrapped":
307
+ shown_theta = float(np.mod(theta, 2 * np.pi))
308
+ display_label = "θ mod 2π (radians)"
309
+ else:
310
+ shown_theta = theta
311
+ display_label = "θ (radians)"
312
+ omega = float(pair_frequencies(dim, base)[pair])
313
+ d0, d1 = pair_dim_labels(pair, dim, style=data["style"])
314
+ detail = fr"""
315
+ ### Selected rotation angle
316
+
317
+ - **Head:** `{int(head)}`
318
+ - **Token position:** `k = {token}`
319
+ - **Pair:** `i = {pair}` → ({d0}, {d1})
320
+ - **Frequency:** `ωᵢ = {omega:.8f}` radians per token position
321
+ - **Raw angle:** `θ({token}, {pair}) = {theta:.8f}` radians
322
+ - **Displayed angle ({display_label}):** `{shown_theta:.8f}`
323
+
324
+ $$\theta(k,i) = k\,base^{{-2i/d}} = {token}\,({base:g})^{{-2\times{pair}/{dim}}}$$
325
+
326
+ Every increase of one token position adds `ωᵢ` radians for this pair. Lower-frequency
327
+ pairs change more slowly as `k` increases.
328
+ """
329
+ return (
330
+ rope_angle_heatmap(seq, dim, base, token, pair, display_mode),
331
+ rope_angle_slices(seq, dim, base, token, pair, display_mode),
332
+ detail,
333
+ gr.update(maximum=_safe_slider_max(seq - 1), value=token),
334
+ gr.update(maximum=_safe_slider_max(n_pairs - 1), value=pair),
335
+ )
336
+
337
+
338
  def update_individual(data, which, head, token, pair, sweep):
339
  if not data:
340
  return "Compute on the Setup tab first.", PLACEHOLDER, PLACEHOLDER
 
502
  bulk_theta = gr.Plot(label="θ(k, i)")
503
  bulk_freq = gr.Plot(label="ω_i")
504
 
505
+ with gr.Tab("RoPE angles"):
506
+ gr.Markdown(
507
+ "Explore how the rotation angle changes with token position `k` and "
508
+ "dimension pair `i`. The selected head is shared with the Setup tab. "
509
+ "A marker identifies the selected `(k, i)` cell in the heatmap."
510
+ )
511
+ with gr.Row():
512
+ angle_token = gr.Slider(0, 15, step=1, value=0, label="Token position k")
513
+ angle_pair = gr.Slider(0, 15, step=1, value=0, label="Pair index i")
514
+ angle_display = gr.Radio(
515
+ [
516
+ ("Absolute angle (radians)", "absolute"),
517
+ ("Angle / 2π (turns)", "turns"),
518
+ ("Wrapped angle mod 2π", "wrapped"),
519
+ ],
520
+ value="absolute",
521
+ label="Angle display",
522
+ )
523
+ angle_heatmap = gr.Plot(label="RoPE angle heatmap")
524
+ angle_slices = gr.Plot(label="Selected pair/token slices")
525
+ angle_detail = gr.Markdown("Compute on the Setup tab first.")
526
+
527
  with gr.Tab("Individual changes"):
528
  with gr.Row():
529
  token_k = gr.Slider(0, 15, step=1, value=0, label="Token index k")
 
627
  for ctrl in bulk_inputs:
628
  ctrl.change(update_bulk, inputs=bulk_inputs, outputs=bulk_outputs)
629
 
630
+ angle_inputs = [state, head, angle_token, angle_pair, angle_display]
631
+ angle_outputs = [angle_heatmap, angle_slices, angle_detail, angle_token, angle_pair]
632
+ for ctrl in angle_inputs:
633
+ ctrl.change(update_angle_explorer, inputs=angle_inputs, outputs=angle_outputs)
634
+
635
  ind_inputs = [state, which, head, token_k, pair_i, sweep]
636
  ind_outputs = [pair_table, pair_plot, sweep_plot]
637
  for ctrl in ind_inputs:
src/plots.py CHANGED
@@ -115,6 +115,128 @@ def theta_heatmap(seq_len: int, dim: int, base: float, mod_2pi: bool = False) ->
115
  return fig
116
 
117
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
118
  def frequency_strip(dim: int, base: float) -> go.Figure:
119
  omega = pair_frequencies(dim, base=base)
120
  fig = go.Figure(data=go.Bar(x=list(range(len(omega))), y=omega, name="ω_i"))
 
115
  return fig
116
 
117
 
118
+ def rope_angle_heatmap(
119
+ seq_len: int,
120
+ dim: int,
121
+ base: float,
122
+ selected_token: int = 0,
123
+ selected_pair: int = 0,
124
+ display_mode: str = "absolute",
125
+ ) -> go.Figure:
126
+ """Interactive θ(k, i) map with token/pair selection marker."""
127
+ grid = theta_grid(seq_len, dim, base=base)
128
+ if display_mode == "turns":
129
+ displayed = grid / (2 * np.pi)
130
+ value_label = "θ / 2π (turns)"
131
+ shown_label = "θ / 2π"
132
+ elif display_mode == "wrapped":
133
+ displayed = np.mod(grid, 2 * np.pi)
134
+ value_label = "θ mod 2π (radians)"
135
+ shown_label = "θ mod 2π"
136
+ else:
137
+ displayed = grid
138
+ value_label = "θ (radians)"
139
+ shown_label = "θ"
140
+ pairs = grid.shape[1]
141
+ tokens = np.arange(seq_len)
142
+ pair_indices = np.arange(pairs)
143
+ selected_token = int(np.clip(selected_token, 0, seq_len - 1))
144
+ selected_pair = int(np.clip(selected_pair, 0, pairs - 1))
145
+ fig = go.Figure(
146
+ go.Heatmap(
147
+ z=displayed.T,
148
+ x=tokens,
149
+ y=pair_indices,
150
+ customdata=grid.T,
151
+ colorbar=dict(title=value_label),
152
+ hovertemplate=(
153
+ "token position k=%{x}<br>pair index i=%{y}<br>"
154
+ "raw θ=%{customdata:.6f} rad<br>"
155
+ + f"shown {shown_label}=%{{z:.6f}}"
156
+ + "<extra></extra>"
157
+ ),
158
+ )
159
+ )
160
+ fig.add_trace(
161
+ go.Scatter(
162
+ x=[selected_token],
163
+ y=[selected_pair],
164
+ mode="markers",
165
+ name="selected (k, i)",
166
+ marker=dict(size=11, color="white", line=dict(color="black", width=2)),
167
+ hovertemplate="selected k=%{x}, i=%{y}<extra></extra>",
168
+ )
169
+ )
170
+ fig.update_layout(
171
+ title="RoPE angle θ(k, i) across token positions and dimension pairs",
172
+ xaxis_title="token position k",
173
+ yaxis_title="pair index i",
174
+ template="plotly_white",
175
+ height=480,
176
+ margin=dict(t=65, b=55),
177
+ )
178
+ return fig
179
+
180
+
181
+ def rope_angle_slices(
182
+ seq_len: int,
183
+ dim: int,
184
+ base: float,
185
+ selected_token: int = 0,
186
+ selected_pair: int = 0,
187
+ display_mode: str = "absolute",
188
+ ) -> go.Figure:
189
+ """Show θ for one pair over positions and one position over pairs."""
190
+ grid = theta_grid(seq_len, dim, base=base)
191
+ if display_mode == "turns":
192
+ displayed = grid / (2 * np.pi)
193
+ value_label = "θ / 2π (turns)"
194
+ elif display_mode == "wrapped":
195
+ displayed = np.mod(grid, 2 * np.pi)
196
+ value_label = "θ mod 2π (radians)"
197
+ else:
198
+ displayed = grid
199
+ value_label = "θ (radians)"
200
+ selected_token = int(np.clip(selected_token, 0, seq_len - 1))
201
+ selected_pair = int(np.clip(selected_pair, 0, grid.shape[1] - 1))
202
+ fig = make_subplots(
203
+ rows=1,
204
+ cols=2,
205
+ subplot_titles=[
206
+ f"Pair i={selected_pair}: angle over token position k",
207
+ f"Token k={selected_token}: angle over pair index i",
208
+ ],
209
+ )
210
+ fig.add_trace(
211
+ go.Scatter(
212
+ x=np.arange(seq_len),
213
+ y=displayed[:, selected_pair],
214
+ mode="lines",
215
+ name=f"pair {selected_pair}",
216
+ hovertemplate="k=%{x}<br>θ=%{y:.6f} rad<extra></extra>",
217
+ ),
218
+ row=1,
219
+ col=1,
220
+ )
221
+ fig.add_trace(
222
+ go.Scatter(
223
+ x=np.arange(grid.shape[1]),
224
+ y=displayed[selected_token, :],
225
+ mode="lines+markers",
226
+ name=f"token {selected_token}",
227
+ hovertemplate="i=%{x}<br>θ=%{y:.6f} rad<extra></extra>",
228
+ ),
229
+ row=1,
230
+ col=2,
231
+ )
232
+ fig.update_xaxes(title_text="token position k", row=1, col=1)
233
+ fig.update_yaxes(title_text=value_label, row=1, col=1)
234
+ fig.update_xaxes(title_text="pair index i", row=1, col=2)
235
+ fig.update_yaxes(title_text=value_label, row=1, col=2)
236
+ fig.update_layout(template="plotly_white", height=360, showlegend=False)
237
+ return fig
238
+
239
+
240
  def frequency_strip(dim: int, base: float) -> go.Figure:
241
  omega = pair_frequencies(dim, base=base)
242
  fig = go.Figure(data=go.Bar(x=list(range(len(omega))), y=omega, name="ω_i"))