diff options
| author | Void Agent <void@jayrup.hermes> | 2026-07-30 15:33:57 +0100 |
|---|---|---|
| committer | Void Agent <void@jayrup.hermes> | 2026-07-30 15:33:57 +0100 |
| commit | 6216c8b69d4d3ed37201f8ed6fa97a6eee854e28 (patch) | |
| tree | 19c7eb22f5c4ca4be1e39a31704d547de4e2b785 /src | |
| parent | b25662f7d37a13377da166c93d227d0c810df7d8 (diff) | |
Add dimensional starvation experiment (d_model=16,32,64,128 vs 384)
Diffstat (limited to 'src')
| -rw-r--r-- | src/dim_starvation.py | 207 |
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") |
