import sys, json, time from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, sweep_baseline, make_report SEEDS=list(range(8)); EPOCHS=5; NTRAIN=600; NTEST=300 LRS=[0.0015,0.003,0.006] class PGWAttention(nn.Module): def __init__(self,d,heads=2,beta=.3,eps=.15,iters=2): super().__init__(); assert d%heads==0 self.d=d; self.h=heads; self.dk=d//heads; self.beta=beta; self.eps=eps; self.iters=iters self.q=nn.Linear(d,d); self.k=nn.Linear(d,d); self.v=nn.Linear(d,d); self.o=nn.Linear(d,d) self.stats={} def forward(self,x): b,l,_=x.shape; q=self.q(x).view(b,l,self.h,self.dk).transpose(1,2); k=self.k(x).view(b,l,self.h,self.dk).transpose(1,2); v=self.v(x).view(b,l,self.h,self.dk).transpose(1,2) c=1-F.normalize(q,dim=-1)@F.normalize(k,dim=-1).transpose(-1,-2) A=torch.cdist(q,q); B=torch.cdist(k,k); bound=1/l T=torch.softmax(-c/self.eps,dim=-1)/l for _ in range(self.iters): # Equivalent expansion of sum_kl (A_ik-B_jl)^2 T_kl. L=(A.square()@T.sum(-1).unsqueeze(-1)+(T@B.transpose(-1,-2).square()).transpose(-1,-2)-2*torch.einsum('bhik,bhkl,bhjl->bhij',A,T,B)).clamp_min(0).detach() T=torch.exp((-(c+self.beta*L)/self.eps).clamp(-25,8)) for __ in range(4): T=T*torch.minimum(torch.ones_like(T.sum(-1)),bound/T.sum(-1).clamp_min(1e-8))[...,None] T=T*torch.minimum(torch.ones_like(T.sum(-2)),bound/T.sum(-2).clamp_min(1e-8))[...,None,:] z=(T@v)/T.sum(-1,keepdim=True).clamp_min(1e-6) self.stats={'mass':float(T.sum(-1).mean().detach()),'entropy':float((-(T.clamp_min(1e-12)*T.clamp_min(1e-12).log()).sum(-1)).mean().detach())} return self.o(z.transpose(1,2).reshape(b,l,self.d)) class Block(nn.Module): def __init__(self,beta,eps): super().__init__(); self.n1=nn.LayerNorm(64); self.a=PGWAttention(64,beta=beta,eps=eps); self.n2=nn.LayerNorm(64); self.ff=nn.Sequential(nn.Linear(64,128),nn.ReLU(),nn.Linear(128,64)) def forward(self,x): x=x+self.a(self.n1(x)); return x+self.ff(self.n2(x)) class PGWTransformer(nn.Module): def __init__(self,beta=.3,eps=.15): super().__init__(); self.inp=nn.Linear(1,64); self.pos=nn.Parameter(torch.zeros(1,32,64)); nn.init.normal_(self.pos,std=.02); self.blocks=nn.ModuleList([Block(beta,eps) for _ in range(2)]); self.head=nn.Linear(32*64,1) def forward(self,x): h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]] for z in self.blocks:h=z(h) return self.head(h.reshape(x.shape[0],-1)) def run(fn,seed,cfg): torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('sequence',seed,n_train=NTRAIN,n_test=NTEST) net,m,_=train_model(fn(cfg),d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a:None) return float(m) if net is not None else float('inf') def base(cfg): return make_model('transformer_tiny',(32,),1) def idea(cfg): return PGWTransformer(cfg['beta'],cfg['eps']) def main(): t=time.time(); grid=[{'lr':v} for v in LRS] base_block=sweep_baseline(lambda c:lambda s:run(base,s,c),grid,seeds=SEEDS[:4]) # exact learning-rate parity: all idea lr values are in baseline grid; beta/eps are idea knobs. ig=[{'lr':.0015,'beta':.2,'eps':.15},{'lr':.003,'beta':.3,'eps':.15},{'lr':.006,'beta':.4,'eps':.20}] dev=[{'cfg':c,'mean':float(np.mean([run(idea,s,c) for s in SEEDS[:4]]))} for c in ig] best=min(dev,key=lambda z:z['mean'])['cfg']; vals=[run(idea,s,best) for s in SEEDS] ir={'best_cfg':best,'sweep':dev,'per_seed':vals} obs={'baseline_test_mse_mean':float(np.mean(base_block['full']['per_seed'])),'idea_test_mse_mean':float(np.mean(vals))} rep=make_report('sequence','transformer_tiny',base_block,ir,{'prediction':'relational compatibility reduces incompatible-token attention','predicted':{'partial_row_upper_bound':1/32,'beta':best['beta']},'observed':obs,'confirmed':False}) rep['protocol_notes']={'track_reason':'The sequence track is multi-token temporal forecasting with transformer attention, directly matching the proposed mechanism.','reduced_budget':'600 train/300 test and 5 epochs due to PGW O(L^4) cost; paired seeds remain 8; baseline sweep uses the same lr union.','elapsed_sec':time.time()-t} Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2)) if __name__=='__main__':main()