diff options
Diffstat (limited to 'src/freq_ablation.py')
| -rw-r--r-- | src/freq_ablation.py | 139 |
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'}") |
