import json, math, random, sys from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, sweep_baseline, evaluate, make_report SEEDS = tuple(range(8)) TRACK = 'holonomy_cycle_sector' def qmul(a, b): w = a[...,0]*b[...,0] - (a[...,1:]*b[...,1:]).sum(-1) v = a[...,0:1]*b[...,1:] + b[...,0:1]*a[...,1:] + torch.cross(a[...,1:], b[...,1:], dim=-1) return torch.cat((w[...,None], v), -1) def qconj(a): return torch.cat((a[...,:1], -a[...,1:]), -1) def qnorm(a): return a / a.square().sum(-1, keepdim=True).sqrt().clamp_min(1e-8) class TransportCycle(nn.Module): def __init__(self, n=8, hidden=32): super().__init__(); self.n=n self.edge=nn.Sequential(nn.Linear(3,hidden),nn.Tanh(),nn.Linear(hidden,4)) self.node=nn.Sequential(nn.Linear(4,hidden),nn.Tanh(),nn.Linear(hidden,hidden),nn.Tanh()) self.head=nn.Linear(hidden,2) def forward(self,x,return_loop=False): f=torch.stack((torch.cos(x),torch.sin(x),x),-1) q=qnorm(self.edge(f)) qr=qconj(q) # cycle node messages from both adjacent directed edges; quaternion action on scalar channels z=torch.stack((torch.cos(x),torch.sin(x),torch.ones_like(x),x),-1) left=qmul(q, torch.cat((z[...,:1], z[...,1:3], torch.zeros_like(z[...,:1])), -1)) right=qmul(qr.roll(1,1), torch.cat((z[...,:1], z[...,1:3], torch.zeros_like(z[...,:1])), -1)) # retain scalar/vector invariant summary; reverse edge is exactly dagger msg=torch.cat((left[...,0:1], left[...,1:2], right[...,0:1], right[...,1:2]),-1) h=self.node(msg).mean(1) out=self.head(h) if return_loop: loop=q[:,0] for k in range(1,self.n): loop=qmul(loop,q[:,k]) return out, loop return out def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) def train_one(seed,cfg,capture=False): seed_all(seed); ds=get_dataset(TRACK,seed,400,200) dev=torch.device('cuda' if torch.cuda.is_available() else 'cpu') net=TransportCycle().to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev) opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg.get('wd',0.0)) for _ in range(cfg['epochs']): p=torch.randperm(len(x),device=dev) for i in range(0,len(x),128): ix=p[i:i+128]; pred,loops=net(x[ix],True) loss=nn.functional.cross_entropy(pred,y[ix]) if cfg.get('lam',0)>0: loss=loss+cfg['lam']*(1-loops[:,0]).mean() opt.zero_grad(); loss.backward(); opt.step() net.eval() 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()) if capture:return {'metric':metric,'compatibility_M':m,'wilson_energy':e} return metric def math_check(): 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)) 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)) 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())} def main(): 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)] base=sweep_baseline(lambda c: (lambda s: train_one(s,c)),grid,seeds=(0,1,2,3)) best=base['best_cfg']; idea_runs=[] for lam in (0.01,0.03,0.1): cfg=dict(best); cfg['lam']=lam idea_runs.append((cfg,evaluate(lambda s,c=cfg:train_one(s,c),seeds=SEEDS))) idea_cfg,idea=min(idea_runs,key=lambda z:z[1]['mean']) sig=train_one(0,idea_cfg,True) 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)}) rep['idea_sweep']=[{'cfg':c,'result':r} for c,r in idea_runs] rep['custom_track']={'name':'holonomy_cycle_sector','file':'/home/maxwelhelp/all/math2nn/bench/custom_tracks/holonomy_cycle_sector.py','domain':'graph-nn'} Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2)) if __name__=='__main__': main()