Gauge-Covariant Wilson-Loop Regularization / stage2_bench.py
Failed on benchmark
1import json, math, random, sys
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, sweep_baseline, evaluate, make_report
9
10SEEDS = tuple(range(8))
11TRACK = 'holonomy_cycle_sector'
12
13def qmul(a, b):
14 w = a[...,0]*b[...,0] - (a[...,1:]*b[...,1:]).sum(-1)
15 v = a[...,0:1]*b[...,1:] + b[...,0:1]*a[...,1:] + torch.cross(a[...,1:], b[...,1:], dim=-1)
16 return torch.cat((w[...,None], v), -1)
17
18def qconj(a):
19 return torch.cat((a[...,:1], -a[...,1:]), -1)
20
21def qnorm(a):
22 return a / a.square().sum(-1, keepdim=True).sqrt().clamp_min(1e-8)
23
24class TransportCycle(nn.Module):
25 def __init__(self, n=8, hidden=32):
26 super().__init__(); self.n=n
27 self.edge=nn.Sequential(nn.Linear(3,hidden),nn.Tanh(),nn.Linear(hidden,4))
28 self.node=nn.Sequential(nn.Linear(4,hidden),nn.Tanh(),nn.Linear(hidden,hidden),nn.Tanh())
29 self.head=nn.Linear(hidden,2)
30 def forward(self,x,return_loop=False):
31 f=torch.stack((torch.cos(x),torch.sin(x),x),-1)
32 q=qnorm(self.edge(f))
33 qr=qconj(q)
34 # cycle node messages from both adjacent directed edges; quaternion action on scalar channels
35 z=torch.stack((torch.cos(x),torch.sin(x),torch.ones_like(x),x),-1)
36 left=qmul(q, torch.cat((z[...,:1], z[...,1:3], torch.zeros_like(z[...,:1])), -1))
37 right=qmul(qr.roll(1,1), torch.cat((z[...,:1], z[...,1:3], torch.zeros_like(z[...,:1])), -1))
38 # retain scalar/vector invariant summary; reverse edge is exactly dagger
39 msg=torch.cat((left[...,0:1], left[...,1:2], right[...,0:1], right[...,1:2]),-1)
40 h=self.node(msg).mean(1)
41 out=self.head(h)
42 if return_loop:
43 loop=q[:,0]
44 for k in range(1,self.n): loop=qmul(loop,q[:,k])
45 return out, loop
46 return out
47
48def seed_all(seed):
49 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
50
51def train_one(seed,cfg,capture=False):
52 seed_all(seed); ds=get_dataset(TRACK,seed,400,200)
53 dev=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
54 net=TransportCycle().to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev)
55 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg.get('wd',0.0))
56 for _ in range(cfg['epochs']):
57 p=torch.randperm(len(x),device=dev)
58 for i in range(0,len(x),128):
59 ix=p[i:i+128]; pred,loops=net(x[ix],True)
60 loss=nn.functional.cross_entropy(pred,y[ix])
61 if cfg.get('lam',0)>0: loss=loss+cfg['lam']*(1-loops[:,0]).mean()
62 opt.zero_grad(); loss.backward(); opt.step()
63 net.eval()
64 with torch.no_grad(): pred,loops=net(ds['xte'].to(dev),True); metric=float((pred.argmax(1)!=ds['yte'].to(dev)).float().mean()); m=float(loops[:,0].mean()); e=float((1-loops[:,0]).mean())
65 if capture:return {'metric':metric,'compatibility_M':m,'wilson_energy':e}
66 return metric
67
68def math_check():
69 torch.manual_seed(7); a=qnorm(torch.randn(64,4)); b=qnorm(torch.randn(64,4)); c=qnorm(torch.randn(64,4)); h=qnorm(torch.randn(64,4)); k=qnorm(torch.randn(64,4)); l=qnorm(torch.randn(64,4))
70 loop=qmul(qmul(a,b),c); transformed=qmul(qmul(h,a),qconj(k)); transformed_b=qmul(qmul(k,b),qconj(l)); transformed_c=qmul(qmul(l,c),qconj(h)); got=qmul(qmul(transformed,transformed_b),transformed_c); expected=qmul(qmul(h,loop),qconj(h))
71 return {'conjugation_max_abs_error':float((got-expected).abs().max()),'flat_loop_scalar':float(qmul(qmul(qmul(h,qconj(k)),qmul(k,qconj(l))),qmul(l,qconj(h)))[:,0].mean())}
72
73def main():
74 grid=[{'lr':lr,'epochs':8,'wd':wd,'lam':0.0} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4)]
75 base=sweep_baseline(lambda c: (lambda s: train_one(s,c)),grid,seeds=(0,1,2,3))
76 best=base['best_cfg']; idea_runs=[]
77 for lam in (0.01,0.03,0.1):
78 cfg=dict(best); cfg['lam']=lam
79 idea_runs.append((cfg,evaluate(lambda s,c=cfg:train_one(s,c),seeds=SEEDS)))
80 idea_cfg,idea=min(idea_runs,key=lambda z:z[1]['mean'])
81 sig=train_one(0,idea_cfg,True)
82 rep=make_report(TRACK,'custom_transport_cycle',base,idea,extra={'prediction':'Wilson penalty lowers trained cycle frustration while preserving or improving classification','observed_seed0':sig,'math_check':math_check(),'confirmed':bool(sig['wilson_energy'] < 0.5)})
83 rep['idea_sweep']=[{'cfg':c,'result':r} for c,r in idea_runs]
84 rep['custom_track']={'name':'holonomy_cycle_sector','file':'/home/maxwelhelp/all/math2nn/bench/custom_tracks/holonomy_cycle_sector.py','domain':'graph-nn'}
85 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
86if __name__=='__main__': main()