summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--src/dim_starvation.py207
1 files changed, 207 insertions, 0 deletions
diff --git a/src/dim_starvation.py b/src/dim_starvation.py
new file mode 100644
index 0000000..edc00be
--- /dev/null
+++ b/src/dim_starvation.py
@@ -0,0 +1,207 @@
+"""
+Dimensional starvation test for J-space bottleneck.
+
+Tests whether Anthropic's "limited capacity" finding is actually
+just geometric compression when vocab_size >> d_model.
+
+Experiment:
+ A) nanoGPT baseline: vocab=65, d_model=384 (d_model >> V — no pressure)
+ B) nanoGPT starved: vocab=65, d_model=16 (V >> d_model — forced compression)
+ C) nanoGPT starved: vocab=65, d_model=32 (intermediate)
+
+If bottleneck (reduced J-space effective rank) only appears when
+d_model shrinks, then Anthropic's finding is geometric, not cognitive.
+
+Usage (inside Docker on meru):
+ python3 src/dim_starvation.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'
+DATA_DIR = 'data/shakespeare_char'
+BLOCK_SIZE = 128
+BATCH_SIZE = 32
+MAX_ITERS = 5000
+N_LAYERS = 6
+N_HEADS = {16: 4, 32: 4, 64: 4, 128: 4, 384: 6} # d_model -> n_head
+JLENS_BATCHES = 10
+JLENS_BS = 16
+
+def train_model(out_dir, d_model, data_dir=DATA_DIR):
+ """Train nanoGPT with specified d_model and return model + config."""
+ 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)
+
+ n_head = N_HEADS[d_model]
+ model_args = dict(n_layer=N_LAYERS, n_head=n_head, n_embd=d_model,
+ block_size=BLOCK_SIZE, bias=False,
+ vocab_size=meta['vocab_size'], dropout=0.1)
+
+ config = GPTConfig(**model_args)
+ model = GPT(config).to(device)
+ n_params = sum(p.numel() for p in model.parameters())
+ print(f" d_model={d_model}, n_head={n_head}, params={n_params/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):
+ 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)
+ _, 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')
+
+ data = train_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)
+ 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, model_args
+
+
+def analyze_jlens(model, data_dir, vocab_size):
+ """Run J-lens and compute effective rank per layer."""
+ 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)
+ itos = meta['itos']
+ d_model = model.config.n_embd
+
+ batch_size = JLENS_BS
+ block_size = BLOCK_SIZE
+ batches = []
+ for _ in range(JLENS_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))
+
+ layer_stats = {}
+ for layer_idx in range(N_LAYERS):
+ jlens = compute_jlens_layer(model, layer_idx, batches, device)
+
+ # Stack all token vectors
+ V = torch.stack([jlens[tid] for tid in range(vocab_size)])
+ U, S, Vt = torch.linalg.svd(V.float(), full_matrices=False)
+
+ eff_rank = (S > 0.01 * S[0]).sum().item()
+ pr = (S.sum()**2 / (S**2).sum()).item()
+
+ # Top tokens by norm
+ norms = {tid: jlens[tid].norm().item() for tid in range(vocab_size)}
+ sorted_toks = sorted(norms.items(), key=lambda x: x[1], reverse=True)
+
+ layer_stats[layer_idx] = {
+ 'eff_rank': eff_rank,
+ 'participation_ratio': pr,
+ 'top_tokens': [(itos[tid], norms[tid]) for tid, _ in sorted_toks[:5]],
+ }
+
+ return layer_stats
+
+
+# ── MAIN ───────────────────────────────────────────────
+print("=" * 60)
+print("DIMENSIONAL STARVATION TEST")
+print("=" * 60)
+print(f"Vocabulary size: 65")
+print()
+
+# We already have d_model=384 results from earlier
+# Test d_model values: 16, 32, 64, 128
+dims_to_test = [16, 32, 64, 128]
+
+results = {}
+# Baseline (already computed)
+results[384] = {'eff_rank': 65, 'pr': 47.1} # from previous run
+
+for d_model in dims_to_test:
+ print(f"\n{'='*60}")
+ print(f"Testing d_model={d_model} (V/d_model = {65/d_model:.1f}x)")
+ print(f"{'='*60}")
+
+ out_dir = f'out-starved-d{d_model}'
+
+ # Train (skip if checkpoint exists)
+ ckpt_path = f'{out_dir}/ckpt.pt'
+ if os.path.exists(ckpt_path):
+ print(f" Loading existing checkpoint...")
+ model, margs = jlens_v2.load_model(ckpt_path, device)
+ else:
+ print(f" Training...")
+ model, margs = train_model(out_dir, d_model)
+ model.eval()
+
+ # J-lens analysis
+ print(f" Running J-lens...")
+ stats = analyze_jlens(model, DATA_DIR, margs['vocab_size'])
+ results[d_model] = {layer: stats[layer] for layer in stats}
+
+ # Quick summary
+ for layer_idx in range(N_LAYERS):
+ s = stats[layer_idx]
+ print(f" L{layer_idx}: eff_rank={s['eff_rank']}, "
+ f"pr={s['participation_ratio']:.1f}, "
+ f"top={', '.join([t[0] for t in s['top_tokens'][:3]])}")
+
+# ── FINAL COMPARISON ────────────────────────────────────
+print(f"\n{'='*60}")
+print("FINAL COMPARISON: Effective Rank vs d_model")
+print(f"{'='*60}")
+print(f" {'d_model':<10} {'V/d_model':>10} {'L2 rank':>10} {'L3 rank':>10} {'L4 rank':>10} {'PR(L3)':>10}")
+print(f" {'-'*10} {'-'*10} {'-'*10} {'-'*10} {'-'*10} {'-'*10}")
+
+for d_model in sorted(results.keys()):
+ r = results[d_model]
+ ratio = 65 / d_model
+ r2 = r.get(2, {}).get('eff_rank', '?') if isinstance(r.get(2), dict) else '?'
+ r3 = r.get(3, {}).get('eff_rank', '?') if isinstance(r.get(3), dict) else '?'
+ r4 = r.get(4, {}).get('eff_rank', '?') if isinstance(r.get(4), dict) else '?'
+ pr3 = r.get(3, {}).get('participation_ratio', 0) if isinstance(r.get(3), dict) else 0
+ print(f" {d_model:<10} {ratio:>10.1f}x {str(r2):>10} {str(r3):>10} {str(r4):>10} {pr3:>10.1f}")
+
+print()
+print(" If Anthropic's bottleneck is geometric:")
+print(" - d_model=384 (V/d=0.2x): full rank (65/65)")
+print(" - d_model=16 (V/d=4.1x): reduced rank (<< 65)")
+print(" - d_model=32 (V/d=2.0x): intermediate rank")