summaryrefslogtreecommitdiff
path: root/src/freq_ablation.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/freq_ablation.py')
-rw-r--r--src/freq_ablation.py139
1 files changed, 139 insertions, 0 deletions
diff --git a/src/freq_ablation.py b/src/freq_ablation.py
new file mode 100644
index 0000000..c011140
--- /dev/null
+++ b/src/freq_ablation.py
@@ -0,0 +1,139 @@
+"""
+Controlled frequency ablation: train model on doubled-q Shakespeare,
+compare J-lens norms for 'q' vs control.
+
+Usage (inside Docker on meru):
+ python3 src/freq_ablation.py
+"""
+import sys, os
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+import torch, numpy as np, pickle
+
+from model import GPT, GPTConfig
+import jlens_v2
+from jlens_v2 import compute_jlens_layer
+
+device = 'cuda'
+MODEL_ARGS = dict(n_layer=6, n_head=6, n_embd=384, block_size=128,
+ bias=False, vocab_size=65, dropout=0.2)
+
+def get_batch(train_data, val_data, split, block_size=128, batch_size=32):
+ data = train_data if split == 'train' else val_data
+ ix = torch.randint(len(data) - block_size, (batch_size,))
+ x = torch.stack([torch.from_numpy(data[i:i+block_size].astype(np.int64)) for i in ix])
+ y = torch.stack([torch.from_numpy(data[i+1:i+1+block_size].astype(np.int64)) for i in ix])
+ return x.to(device), y.to(device)
+
+def train(data_dir, out_dir, max_iters=5000):
+ """Train a model and return best checkpoint path."""
+ train_data = np.memmap(f'{data_dir}/train.bin', dtype=np.uint16, mode='r')
+ val_data = np.memmap(f'{data_dir}/val.bin', dtype=np.uint16, mode='r')
+
+ config = GPTConfig(**MODEL_ARGS)
+ model = GPT(config).to(device)
+ print(f" Params: {sum(p.numel() for p in model.parameters())/1e6:.2f}M")
+
+ optimizer = model.configure_optimizers(weight_decay=0.1, lr=1e-3,
+ betas=(0.9, 0.99), device_type='cuda')
+ os.makedirs(out_dir, exist_ok=True)
+ best_val = 1e9
+
+ for it in range(max_iters):
+ if it % 500 == 0:
+ model.eval()
+ losses = {}
+ for split in ['train', 'val']:
+ lv = []
+ for _ in range(50):
+ X, Y = get_batch(train_data, val_data, split)
+ _, loss = model(X, Y)
+ lv.append(loss.item())
+ losses[split] = np.mean(lv)
+ model.train()
+ print(f" step {it}: train={losses['train']:.4f}, val={losses['val']:.4f}")
+ if losses['val'] < best_val:
+ best_val = losses['val']
+ torch.save({'model': model.state_dict(), 'model_args': MODEL_ARGS,
+ 'best_val_loss': best_val}, f'{out_dir}/ckpt.pt')
+
+ X, Y = get_batch(train_data, val_data, 'train')
+ logits, loss = model(X, Y)
+ loss.backward()
+ optimizer.step()
+ optimizer.zero_grad(set_to_none=True)
+ if it % 500 == 0:
+ print(f" iter {it}: loss={loss.item():.4f}")
+
+ print(f" Done. Best val: {best_val:.4f}")
+ return model
+
+def run_jlens_on_model(model, data_dir, target_char, n_batches=10):
+ """Compute J-lens and return norm for target_char."""
+ train_data = np.memmap(f'{data_dir}/train.bin', dtype=np.uint16, mode='r')
+ with open(f'{data_dir}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ target_id = meta['stoi'][target_char]
+
+ batch_size = 16
+ block_size = MODEL_ARGS['block_size']
+ batches = []
+ for _ in range(n_batches):
+ ix = torch.randint(len(train_data) - block_size, (batch_size,))
+ x = torch.stack([torch.from_numpy(
+ train_data[i:i+block_size].astype(np.int64)) for i in ix])
+ y = torch.stack([torch.from_numpy(
+ train_data[i+1:i+1+block_size].astype(np.int64)) for i in ix])
+ batches.append((x, y))
+
+ results = {}
+ for layer_idx in [2, 3, 4]:
+ jlens = compute_jlens_layer(model, layer_idx, batches, device)
+ results[layer_idx] = jlens[target_id].norm().item()
+
+ return results
+
+# ── MAIN ───────────────────────────────────────────────
+print("=" * 60)
+print("CONTROLLED FREQUENCY ABLATION: 'q' DOUBLED")
+print("=" * 60)
+
+# 1. Train control model (or use existing)
+print("\n[1/4] Using existing control model...")
+ctrl_model, _ = jlens_v2.load_model('out-shakespeare-char/ckpt.pt', device)
+ctrl_model.eval()
+
+# 2. Train modified model
+print("\n[2/4] Training modified model (doubled 'q')...")
+mod_model = train('data/freq_experiment/doubled_q', 'out-freq-doubled-q')
+mod_model.eval()
+
+# 3. J-lens on both
+print("\n[3/4] Computing J-lens...")
+print(" Control model...")
+ctrl_norms = run_jlens_on_model(ctrl_model, 'data/shakespeare_char', 'q')
+print(" Modified model...")
+mod_norms = run_jlens_on_model(mod_model, 'data/freq_experiment/doubled_q', 'q')
+
+# 4. Compare
+print("\n[4/4] RESULTS")
+print("=" * 60)
+print(f" Character: 'q'")
+print(f" Original frequency: 0.055%")
+print(f" Doubled frequency: 0.109%")
+print()
+print(f" {'Layer':<8} {'Control':>10} {'Modified':>10} {'Change':>10}")
+print(f" {'-'*8} {'-'*10} {'-'*10} {'-'*10}")
+for layer_idx in sorted(ctrl_norms.keys()):
+ c = ctrl_norms[layer_idx]
+ m = mod_norms[layer_idx]
+ pct = (m - c) / c * 100
+ direction = "↓" if pct < 0 else "↑"
+ print(f" L{layer_idx:<7} {c:>10.4f} {m:>10.4f} {direction}{abs(pct):>8.1f}%")
+
+print()
+avg_ctrl = np.mean(list(ctrl_norms.values()))
+avg_mod = np.mean(list(mod_norms.values()))
+avg_pct = (avg_mod - avg_ctrl) / avg_ctrl * 100
+print(f" Average: {avg_ctrl:.4f} → {avg_mod:.4f} ({avg_pct:+.1f}%)")
+print(f" HYPOTHESIS: doubling frequency should REDUCE J-lens norm")
+print(f" RESULT: {'CONFIRMED' if avg_pct < -2 else 'INCONCLUSIVE' if abs(avg_pct) < 2 else 'CONTRADICTED'}")