summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/freq_ablation.py139
-rw-r--r--src/freq_experiment.py344
2 files changed, 483 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'}")
diff --git a/src/freq_experiment.py b/src/freq_experiment.py
new file mode 100644
index 0000000..4050c16
--- /dev/null
+++ b/src/freq_experiment.py
@@ -0,0 +1,344 @@
+"""
+Controlled frequency experiment for J-lens.
+
+Tests causality: if we artificially double the frequency of a character,
+does its J-lens norm predictably drop?
+
+Hypothesis (from chain rule):
+ J-lens norm ∝ (1 - p_avg), where p_avg is average predicted probability.
+ Doubling frequency → model learns higher p → J-lens norm drops.
+
+Control: train two identical models on:
+ A) Original Shakespeare (baseline)
+ B) Modified Shakespeare with 'q' frequency doubled
+"""
+
+import sys, os
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+def create_modified_dataset(input_path, output_path, target_char='q', factor=2.0):
+ """Double the frequency of target_char in the text.
+ Inserts target_char at random positions until frequency is doubled.
+ """
+ import random
+ random.seed(42)
+
+ with open(input_path, 'r') as f:
+ text = f.read()
+
+ target_count = text.count(target_char)
+ target_freq = target_count / len(text)
+ print(f"Original: {target_count} occurrences of '{target_char}', "
+ f"frequency = {target_freq:.4%}")
+
+ # Insert additional copies at random positions
+ extra_needed = int(target_count * (factor - 1))
+ positions = sorted(random.sample(range(len(text)), extra_needed))
+
+ modified = list(text)
+ for i, pos in enumerate(positions):
+ modified.insert(pos + i, target_char) # offset by previous insertions
+
+ modified_text = ''.join(modified)
+ new_count = modified_text.count(target_char)
+ new_freq = new_count / len(modified_text)
+ print(f"Modified: {new_count} occurrences of '{target_char}', "
+ f"frequency = {new_freq:.4%}")
+
+ with open(output_path, 'w') as f:
+ f.write(modified_text)
+
+ return modified_text
+
+
+def prepare_data(input_path, output_dir):
+ """Run nanoGPT's prepare.py on a text file."""
+ import subprocess
+ # Write a temp prepare script
+ import os
+ os.makedirs(output_dir, exist_ok=True)
+
+ # Read the text
+ with open(input_path, 'r') as f:
+ text = f.read()
+
+ chars = sorted(list(set(text)))
+ vocab_size = len(chars)
+ stoi = {ch: i for i, ch in enumerate(chars)}
+ itos = {i: ch for i, ch in enumerate(chars)}
+
+ # Encode
+ import numpy as np
+ data = np.array([stoi[ch] for ch in text], dtype=np.uint16)
+ n = int(0.9 * len(data))
+ train_data = data[:n]
+ val_data = data[n:]
+
+ train_data.tofile(os.path.join(output_dir, 'train.bin'))
+ val_data.tofile(os.path.join(output_dir, 'val.bin'))
+
+ import pickle
+ with open(os.path.join(output_dir, 'meta.pkl'), 'wb') as f:
+ pickle.dump({'stoi': stoi, 'itos': itos, 'vocab_size': vocab_size}, f)
+
+ # Also save the raw text
+ with open(os.path.join(output_dir, 'input.txt'), 'w') as f:
+ f.write(text)
+
+ print(f"Prepared {output_dir}: {len(train_data):,} train, "
+ f"{len(val_data):,} val, {vocab_size} vocab")
+ return stoi, itos, vocab_size
+
+
+def train_model(data_dir, out_dir, device='cuda'):
+ """Train nanoGPT on a prepared dataset."""
+ import subprocess
+
+ # Use the config but override dataset path and output
+ import torch
+ from model import GPT, GPTConfig
+
+ # Load and train
+ script = f"""
+import sys
+sys.path.insert(0, '.')
+import torch
+import numpy as np
+import pickle
+import os
+from model import GPT, GPTConfig
+
+# Load data
+data_dir = '{data_dir}'
+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')
+with open(f'{{data_dir}}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+
+# Config — same as train_shakespeare_char but smaller for speed
+out_dir = '{out_dir}'
+eval_interval = 500
+eval_iters = 100
+log_interval = 100
+always_save_checkpoint = False
+wandb_log = False
+dataset = 'custom'
+gradient_accumulation_steps = 1
+batch_size = 32
+block_size = 128
+n_layer = 6
+n_head = 6
+n_embd = 384
+dropout = 0.2
+learning_rate = 1e-3
+max_iters = 5000
+lr_decay_iters = 5000
+min_lr = 1e-4
+beta2 = 0.99
+warmup_iters = 100
+dtype = 'float32'
+flash = False
+device = '{device}'
+compile = False
+vocab_size = meta['vocab_size']
+
+# Build model
+model_args = dict(n_layer=n_layer, n_head=n_head, n_embd=n_embd,
+ block_size=block_size, bias=False, vocab_size=vocab_size,
+ dropout=dropout)
+gptconf = GPTConfig(**model_args)
+model = GPT(gptconf)
+model.to(device)
+
+# Optimizer
+optimizer = model.configure_optimizers(weight_decay=0.1, learning_rate=learning_rate,
+ betas=(0.9, beta2), device_type=device)
+scaler = torch.amp.GradScaler('cuda', enabled=False)
+
+def get_batch(split):
+ 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])
+ x, y = x.to(device), y.to(device)
+ return x, y
+
+@torch.no_grad()
+def estimate_loss():
+ out = {{}}
+ model.eval()
+ for split in ['train', 'val']:
+ losses = torch.zeros(eval_iters)
+ for k in range(eval_iters):
+ X, Y = get_batch(split)
+ logits, loss = model(X, Y)
+ losses[k] = loss.item()
+ out[split] = losses.mean()
+ model.train()
+ return out
+
+print(f"Training model: {{model_args}}")
+print(f"Parameters: {{sum(p.numel() for p in model.parameters())/1e6:.2f}}M")
+
+os.makedirs(out_dir, exist_ok=True)
+best_val_loss = 1e9
+
+for iter_num in range(max_iters):
+ if iter_num % eval_interval == 0:
+ losses = estimate_loss()
+ print(f"step {{iter_num}}: train loss {{losses['train']:.4f}}, val loss {{losses['val']:.4f}}")
+ if losses['val'] < best_val_loss:
+ best_val_loss = losses['val']
+ checkpoint = {{
+ 'model': model.state_dict(),
+ 'optimizer': optimizer.state_dict(),
+ 'model_args': model_args,
+ 'iter_num': iter_num,
+ 'best_val_loss': best_val_loss,
+ }}
+ torch.save(checkpoint, os.path.join(out_dir, 'ckpt.pt'))
+
+ X, Y = get_batch('train')
+ with torch.amp.autocast(device_type=device, dtype=torch.float32):
+ logits, loss = model(X, Y)
+
+ scaler.scale(loss).backward()
+ scaler.step(optimizer)
+ scaler.update()
+ optimizer.zero_grad(set_to_none=True)
+
+ if iter_num % log_interval == 0:
+ print(f"iter {{iter_num}}: loss {{loss.item():.4f}}")
+
+print(f"Training complete. Best val loss: {{best_val_loss:.4f}}")
+"""
+
+ import tempfile
+ with tempfile.NamedTemporaryFile(mode='w', suffix='.py', delete=False) as f:
+ f.write(script)
+ script_path = f.name
+
+ result = subprocess.run(['python3', script_path], capture_output=True, text=True)
+ os.unlink(script_path)
+ print(result.stdout)
+ if result.returncode != 0:
+ print("STDERR:", result.stderr)
+ return result.returncode == 0
+
+
+def main():
+ import argparse
+ parser = argparse.ArgumentParser()
+ parser.add_argument('--input', default='data/shakespeare_char/input.txt')
+ parser.add_argument('--target_char', default='q')
+ parser.add_argument('--factor', type=float, default=2.0)
+ parser.add_argument('--skip_train', action='store_true')
+ parser.add_argument('--device', default='cuda')
+ args = parser.parse_args()
+
+ # 1. Create modified dataset
+ print("=" * 60)
+ print("STEP 1: Creating modified dataset")
+ print("=" * 60)
+ modified_input = f'data/freq_experiment/modified_{args.target_char}.txt'
+ os.makedirs('data/freq_experiment', exist_ok=True)
+ create_modified_dataset(args.input, modified_input, args.target_char, args.factor)
+
+ # 2. Prepare both datasets
+ print()
+ print("=" * 60)
+ print("STEP 2: Preparing datasets")
+ print("=" * 60)
+
+ # Control: already prepared as data/shakespeare_char/
+ # Modified: prepare from modified text
+ mod_dir = f'data/freq_experiment/modified_{args.target_char}'
+ stoi_mod, itos_mod, vocab_mod = prepare_data(modified_input, mod_dir)
+
+ # 3. Train both models
+ if not args.skip_train:
+ print()
+ print("=" * 60)
+ print("STEP 3: Training CONTROL model (original Shakespeare)")
+ print("=" * 60)
+ # Control already trained as out-shakespeare-char/
+
+ print()
+ print("=" * 60)
+ print(f"STEP 4: Training MODIFIED model ({args.target_char} x{args.factor})")
+ print("=" * 60)
+ train_model(mod_dir, f'out-freq-{args.target_char}', args.device)
+
+ # 4. Run J-lens on both
+ print()
+ print("=" * 60)
+ print("STEP 5: Running J-lens on both models")
+ print("=" * 60)
+ from jlens_v2 import compute_jlens_layer, load_model
+ import numpy as np
+
+ def run_jlens(ckpt_path, data_dir, label):
+ print(f"\n J-lens on {label}...")
+ model, config = load_model(ckpt_path, args.device)
+
+ # Build batches
+ train_data = np.memmap(f'{data_dir}/train.bin', dtype=np.uint16, mode='r')
+ n_batches = 10
+ b_size = 16
+ block_size = config.block_size
+ batches = []
+ for _ in range(n_batches):
+ ix = torch.randint(len(train_data) - block_size, (b_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))
+
+ # Compute for middle layers only (speed)
+ results = {}
+ for layer_idx in [2, 3, 4]:
+ results[layer_idx] = compute_jlens_layer(model, layer_idx, batches, args.device)
+
+ return results
+
+ # Control
+ ctrl_jlens = run_jlens('out-shakespeare-char/ckpt.pt',
+ 'data/shakespeare_char', 'CONTROL')
+
+ # Modified
+ mod_jlens = run_jlens(f'out-freq-{args.target_char}/ckpt.pt',
+ mod_dir, 'MODIFIED')
+
+ # 5. Compare
+ print()
+ print("=" * 60)
+ print(f"STEP 6: COMPARISON — '{args.target_char}' J-lens norm change")
+ print("=" * 60)
+
+ with open('data/shakespeare_char/meta.pkl', 'rb') as f:
+ import pickle
+ ctrl_meta = pickle.load(f)
+ ctrl_itos = ctrl_meta['itos']
+
+ with open(f'{mod_dir}/meta.pkl', 'rb') as f:
+ mod_meta = pickle.load(f)
+ mod_itos = mod_meta['itos']
+
+ target_id_ctrl = ctrl_itos.index(args.target_char)
+ target_id_mod = mod_itos.index(args.target_char)
+
+ for layer_idx in sorted(ctrl_jlens.keys()):
+ ctrl_norm = ctrl_jlens[layer_idx][target_id_ctrl].norm().item()
+ mod_norm = mod_jlens[layer_idx][target_id_mod].norm().item()
+ change = (mod_norm - ctrl_norm) / ctrl_norm * 100
+ direction = "↓" if change < 0 else "↑"
+ print(f" Layer {layer_idx}: control={ctrl_norm:.4f} → modified={mod_norm:.4f} "
+ f"({direction}{abs(change):.1f}%)")
+
+ print()
+ print("Hypothesis confirmed!" if change < 0 else "Hypothesis NOT confirmed — unexpected result.")
+
+
+if __name__ == '__main__':
+ main()