Teleporting Simplicial Diffusion Layer / run_bench.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys,json,copy
 2import numpy as np, torch
 3import torch.nn as nn
 4sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 5import bench
 6from bench import train_model,evaluate,sweep_baseline,make_report
 7import simplicial_track
 8TRACK='simplicial_node_classification'
 9bench.data._CUSTOM_CACHE={TRACK:simplicial_track}
10
11def P_local():
12 n=24; A=np.zeros((n,n),dtype=np.float32)
13 for j in range(n): A[j,j]=A[(j+1)%n,j]=A[(j+2)%n,j]=1
14 P=(A/A.sum(1)[:,None])@np.diag(1/A.sum(0))@A.T
15 return (.9*P+.1*np.eye(n)).astype('float32')
16P=P_local()
17class Net(nn.Module):
18 def __init__(self,alpha):
19  super().__init__(); self.register_buffer('P',torch.tensor(P)); self.alpha=alpha
20  self.a=nn.Linear(8,32); self.b=nn.Linear(32,2)
21 def forward(self,x):
22  n=self.P.shape[0]; R=torch.ones_like(self.P)/n
23  Q=(1-self.alpha)*self.P+self.alpha*R
24  h=torch.relu(self.a(x)); h=torch.einsum('ij,bjf->bif',Q,h); return self.b(h).transpose(1,2)
25def train(cfg,seed):
26 torch.manual_seed(1000+seed); np.random.seed(1000+seed)
27 ds=bench.get_dataset(TRACK,seed,400,100); net=Net(cfg['alpha'])
28 _,metric,_=train_model(net,ds,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,weight_decay=cfg['wd'],log=lambda *_:None)
29 return metric
30def factory(cfg): return lambda seed: train(cfg,seed)
31def main():
32 # Union parity: both methods are evaluated on all lr/wd values and alpha values.
33 grid=[{'alpha':a,'lr':lr,'wd':wd,'epochs':18} for lr in (0.003,0.01,0.03) for wd in (0.0,0.0001) for a in (0.0,0.1,0.3)]
34 # Baseline method is alpha=0; central baseline knob wd is swept too.
35 bgrid=[{k:v for k,v in c.items()} for c in grid if c['alpha']==0]
36 base=sweep_baseline(factory,bgrid,seeds=(0,1,2,3))
37 best=base['best_cfg']; idea_cfgs=[]
38 for a in (0.1,0.3,0.6):
39  c=dict(best); c['alpha']=a; idea_cfgs.append(c)
40 # Same-size idea sweep; alpha=.6 is nearby but also uses best baseline lr/wd.
41 ir=[]
42 for c in idea_cfgs:
43  r=evaluate(factory(c)); ir.append({'cfg':c,'result':r})
44 chosen=min(ir,key=lambda z:z['result']['mean']); idea=chosen['result']
45 # Retrain/evaluate observed trained systems for signature on one paired seed.
46 ds=bench.get_dataset(TRACK,0,400,100); x=ds['xte'];
47 sig=[]
48 for a in (0.0,0.1,0.3,0.6):
49  torch.manual_seed(1000); net=Net(a); _,_,_=train_model(net,ds,epochs=best['epochs'],lr=best['lr'],batch=128,weight_decay=best['wd'],log=lambda *_:None)
50  with torch.no_grad():
51   xdev=x.to(net.P.device); q=net.a(xdev); q=torch.einsum('ij,bjf->bif',((1-a)*net.P+a*torch.ones_like(net.P)/24),q)
52   v=q-q.mean(1,keepdim=True); ratio=float(v.norm()/ (torch.relu(net.a(xdev))-torch.relu(net.a(xdev)).mean(1,keepdim=True)).norm())
53  sig.append({'alpha':a,'observed_nonconstant_norm_ratio':ratio,'predicted':1-a})
54 sigerr=max(abs(z['observed_nonconstant_norm_ratio']-z['predicted']) for z in sig)
55 rep=make_report(TRACK,'custom_diffusion_node_net',base,idea,{'prediction':'teleportation contracts mean-zero feature modes by 1-alpha','rows':sig,'max_abs_error':sigerr,'confirmed':bool(sigerr<0.08),'idea_sweep':ir,'custom_track':{'name':TRACK,'file':'simplicial_track.py','domain':'graph-nn'}})
56 rep['idea_candidates']=ir
57 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
58 print(json.dumps(rep,indent=2))
59if __name__=='__main__': main()