Signed spectral attention / signed_attention_bench.py

Mechanism confirmed, baseline not beaten

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