Fundamental-Cycle Compatibility Basis / run_bench.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, time, importlib.util
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import make_report, get_dataset as bench_get_dataset
7from custom_cycle_track import get_dataset, EDGES, N
8
9SEEDS = tuple(range(8))
10EPOCHS = 18
11BATCH = 128
12LRS = [1e-3, 3e-3, 1e-2]
13IDEA_LAMBDAS = [0.01, 0.05, 0.2]
14
15def basis_matrix(n, edges):
16 m = len(edges); adj = {i: [] for i in range(2*n)}
17 for i,(x,y) in enumerate(edges):
18 u,v=x,n+y; adj[u].append((v,i,1)); adj[v].append((u,i,-1))
19 parent={}; pe={}; ps={}; seen=set(); tree=set()
20 for root in range(2*n):
21 if root in seen or not adj[root]: continue
22 seen.add(root); stack=[root]
23 while stack:
24 u=stack.pop()
25 for v,e,s in adj[u]:
26 if v not in seen:
27 seen.add(v); parent[v]=u; pe[v]=e; ps[v]=s; tree.add(e); stack.append(v)
28 def path(a,b):
29 anc=set(); u=a
30 while True:
31 anc.add(u)
32 if u not in parent: break
33 u=parent[u]
34 down=[]; u=b
35 while u not in anc:
36 down.append((pe[u],ps[u])); u=parent[u]
37 lca=u; up=[]; u=a
38 while u != lca:
39 up.append((pe[u],-ps[u])); u=parent[u]
40 return up + list(reversed(down))
41 rows=[]
42 for e,(x,y) in enumerate(edges):
43 if e in tree: continue
44 c=np.zeros(m); c[e]=1
45 for j,s in path(n+y,x): c[j]+=s
46 rows.append(c)
47 return np.asarray(rows), tree
48
49def all_simple_cycles(n, edges, cap=10000):
50 adj={i:[] for i in range(2*n)}
51 for i,(x,y) in enumerate(edges):
52 u,v=x,n+y; adj[u].append((v,i,1)); adj[v].append((u,i,-1))
53 found={}
54 def dfs(start,u,vis,c):
55 if len(found)>=cap:return
56 for v,e,s in adj[u]:
57 if v==start and len(vis)>=4:
58 z=c.copy(); z[e]+=s; key=tuple(np.flatnonzero(z))
59 found.setdefault(key,z)
60 elif v not in vis and v>=start:
61 c[e]+=s; dfs(start,v,vis|{v},c); c[e]-=s
62 for s in range(2*n): dfs(s,s,{s},np.zeros(len(edges)))
63 return list(found.values())
64
65class PairNet(nn.Module):
66 def __init__(self):
67 super().__init__()
68 self.q=nn.Sequential(nn.Linear(N,32),nn.Tanh(),nn.Linear(32,N))
69 self.r=nn.Sequential(nn.Linear(N,32),nn.Tanh(),nn.Linear(32,N))
70 def forward(self, y, x): return self.q(y), self.r(x)
71
72def train_one(seed, lr, lam, return_net=False):
73 torch.manual_seed(seed); np.random.seed(seed)
74 d=get_dataset(seed, 400, 400)
75 dev='cuda' if torch.cuda.is_available() else 'cpu'
76 try:
77 net=PairNet().to(dev)
78 opt=torch.optim.Adam(net.parameters(),lr=lr)
79 xt=torch.tensor(d['train_x'],dtype=torch.long,device=dev)
80 yt=torch.tensor(d['train_y'],dtype=torch.long,device=dev)
81 ce=nn.CrossEntropyLoss()
82 edge_x=torch.tensor(EDGES[:,0],dtype=torch.long,device=dev)
83 edge_y=torch.tensor(EDGES[:,1],dtype=torch.long,device=dev)
84 B=torch.tensor(basis_matrix(N,EDGES)[0],dtype=torch.float32,device=dev)
85 for ep in range(EPOCHS):
86 perm=torch.randperm(len(xt),device=dev)
87 for st in range(0,len(xt),BATCH):
88 ii=perm[st:st+BATCH]; lq,lrlog=net(torch.eye(N,device=dev)[yt[ii]],torch.eye(N,device=dev)[xt[ii]])
89 loss=ce(lq,xt[ii])+ce(lrlog,yt[ii])
90 if lam:
91 # All support edges are evaluated, so cycle penalty is independent of sample frequency.
92 aq=net.q(torch.eye(N,device=dev)[edge_y]).log_softmax(1)[range(len(EDGES)),edge_x]
93 ar=net.r(torch.eye(N,device=dev)[edge_x]).log_softmax(1)[range(len(EDGES)),edge_y]
94 loss=loss+lam*(B@(aq-ar)).square().mean()
95 opt.zero_grad(); loss.backward(); opt.step()
96 net.eval()
97 with torch.no_grad():
98 xe=torch.tensor(d['test_x'],dtype=torch.long,device=dev); ye=torch.tensor(d['test_y'],dtype=torch.long,device=dev)
99 oq,orr=net(torch.eye(N,device=dev)[ye],torch.eye(N,device=dev)[xe])
100 metric=float(0.5*((oq.argmax(1)!=xe).float().mean()+(orr.argmax(1)!=ye).float().mean()))
101 if return_net: return metric, net, d, dev
102 return metric
103 except RuntimeError:
104 if dev=='cuda':
105 torch.cuda.empty_cache()
106 os.environ['CUDA_VISIBLE_DEVICES']=''
107 return train_one(seed,lr,lam,return_net)
108 raise
109
110def evaluate(cfg, seeds=SEEDS):
111 vals=[train_one(s,cfg['lr'],cfg.get('lambda',0.0)) for s in seeds]
112 return {'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)}
113
114def main():
115 # Cheap numerical verification before neural training.
116 B,_=basis_matrix(N,EDGES); rng=np.random.default_rng(123); ux=rng.normal(size=N); vy=rng.normal(size=N)
117 a=np.array([ux[x]-vy[y] for x,y in EDGES]); math_max=float(np.max(np.abs(B@a)))
118 cycles=all_simple_cycles(N,EDGES); assert B.shape[0]==len(EDGES)-2*N+1
119 # Baseline sweep includes every lr used by the idea; lambda=0 is standard CE.
120 sweep=[]
121 for lr in LRS: sweep.append({'lr':lr,'lambda':0.0})
122 sweep_res=[]
123 for cfg in sweep:
124 r=evaluate(cfg,tuple(range(4))); sweep_res.append({'cfg':cfg,'mean':r['mean']})
125 best=min(sweep_res,key=lambda z:z['mean'])['cfg']
126 base_full=evaluate(best)
127 base_block={'best_cfg':best,'sweep':sweep_res,'full':base_full}
128 idea_runs=[]
129 for lam in IDEA_LAMBDAS:
130 cfg={'lr':best['lr'],'lambda':lam}; r=evaluate(cfg)
131 idea_runs.append((r,cfg))
132 idea,idea_cfg=min(idea_runs,key=lambda z:z[0]['mean'])
133 # Signature is measured from a trained benchmark model, not toy algebra.
134 metric,net,d,dev=train_one(0,idea_cfg['lr'],idea_cfg['lambda'],True)
135 with torch.no_grad():
136 ex=torch.tensor(EDGES[:,0],device=dev); ey=torch.tensor(EDGES[:,1],device=dev)
137 aq=net.q(torch.eye(N,device=dev)[ey]).log_softmax(1)[range(len(EDGES)),ex]
138 ar=net.r(torch.eye(N,device=dev)[ex]).log_softmax(1)[range(len(EDGES)),ey]
139 avec=(aq-ar).cpu().numpy()
140 basis_res=np.abs(B@avec); held=np.array([abs(c@avec) for c in cycles])
141 sig={'prediction':'basis constraints span all cycle constraints; basis count equals cycle rank',
142 'predicted_cycle_rank':int(B.shape[0]),'observed_exhaustive_cycles':len(cycles),
143 'observed_basis_mean_abs':float(basis_res.mean()),
144 'observed_unseen_cycle_mean_abs':float(held.mean()) if len(held) else 0.0,
145 'observed_unseen_to_basis_ratio':float(held.mean()/(basis_res.mean()+1e-12)) if len(held) else 0.0,
146 'confirmed':bool(B.shape[0]==len(EDGES)-2*N+1 and (not len(held) or held.mean() <= 10*(basis_res.mean()+1e-8)))}
147 rep=make_report('sparse_cycle_compatibility','pair_mlp',base_block,idea,{'custom_track':{'name':'sparse_cycle_compatibility','file':'custom_cycle_track.py','domain':'masked_categorical_compatibility'},'idea_cfg':idea_cfg,'mechanism_signature':sig,'math_check':{'basis_rank':int(B.shape[0]),'expected_rank':int(len(EDGES)-2*N+1),'compatible_max_residual':math_max}})
148 open('bench_report.json','w').write(json.dumps(rep,indent=2))
149 print(json.dumps(rep,indent=2))
150if __name__=='__main__': main()