Manifold-kernel attention / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()