Manifold-kernel attention / bench_experiment.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, math, 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, make_model, train_model, sweep_baseline, make_report
9
10SEEDS = tuple(range(8))
11SWEEP_SEEDS = (0, 1, 2, 3)
12EPOCHS = 12
13BATCH = 128
14LR_GRID = [1e-3, 3e-3, 6e-3]
15
16
17def kernel_attention_math_check(seed=1180):
18 rng = np.random.default_rng(seed); rows = []
19 for k in (1, 2, 3):
20 x = rng.random((1800, k)); q = x[:40]
21 d = np.sqrt(((q[:, None, :] - x[None, :, :]) ** 2).sum(-1))
22 radii = np.geomspace(.04, .16, 7); mass, wr = [], []
23 for r in radii:
24 w = np.exp(-(d / r) ** 2); mass.append(w.mean())
25 a = w / w.sum(1, keepdims=True); wr.append(np.sqrt((a*d*d).sum(1).mean()))
26 rows.append({'dimension': k, 'predicted_mass_slope': k,
27 'observed_mass_slope': float(np.polyfit(np.log(radii), np.log(mass), 1)[0]),
28 'predicted_radius_slope': 1,
29 'observed_radius_slope': float(np.polyfit(np.log(radii), np.log(wr), 1)[0])})
30 return rows
31
32
33class LocalKernelSelfAttention(nn.Module):
34 def __init__(self, d, nhead=2, neighbors=16):
35 super().__init__(); assert d % nhead == 0
36 self.d, self.nhead, self.hd, self.neighbors = d, nhead, d//nhead, neighbors
37 self.qkv, self.out = nn.Linear(d, 3*d), nn.Linear(d, d)
38
39 def forward(self, x):
40 b, t, d = x.shape; z = torch.nn.functional.layer_norm(x, (d,))
41 dist = torch.cdist(z, z).clamp_min(1e-7)
42 eye = torch.eye(t, device=x.device, dtype=x.dtype)
43 masked = dist + eye[None] * 1e6; m = min(self.neighbors, max(1, t-1))
44 knn = masked.topk(m, dim=-1, largest=False).values; r = knn[..., -1:].detach()
45 khat = (1.0/(torch.log((r+1e-7)/(knn[..., :-1]+1e-7)).mean(-1, keepdim=True)+1e-7)).clamp(1., float(d))
46 g = torch.exp(-((dist/r)**2))/(r**khat); g = g.masked_fill(eye.bool()[None], 0.)
47 a = g/(g.sum(-1, keepdim=True)+1e-7)
48 v = self.qkv(x)[..., 2*d:].reshape(b,t,self.nhead,self.hd).transpose(1,2)
49 y = torch.matmul(a[:,None],v).transpose(1,2).reshape(b,t,d)
50 return self.out(y), a
51
52
53class KernelBlock(nn.Module):
54 def __init__(self, d=64, nhead=2, neighbors=16):
55 super().__init__(); self.attn=LocalKernelSelfAttention(d,nhead,neighbors)
56 self.norm1=nn.LayerNorm(d); self.ff=nn.Sequential(nn.Linear(d,128),nn.ReLU(),nn.Linear(128,d)); self.norm2=nn.LayerNorm(d)
57 def forward(self,x):
58 y,a=self.attn(x); x=self.norm1(x+y); return self.norm2(x+self.ff(x)),a
59
60
61class KernelTransformer(nn.Module):
62 def __init__(self, win, out_dim=1, d=64, depth=2, neighbors=16):
63 super().__init__(); self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02)
64 self.blocks=nn.ModuleList([KernelBlock(d,2,neighbors) for _ in range(depth)]); self.head=nn.Linear(win*d,out_dim); self.last_attention=None
65 def forward(self,x):
66 h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]; aa=[]
67 for block in self.blocks: h,a=block(h); aa.append(a)
68 self.last_attention=aa[-1].detach(); return self.head(h.reshape(h.shape[0],-1))
69
70
71def make_ds(seed): return get_dataset('sequence',seed,n_train=400,n_test=200)
72
73
74def trained_signature(net, x):
75 device=next(net.parameters()).device; x=x[:64].to(device)
76 with torch.no_grad():
77 h=net.inp(x.unsqueeze(-1))+net.pos[:,:x.shape[1]]; h=torch.nn.functional.layer_norm(h,(h.shape[-1],))
78 d=torch.cdist(h,h); eye=torch.eye(d.shape[-1],device=device,dtype=d.dtype); md=d+eye[None]*1e6
79 m=min(net.blocks[0].attn.neighbors,d.shape[-1]-1); knn=md.topk(m,dim=-1,largest=False).values; r=knn[...,-1:]
80 kh=(1/(torch.log((r+1e-7)/(knn[...,:-1]+1e-7)).mean(-1,keepdim=True)+1e-7)).clamp(1.,float(h.shape[-1]))
81 gam=np.geomspace(.7,2.,6); eff=[]
82 for g in gam:
83 w=torch.exp(-((d/(r*g))**2)); w=w.masked_fill(eye.bool()[None],0); w=w/(w.sum(-1,keepdim=True)+1e-7); eff.append(float((1/(w.square().sum(-1)+1e-9)).mean()))
84 slope=float(np.polyfit(np.log(gam),np.log(eff),1)[0]); pred=float(kh.mean())
85 ent=float(-(net.last_attention.clamp_min(1e-9)*net.last_attention.clamp_min(1e-9).log()).sum(-1).mean())
86 return {'predicted_dimension':pred,'predicted_support_slope':pred,'observed_support_slope':slope,'kernel_entropy_observed':ent,'effective_neighbors_observed':eff[2],'confirmed':bool(abs(slope-pred)<max(.75,.35*pred))}
87
88
89def train_eval(seed, idea, lr, capture=False):
90 torch.manual_seed(seed); np.random.seed(seed); ds=make_ds(seed)
91 net=KernelTransformer(ds['input_shape'][0],ds['out_dim']) if idea else make_model('transformer_tiny',ds['input_shape'],ds['out_dim'])
92 net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
93 sig={}
94 if capture and idea:
95 with torch.no_grad(): _=net(ds['xte'][:64].to(next(net.parameters()).device))
96 sig=trained_signature(net,ds['xte'])
97 return float(metric),sig
98
99
100def main():
101 t0=time.time(); math_rows=kernel_attention_math_check()
102 grid=[{'lr':x} for x in LR_GRID]
103 base=sweep_baseline(lambda c:lambda s:train_eval(s,False,c['lr'])[0],grid,seeds=SWEEP_SEEDS)
104 idea_sweep=[]
105 for lr in LR_GRID:
106 vals=[train_eval(s,True,lr)[0] for s in SWEEP_SEEDS]; idea_sweep.append({'cfg':{'lr':lr,'neighbors':16},'mean':float(np.mean(vals))})
107 best_lr=min(idea_sweep,key=lambda z:z['mean'])['cfg']['lr']; vals=[]; sig={}
108 for s in SEEDS:
109 v,sg=train_eval(s,True,best_lr,capture=(s==0)); vals.append(v)
110 if sg: sig=sg
111 idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals),'cfg':{'lr':best_lr,'neighbors':16},'sweep':idea_sweep}
112 report=make_report('sequence','transformer_tiny',base,idea,{'prediction':'support slope versus bandwidth equals local intrinsic dimension','observed':sig,'confirmed':sig.get('confirmed',False)})
113 report['math_check']=math_rows; report['runtime_sec']=time.time()-t0; Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
114
115if __name__=='__main__': main()