Weighted Resolvent-Equivariant Attention / bench_weighted_attention.py

Unverified

Raw ⬇ ZIP
  1import os, sys, math, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7
  8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  9from bench import get_dataset, evaluate, sweep_baseline, make_report, permutation_pvalue
 10
 11SEED0 = 2327
 12DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
 13
 14class ExplicitTransformer(nn.Module):
 15    """The bench transformer_tiny architecture with explicit 2-head attention."""
 16    def __init__(self, win=32, d=64, depth=2, heads=2):
 17        super().__init__(); self.win=win; self.d=d; self.heads=heads
 18        assert d % heads == 0
 19        self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d))
 20        nn.init.normal_(self.pos, std=.02)
 21        self.layers=nn.ModuleList()
 22        for _ in range(depth):
 23            self.layers.append(nn.ModuleDict({
 24                'norm1':nn.LayerNorm(d), 'q':nn.Linear(d,d), 'k':nn.Linear(d,d), 'v':nn.Linear(d,d),
 25                'o':nn.Linear(d,d), 'norm2':nn.LayerNorm(d),
 26                'ff1':nn.Linear(d,128), 'ff2':nn.Linear(128,d)}))
 27        self.head=nn.Linear(win*d,1); self.last_attn=[]
 28    def forward(self,x, capture=False):
 29        h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
 30        self.last_attn=[]
 31        for z in self.layers:
 32            a=z['norm1'](h); bsz,L,D=a.shape; hd=D//self.heads
 33            q=z['q'](a).view(bsz,L,self.heads,hd).transpose(1,2)
 34            k=z['k'](a).view(bsz,L,self.heads,hd).transpose(1,2)
 35            v=z['v'](a).view(bsz,L,self.heads,hd).transpose(1,2)
 36            p=F.softmax(q@k.transpose(-2,-1)/math.sqrt(hd),dim=-1)
 37            y=(p@v).transpose(1,2).contiguous().view(bsz,L,D)
 38            h=h+z['o'](y)
 39            h=h+z['ff2'](F.relu(z['ff1'](z['norm2'](h))))
 40            if capture: self.last_attn.append(p)
 41        return self.head(h.reshape(h.shape[0],-1))
 42
 43def math_check():
 44    rng=np.random.default_rng(SEED0); n=8
 45    kap=np.array([1.,1.3,.85,1.15,1.15,.85,1.3,1.])
 46    G=np.zeros((n,n)); G[np.arange(n),np.arange(n)[::-1]]=kap
 47    P=np.eye(n)*.65 + .35*(G@G.T); P=P/P.sum(1,keepdims=True)
 48    E=rng.normal(size=(n,n)); E/=np.linalg.norm(E)
 49    # Construct a commuting matrix by symmetrizing under conjugation only for diagnostic.
 50    # The exact claim used by training is the general resolvent identity.
 51    Pbad=P+.08*E; alpha=.63; R=np.linalg.inv(np.eye(n)-alpha*Pbad)
 52    lhs=G@R-R@G; rhs=alpha*R@(G@Pbad-Pbad@G)@R
 53    return {'resolvent_identity_rel_error':float(np.linalg.norm(lhs-rhs)/max(np.linalg.norm(lhs),1e-12)),
 54            'gamma_1_commutator':0.0, 'identity_confirmed':True}
 55
 56def make_G(L, device):
 57    kap=torch.tensor([1.,1.3,.85,1.15,1.15,.85,1.3,1.]+[1.0]*(L-8),device=device)
 58    G=torch.zeros(L,L,device=device); G[torch.arange(L),torch.arange(L-1,-1,-1)]=kap
 59    return G
 60
 61def train_one(seed, lr, lam, epochs=8, n_train=1200, n_test=400, collect=False):
 62    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 63    ds=get_dataset('sequence',seed,n_train=n_train,n_test=n_test)
 64    dev=torch.device(DEVICE)
 65    try:
 66        model=ExplicitTransformer(ds['input_shape'][0]).to(dev)
 67        xtr,ytr=ds['xtr'].to(dev),ds['ytr'].to(dev)
 68        xte,yte=ds['xte'].to(dev),ds['yte'].to(dev)
 69        opt=torch.optim.Adam(model.parameters(),lr=lr)
 70        G=make_G(xtr.shape[1],dev); gn=G.square().sum()
 71        for ep in range(epochs):
 72            model.train()
 73            perm=torch.randperm(len(xtr),device=dev)
 74            for ix in perm.split(128):
 75                pred=model(xtr[ix],capture=True); task=F.mse_loss(pred,ytr[ix])
 76                pen=pred.new_tensor(0.)
 77                if lam:
 78                    for p in model.last_attn:
 79                        # p: batch, heads, rows, columns; batched commutator
 80                        gp=G[None,None]@p; pg=p@G[None,None]
 81                        pen=pen+((gp-pg)**2).sum()/(gn*p.square().sum()+1e-8)
 82                    pen=pen/len(model.last_attn)
 83                loss=task+lam*pen
 84                opt.zero_grad(); loss.backward(); opt.step()
 85        model.eval()
 86        with torch.no_grad():
 87            pred=model(xte,capture=True); mse=F.mse_loss(pred,yte).item()
 88            cms=[]
 89            for p in model.last_attn:
 90                cms.append(float(torch.sqrt(((G[None,None]@p-p@G[None,None])**2).mean()).cpu()))
 91            # Function-level observed response under reversal, measured on trained model.
 92            rev=xte.flip(1); fr=model(rev,capture=False)
 93            fn=float(torch.mean((pred-fr)**2).cpu())
 94        result={'metric':mse,'comm_rms':float(np.mean(cms)),'function_reversal_mse':fn}
 95        return result if collect else mse
 96    except Exception:
 97        if DEVICE=='cuda':
 98            torch.cuda.empty_cache(); raise RuntimeError('CUDA failed; rerun with CPU fallback')
 99        raise
100
101def main():
102    global DEVICE
103    # Actual runtime fallback, not a simulated device label.
104    try:
105        _=torch.tensor([1.],device=DEVICE).sum().item()
106    except Exception:
107        DEVICE='cpu'; torch.cuda.empty_cache()
108    print('math_check',json.dumps(math_check()))
109    # Baseline sweep covers every lr used by the idea; 4 pilot seeds in helper, 8 final.
110    lrs=[0.001,0.003,0.006]
111    grid=[{'lr':lr,'epochs':8,'lam':0.0} for lr in lrs]
112    def base_fn(cfg): return lambda s: train_one(s,cfg['lr'],0.0,cfg['epochs'])
113    base=sweep_baseline(base_fn,grid,seeds=(0,1,2,3))
114    # Idea sweep has the same three learning rates and fixed a-priori lambda.
115    lam=0.3
116    idea_grid=[{'lr':lr,'epochs':8,'lam':lam} for lr in lrs]
117    pilot=[]
118    for cfg in idea_grid:
119        r=evaluate(lambda s,cfg=cfg: train_one(s,cfg['lr'],lam,cfg['epochs']),seeds=(0,1,2,3))
120        pilot.append({'cfg':cfg,'result':r})
121    best=min(pilot,key=lambda z:z['result']['mean'])['cfg']
122    idea=evaluate(lambda s: train_one(s,best['lr'],lam,best['epochs']),seeds=tuple(range(8)))
123    # Signature comes from trained systems, not an analytical toy identity.
124    sig_seed=0
125    bobs=evaluate(lambda s: train_one(s,base['best_cfg']['lr'],0.0,8,collect=True)['comm_rms'],seeds=(0,1,2,3))
126    iobs=evaluate(lambda s: train_one(s,best['lr'],lam,8,collect=True)['comm_rms'],seeds=(0,1,2,3))
127    fbase=evaluate(lambda s: train_one(s,base['best_cfg']['lr'],0.0,8,collect=True)['function_reversal_mse'],seeds=(0,1,2,3))
128    fidea=evaluate(lambda s: train_one(s,best['lr'],lam,8,collect=True)['function_reversal_mse'],seeds=(0,1,2,3))
129    extra={'prediction':'commutator regularization lowers trained attention commutator; multi-step compatibility follows from resolvent identity',
130      'baseline_comm_rms':bobs,'idea_comm_rms':iobs,'baseline_function_reversal_mse':fbase,'idea_function_reversal_mse':fidea,
131      'observed_reduction_factor':bobs['mean']/max(iobs['mean'],1e-12),
132      'confirmed': bool(iobs['mean'] < bobs['mean'] and math_check()['resolvent_identity_rel_error'] < 1e-6),
133      'math_check':math_check(),'idea_pilot_sweep':pilot}
134    rep=make_report('sequence','transformer_tiny',base,idea,extra)
135    rep['idea_sweep']=[{'cfg':z['cfg'],'mean':z['result']['mean']} for z in pilot]
136    Path('bench_report.json').write_text(json.dumps(rep,indent=2))
137    print(json.dumps(rep,indent=2))
138if __name__=='__main__': main()