summaryrefslogtreecommitdiff
path: root/viz/gen_fig2.py
diff options
context:
space:
mode:
Diffstat (limited to 'viz/gen_fig2.py')
-rw-r--r--viz/gen_fig2.py41
1 files changed, 41 insertions, 0 deletions
diff --git a/viz/gen_fig2.py b/viz/gen_fig2.py
new file mode 100644
index 0000000..c65f8b5
--- /dev/null
+++ b/viz/gen_fig2.py
@@ -0,0 +1,41 @@
+import matplotlib.pyplot as plt
+import numpy as np
+
+fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5), dpi=300)
+
+layers = [f'L{i}' for i in range(6)]
+partial_r = [-0.307, -0.319, -0.338, -0.364, -0.240, +0.056]
+
+# Left: Partial correlation across layers
+ax1.plot(layers, partial_r, marker='o', linewidth=2.5, markersize=8, color='#ef4444', label='Partial r (frequency | ||W_U||)')
+ax1.axhline(0, color='gray', linestyle='--', alpha=0.7)
+ax1.set_ylabel('Partial Correlation', fontsize=11, fontweight='bold')
+ax1.set_xlabel('Transformer Layer', fontsize=11, fontweight='bold')
+ax1.set_title('Layer Dynamics Component\n(W_U Regressed Out)', fontsize=12, fontweight='bold', pad=12)
+ax1.set_ylim(-0.5, 0.2)
+ax1.grid(True, linestyle=':', alpha=0.6)
+ax1.legend(loc='lower left', frameon=True)
+
+for i, txt in enumerate(partial_r):
+ ax1.annotate(f'{txt:+.3f}', (layers[i], partial_r[i]), textcoords="offset points", xytext=(0,10 if txt>0 else -15), ha='center', fontweight='bold', fontsize=9)
+
+# Right: Static geometry W_U emergence (Init vs Trained)
+stages = ['Init (Seed 0)', 'Init (Seed 1)', 'Init (Seed 2)', 'Trained Model']
+wu_r = [+0.002, -0.189, -0.166, -0.606]
+colors = ['#94a3b8', '#94a3b8', '#94a3b8', '#6366f1']
+
+bars = ax2.bar(stages, wu_r, color=colors, width=0.55, alpha=0.9)
+ax2.axhline(0, color='gray', linestyle='--', alpha=0.7)
+ax2.set_ylabel('r( ||W_U[k]||, Token Frequency )', fontsize=11, fontweight='bold')
+ax2.set_title('Static Unembedding Geometry\n(Learned vs Initialized)', fontsize=12, fontweight='bold', pad=12)
+ax2.set_ylim(-0.8, 0.2)
+ax2.grid(True, linestyle=':', alpha=0.6)
+
+for bar in bars:
+ yval = bar.get_height()
+ ax2.text(bar.get_x() + bar.get_width()/2.0, yval - 0.05 if yval < 0 else yval + 0.02, f'{yval:+.3f}', ha='center', va='top' if yval < 0 else 'bottom', fontweight='bold', fontsize=9)
+
+plt.suptitle('Decomposition of Frequency Confound: Static Geometry vs. Layer Dynamics', fontsize=14, fontweight='bold', y=0.98)
+plt.tight_layout(rect=[0, 0, 1, 0.93])
+plt.savefig('docs/assets/fig2_wu_decomposition.png', dpi=300)
+print("Saved Fig 2")