Gauge-Covariant Wilson-Loop Regularization / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
 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()