Kernel-Prompted Random Transformer / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import json, sys, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
9
10SEED = 2052
11N_SUPPORT = 64
12DH = 64
13EPOCHS = 20
14# Union of step sizes is shared by baseline and idea configs.
15LR_GRID = [1e-3, 3e-3, 1e-2]
16SIGMA_GRID = [0.4, 0.6, 0.8]
17records = {'baseline': {}, 'idea': {}}
18
19class FrozenAttentionPrompt(nn.Module):
20 """One frozen single-head attention layer, with trainable prompt only."""
21 def __init__(self, d_in=32, dh=DH, n_prompt=N_SUPPORT, seed=0):
22 super().__init__()
23 r = np.random.RandomState(seed)
24 q = r.normal(0, 1/np.sqrt(dh), (dh, dh)).astype('float32')
25 k = r.normal(0, 1/np.sqrt(dh), (dh, dh)).astype('float32')
26 v = r.normal(0, 1/np.sqrt(dh), (dh, 1)).astype('float32')
27 self.register_buffer('Q', torch.from_numpy(q))
28 self.register_buffer('K', torch.from_numpy(k))
29 self.register_buffer('V', torch.from_numpy(v))
30 self.prompt = nn.Parameter(torch.zeros(n_prompt, dh))
31 nn.init.normal_(self.prompt, std=0.05)
32 self.d_in, self.dh = d_in, dh
33
34 def embed(self, x):
35 z = torch.zeros(x.shape[0], self.dh, device=x.device, dtype=x.dtype)
36 z[:, :self.d_in] = x
37 z[:, self.d_in] = 1.
38 return z
39
40 def forward(self, x):
41 e = self.embed(x)
42 q = e @ self.Q
43 keys = self.prompt @ self.K
44 logits = q @ keys.T / np.sqrt(self.dh)
45 w = torch.softmax(logits, dim=1)
46 vals = self.prompt @ self.V
47 return w @ vals
48
49def set_seed(seed):
50 np.random.seed(seed); torch.manual_seed(seed)
51 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
52
53def dataset(seed):
54 d = get_dataset('sequence', seed, n_train=400, n_test=400)
55 # Bench train_model expects torch tensors; get_dataset already supplies them.
56 return d
57
58def baseline_fn(cfg):
59 def run(seed):
60 set_seed(seed)
61 ds = dataset(seed)
62 model = FrozenAttentionPrompt(seed=seed)
63 model, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
64 batch=128, weight_decay=0.0, log=lambda *_: None)
65 if model is None: return float('nan')
66 with torch.no_grad():
67 device = next(model.parameters()).device
68 x = ds['xte'].to(device); pred = model(x)
69 # observed trained-model attention entropy and prompt norm
70 e = model.embed(x[:32]); logits = (e @ model.Q) @ (model.prompt @ model.K).T / np.sqrt(DH)
71 records['baseline'][int(seed)] = {'prompt_norm': float(model.prompt.norm()),
72 'attention_entropy': float((-(torch.softmax(logits,1)*torch.log_softmax(logits,1)).sum(1)).mean())}
73 return float(metric)
74 return run
75
76def analytic_prompt(xs, ys, sigma, seed):
77 r = np.random.RandomState(seed)
78 q = r.normal(0, 1/np.sqrt(DH), (DH, DH))
79 k = r.normal(0, 1/np.sqrt(DH), (DH, DH))
80 v = r.normal(0, 1/np.sqrt(DH), (DH, 1))
81 A = (q.T @ k)[:33] / np.sqrt(DH)
82 B = np.vstack([A, v.T])
83 c = np.column_stack([xs / sigma**2, -np.sum(xs*xs, axis=1)/(2*sigma**2), ys])
84 P = np.linalg.lstsq(B, c.T, rcond=None)[0].T
85 return P, q, k, v
86
87def idea_fn(cfg):
88 def run(seed):
89 ds = dataset(seed)
90 xs = ds['xtr'][:N_SUPPORT].numpy().astype('float64')
91 ys = ds['ytr'][:N_SUPPORT].numpy().astype('float64').reshape(-1,1)
92 P,q,k,v = analytic_prompt(xs, ys, cfg['sigma'], seed)
93 xt = ds['xte'].numpy().astype('float64')
94 e = np.zeros((len(xt), DH)); e[:,:32] = xt; e[:,32] = 1
95 logits = e @ ((q.T @ k) @ P.T) / np.sqrt(DH)
96 logits -= logits.max(1, keepdims=True)
97 w = np.exp(logits); w /= w.sum(1, keepdims=True)
98 pred = w @ (P @ v)
99 metric = float(np.mean((pred[:,0] - ds['yte'].numpy())**2))
100 target = xt @ xs.T / cfg['sigma']**2 - (xs*xs).sum(1)[None,:]/(2*cfg['sigma']**2)
101 # common query term is intentionally omitted, as it cancels in softmax.
102 residual = float(np.sqrt(np.mean((e @ ((q.T@k)@P.T)/np.sqrt(DH)-target)**2)))
103 records['idea'].setdefault(cfg['sigma'], {})[int(seed)] = {
104 'prompt_norm': float(np.linalg.norm(P,axis=1).mean()),
105 'logit_residual': residual,
106 'entropy': float(np.mean(-(w*np.log(np.maximum(w,1e-30))).sum(1))) }
107 return metric
108 return run
109
110def main():
111 # Baseline sweep covers all rates used by idea-side comparison.
112 base_grid = [{'lr': x} for x in LR_GRID]
113 base = sweep_baseline(baseline_fn, base_grid)
114 # Idea sweep has three nearby bandwidth settings; each is evaluated on all 8 seeds.
115 idea_runs = []
116 for sigma, lr in zip(SIGMA_GRID, LR_GRID):
117 r = evaluate(idea_fn({'sigma': sigma, 'lr': lr}))
118 idea_runs.append({'cfg': {'sigma': sigma, 'lr': lr}, 'result': r})
119 best = min(idea_runs, key=lambda z: z['result']['mean'])
120 extra = {'track_choice': 'sequence: multi-token temporal-window correlations require sequence attention',
121 'best_sigma': best['cfg']['sigma'], 'idea_sweep': idea_runs,
122 'trained_baseline_behavior': records['baseline']}
123 sigs = [records['idea'][s] for s in SIGMA_GRID if s in records['idea']]
124 if sigs:
125 means = [float(np.mean([x['prompt_norm'] for x in z.values()])) for z in sigs]
126 ref = means[1]
127 best_resid = float(np.mean([x['logit_residual'] for x in records['idea'][best['cfg']['sigma']].values()]))
128 extra['mechanism_prediction'] = {'prediction': 'prompt norm scales as sigma^-2; affine logit residual is near zero',
129 'observed_prompt_norms': dict(zip(SIGMA_GRID, means)),
130 'predicted_norm_ratio_sigma_0.4_to_0.8': 4.0,
131 'observed_norm_ratio_sigma_0.4_to_0.8': means[0]/means[-1],
132 'observed_best_logit_residual': best_resid,
133 'confirmed': bool(abs(means[0]/means[-1]-4.0) < 0.5 and best_resid < 1e-4)}
134 report = make_report('sequence', 'frozen_random_attention_prompt', base, best['result'], extra)
135 report['baseline_sweep_shared_lr_grid'] = LR_GRID
136 report['cut'] = 'Only sequence track tested; no CIFAR/MNIST transfer and no latency/FLOP study.'
137 Path('bench_report.json').write_text(json.dumps(report, indent=2))
138 print(json.dumps(report, indent=2))
139
140if __name__ == '__main__': main()