""" 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'}")