Fundamental-Cycle Compatibility Basis / run_bench.py

Mechanism confirmed, baseline not beaten

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