Partial Gromov-Wasserstein Cross-Attention / bench_pgw.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 1import sys, json, time
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6import torch.nn.functional as F
 7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 8from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
 9
10SEEDS=list(range(8)); EPOCHS=5; NTRAIN=600; NTEST=300
11LRS=[0.0015,0.003,0.006]
12
13class PGWAttention(nn.Module):
14 def __init__(self,d,heads=2,beta=.3,eps=.15,iters=2):
15  super().__init__(); assert d%heads==0
16  self.d=d; self.h=heads; self.dk=d//heads; self.beta=beta; self.eps=eps; self.iters=iters
17  self.q=nn.Linear(d,d); self.k=nn.Linear(d,d); self.v=nn.Linear(d,d); self.o=nn.Linear(d,d)
18  self.stats={}
19 def forward(self,x):
20  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)
21  c=1-F.normalize(q,dim=-1)@F.normalize(k,dim=-1).transpose(-1,-2)
22  A=torch.cdist(q,q); B=torch.cdist(k,k); bound=1/l
23  T=torch.softmax(-c/self.eps,dim=-1)/l
24  for _ in range(self.iters):
25   # Equivalent expansion of sum_kl (A_ik-B_jl)^2 T_kl.
26   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()
27   T=torch.exp((-(c+self.beta*L)/self.eps).clamp(-25,8))
28   for __ in range(4):
29    T=T*torch.minimum(torch.ones_like(T.sum(-1)),bound/T.sum(-1).clamp_min(1e-8))[...,None]
30    T=T*torch.minimum(torch.ones_like(T.sum(-2)),bound/T.sum(-2).clamp_min(1e-8))[...,None,:]
31  z=(T@v)/T.sum(-1,keepdim=True).clamp_min(1e-6)
32  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())}
33  return self.o(z.transpose(1,2).reshape(b,l,self.d))
34class Block(nn.Module):
35 def __init__(self,beta,eps):
36  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))
37 def forward(self,x): x=x+self.a(self.n1(x)); return x+self.ff(self.n2(x))
38class PGWTransformer(nn.Module):
39 def __init__(self,beta=.3,eps=.15):
40  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)
41 def forward(self,x):
42  h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
43  for z in self.blocks:h=z(h)
44  return self.head(h.reshape(x.shape[0],-1))
45def run(fn,seed,cfg):
46 torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('sequence',seed,n_train=NTRAIN,n_test=NTEST)
47 net,m,_=train_model(fn(cfg),d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a:None)
48 return float(m) if net is not None else float('inf')
49def base(cfg): return make_model('transformer_tiny',(32,),1)
50def idea(cfg): return PGWTransformer(cfg['beta'],cfg['eps'])
51def main():
52 t=time.time(); grid=[{'lr':v} for v in LRS]
53 base_block=sweep_baseline(lambda c:lambda s:run(base,s,c),grid,seeds=SEEDS[:4])
54 # exact learning-rate parity: all idea lr values are in baseline grid; beta/eps are idea knobs.
55 ig=[{'lr':.0015,'beta':.2,'eps':.15},{'lr':.003,'beta':.3,'eps':.15},{'lr':.006,'beta':.4,'eps':.20}]
56 dev=[{'cfg':c,'mean':float(np.mean([run(idea,s,c) for s in SEEDS[:4]]))} for c in ig]
57 best=min(dev,key=lambda z:z['mean'])['cfg']; vals=[run(idea,s,best) for s in SEEDS]
58 ir={'best_cfg':best,'sweep':dev,'per_seed':vals}
59 obs={'baseline_test_mse_mean':float(np.mean(base_block['full']['per_seed'])),'idea_test_mse_mean':float(np.mean(vals))}
60 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})
61 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}
62 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
63if __name__=='__main__':main()