1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
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, learning_rate=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")
|