Spaces:
Running on Zero
Running on Zero
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- app.py +80 -0
- 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"))
|