Signed spectral attention / signed_attention_bench.py
Mechanism confirmed, baseline not beaten
1from __future__ import annotations
2import json, math, random, time
3from pathlib import Path
4import numpy as np
5import torch
6import torch.nn as nn
7
8import sys
9sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
10from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
11
12SEED = 2820
13NTR, NTE, EPOCHS, BATCH = 400, 200, 6, 128
14D, DEPTH, HEADS = 32, 1, 2
15A, B, SIGMA_POS, SIGMA_NEG = 1.0, 0.8, 0.5, 2.0
16
17
18def set_seed(seed):
19 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
20 if torch.cuda.is_available():
21 try: torch.cuda.manual_seed_all(seed)
22 except Exception: pass
23
24
25def sanity_check():
26 rng = np.random.default_rng(SEED)
27 u = rng.uniform(-3, 3, 300)
28 truth = A*np.exp(-0.5*(SIGMA_POS*u)**2) - B*np.exp(-0.5*(SIGMA_NEG*u)**2)
29 rows = []
30 for m in (16, 32, 64, 128, 256):
31 es, eu = [], []
32 for t in range(30):
33 rr = np.random.default_rng(SEED + 1000*m + t)
34 neg = rr.random(m) < B/(A+B)
35 om = np.where(neg, rr.normal(0, SIGMA_NEG, m), rr.normal(0, SIGMA_POS, m))
36 c = np.cos(u[:, None]*om)
37 signed = (A+B)*np.mean(c*np.where(neg, -1., 1.), axis=1)
38 unsigned = (A+B)*np.mean(c, axis=1)
39 es.append(np.mean((signed-truth)**2)); eu.append(np.mean((unsigned-truth)**2))
40 rows.append((m, float(np.sqrt(np.mean(es))), float(np.sqrt(np.mean(eu)))))
41 slope = float(np.polyfit(np.log([r[0] for r in rows]), np.log([r[1] for r in rows]), 1)[0])
42 positive_floor = float(np.sqrt(np.mean((A*np.exp(-.5*(SIGMA_POS*u)**2) + B*np.exp(-.5*(SIGMA_NEG*u)**2)-truth)**2)))
43 return {"rows": [{"M":m,"signed_rmse":s,"positive_rmse":p} for m,s,p in rows],
44 "observed_slope": slope, "predicted_slope": -.5,
45 "positive_predicted_floor": positive_floor,
46 "passed": bool(slope < -.25 and rows[-1][1] < rows[0][1] and rows[-1][2] > positive_floor*.95)}
47
48
49class SoftmaxBlock(nn.Module):
50 def __init__(self, d, heads, temperature=1.0):
51 super().__init__(); self.temperature = temperature
52 self.attn = nn.MultiheadAttention(d, heads, dropout=0., batch_first=True)
53 self.n1 = nn.LayerNorm(d); self.ff = nn.Sequential(nn.Linear(d, 64), nn.GELU(), nn.Linear(64,d)); self.n2=nn.LayerNorm(d)
54 def forward(self, x):
55 h = self.n1(x)
56 # Apply the swept temperature to the standard attention logits.
57 w = self.attn.in_proj_weight
58 b = self.attn.in_proj_bias
59 q = torch.nn.functional.linear(h, w[:self.attn.embed_dim], b[:self.attn.embed_dim])
60 k = torch.nn.functional.linear(h, w[self.attn.embed_dim:2*self.attn.embed_dim], b[self.attn.embed_dim:2*self.attn.embed_dim])
61 v = torch.nn.functional.linear(h, w[2*self.attn.embed_dim:], b[2*self.attn.embed_dim:])
62 q = q.view(x.shape[0], x.shape[1], self.attn.num_heads, -1).transpose(1,2)
63 k = k.view(x.shape[0], x.shape[1], self.attn.num_heads, -1).transpose(1,2)
64 v = v.view(x.shape[0], x.shape[1], self.attn.num_heads, -1).transpose(1,2)
65 logits = torch.matmul(q, k.transpose(-2,-1)) / (math.sqrt(q.shape[-1]) * self.temperature)
66 y = torch.matmul(torch.softmax(logits, dim=-1), v).transpose(1,2).reshape_as(h)
67 y = self.attn.out_proj(y)
68 return self.n2(x+y + self.ff(self.n2(x+y)))
69
70
71class SignedBlock(nn.Module):
72 def __init__(self, d, m, seed):
73 super().__init__(); self.d=d; self.m=m; self.C=A+B
74 rr=np.random.default_rng(seed)
75 neg=rr.random(m)<B/(A+B)
76 om=np.where(neg, rr.normal(0,SIGMA_NEG,(m,d)), rr.normal(0,SIGMA_POS,(m,d))).astype(np.float32)
77 self.register_buffer('omega', torch.from_numpy(om)); self.register_buffer('signs', torch.from_numpy(np.where(neg,-1.,1.).astype(np.float32)))
78 self.q=nn.Linear(d,d,bias=False); self.k=nn.Linear(d,d,bias=False); self.v=nn.Linear(d,d,bias=False); self.o=nn.Linear(d,d)
79 self.n1=nn.LayerNorm(d); self.ff=nn.Sequential(nn.Linear(d,64),nn.GELU(),nn.Linear(64,d)); self.n2=nn.LayerNorm(d)
80 def features(self, x, m=None):
81 om=self.omega[:m] if m else self.omega
82 a=x @ om.T
83 return torch.cat((torch.cos(a), torch.sin(a)), dim=-1)
84 def forward(self,x):
85 h=self.n1(x); q=self.features(self.q(h)); k=self.features(self.k(h)); v=self.v(h)
86 s=torch.repeat_interleave(self.signs,2); scale=self.C/self.m
87 # Associative factorization: Zq D (Zk^T V), with positive feature normalization.
88 kv=torch.einsum('bnd,bnr->bdr', k*s, v)
89 out=scale*torch.einsum('bnd,bdr->bnr',q,kv)
90 den=scale*torch.einsum('bnd,bd->bn', q, k.sum(dim=1)).unsqueeze(-1)
91 # Positive denominator uses |spectral signs|, avoiding signed cancellation.
92 den=den.abs().clamp_min(1e-3)
93 y=self.o(out/den)
94 z=x+y; return self.n2(z+self.ff(z))
95
96
97class Net(nn.Module):
98 def __init__(self, m=None, temp=1., seed=0):
99 super().__init__(); self.inp=nn.Linear(1,D); self.pos=nn.Parameter(torch.zeros(1,32,D)); nn.init.normal_(self.pos,std=.02)
100 self.block=SignedBlock(D,m,seed+91) if m else SoftmaxBlock(D,HEADS,temp)
101 self.head=nn.Linear(32*D,1)
102 def forward(self,x):
103 h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
104 return self.head(self.block(h).reshape(x.shape[0],-1))
105
106
107def run(cfg, idea, seed, capture=False):
108 set_seed(seed); ds=get_dataset('sequence',seed,n_train=NTR,n_test=NTE)
109 net=Net(m=cfg['m'], seed=seed) if idea else Net(temp=cfg['temperature'])
110 net, metric, hist=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
111 if capture: captured[(idea,seed)] = net
112 return float(metric)
113
114
115def signature():
116 net=captured.get((True,0));
117 if net is None: return {"confirmed":False,"reason":"model capture failed"}
118 net.eval()
119 dev=next(net.parameters()).device
120 x=get_dataset('sequence',0,n_train=NTR,n_test=32)['xte'].to(dev)
121 with torch.no_grad():
122 h=net.inp(x.unsqueeze(-1))+net.pos[:,:32]
123 q=net.block.q(net.block.n1(h)); k=net.block.k(net.block.n1(h));
124 q=q.reshape(-1,D); k=k.reshape(-1,D); u=q[:128]-k[:128]
125 exact=A*torch.exp(-.5*(SIGMA_POS*u).pow(2).sum(1))-B*torch.exp(-.5*(SIGMA_NEG*u).pow(2).sum(1))
126 M=net.block.m; vals=[]
127 for mm in sorted(set((max(4,M//4),max(8,M//2),M))):
128 zq=net.block.features(q[:128],mm); zk=net.block.features(k[:128],mm); s=torch.repeat_interleave(net.block.signs[:mm],2)
129 est=(A+B)/mm*(zq*(zk*s)).sum(1); vals.append((mm,float(torch.sqrt(torch.mean((est-exact)**2)))))
130 slope=float(np.polyfit(np.log([a for a,b in vals]),np.log([b for a,b in vals]),1)[0]) if len(vals)>1 else float('nan')
131 unsigned=(A+B)/M*(net.block.features(q[:128],M)*(net.block.features(k[:128],M))).sum(1)
132 ur=float(torch.sqrt(torch.mean((unsigned-exact)**2))); sr=vals[-1][1]
133 return {"trained_model":True,"M":M,"observed_rmse_by_M":[{"M":a,"rmse":b} for a,b in vals],"observed_slope":slope,"predicted_slope":-.5,"signed_rmse":sr,"unsigned_rmse":ur,"confirmed":bool(sr<ur and slope < -.15)}
134
135
136if __name__=='__main__':
137 print(json.dumps({"sanity_check":sanity_check()},indent=2))
138 captured={}
139 # Union of all learning rates is evaluated by baseline; temperature is the baseline's method knob.
140 lrs=[1e-3,3e-3,6e-3]; temps=[0.7,1.0]
141 grid=[{"lr":lr,"temperature":t} for lr in lrs for t in temps]
142 base=sweep_baseline(lambda cfg: lambda s: run(cfg,False,s,capture=False),grid)
143 best_lr=base['best_cfg']['lr']
144 # Three-point idea sweep: selected baseline lr plus two nearby rates; all are in baseline grid.
145 idea_cfgs=[{"lr":lr,"m":32} for lr in (1e-3,3e-3,6e-3)]
146 idea_all=[]
147 for cfg in idea_cfgs:
148 r=evaluate(lambda s,cfg=cfg: run(cfg,True,s,capture=False))
149 idea_all.append({"cfg":cfg,"result":r})
150 best=min(idea_all,key=lambda z:z['result']['mean']); bestcfg=best['cfg']
151 idea=evaluate(lambda s: run(bestcfg,True,s,capture=(s==0)))
152 extra=signature()
153 report=make_report('sequence','transformer_tiny',base,idea,extra)
154 report['idea_sweep']=idea_all; report['sanity_check']=sanity_check(); report['protocol']={"seeds":list(range(8)),"n_train":NTR,"n_test":NTE,"epochs":EPOCHS,"baseline_grid":grid,"idea_grid":idea_cfgs}
155 Path('bench_report.json').write_text(json.dumps(report,indent=2))
156 print(json.dumps(report,indent=2))