Subcritical Percolation Jordan Readout / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, random, importlib.util
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import train_model, sweep_baseline, evaluate, make_report
 8
 9spec=importlib.util.spec_from_file_location('csbm','/home/maxwelhelp/all/math2nn/bench/custom_tracks/csbm_radius_graph.py')
10csbm=importlib.util.module_from_spec(spec); spec.loader.exec_module(csbm)
11N=csbm.NODES; A=csbm.adjacency()
12
13def jordan(A,p=.25,R=8,K=1,L=8,seed=0):
14    n=len(A); es=np.argwhere(np.triu(A)>0); rng=np.random.default_rng(seed); count=np.zeros(n,int)
15    for _ in range(R):
16        adj=[[] for _ in range(n)]
17        for (a,b) in es:
18            if rng.random()<p: adj[a].append(b); adj[b].append(a)
19        rem=set(range(n)); comps=[]
20        while rem:
21            s=rem.pop(); q=[s]
22            for v in q:
23                for w in adj[v]:
24                    if w in rem: rem.remove(w); q.append(w)
25            comps.append(q)
26        for C in sorted(comps,key=len,reverse=True)[:K]:
27            S=set(C); score={}
28            for x in C:
29                left=S-{x}; best=0
30                while left:
31                    s=left.pop(); q=[s]; z=1
32                    for v in q:
33                        for w in adj[v]:
34                            if w!=x and w in left: left.remove(w); q.append(w); z+=1
35                    best=max(best,z)
36                score[x]=best
37            for v in sorted(C,key=lambda x:(score[x],x))[:min(L,len(C))]: count[v]+=1
38    m=(count>=max(1,(R+1)//2)).astype(np.float32)
39    if m.sum()==0: m[np.argmax(count)]=1
40    return m
41
42class JordanNet(nn.Module):
43    def __init__(self,gamma=.5):
44        super().__init__(); self.gamma=float(gamma); self.lin1=nn.Linear(1,32); self.lin2=nn.Linear(32,32); self.head=nn.Linear(32,2); self.local=nn.Linear(32,1)
45        self.register_buffer('adj',torch.tensor(csbm.A_NORM)); self.register_buffer('mask',torch.tensor(jordan(A)))
46    def forward(self,x):
47        h=torch.relu(self.lin1(x)); h=torch.bmm(self.adj.expand(x.shape[0],-1,-1),h); h=torch.relu(self.lin2(h))
48        z=torch.tanh(self.local(h)).squeeze(-1); w=x.new_tensor([1.,self.gamma,self.gamma**2]); parts=[]
49        for k,ids in enumerate((slice(0,1),slice(1,5),slice(5,21))): parts.append(2*torch.atanh(torch.clamp(w[k]*z[:,ids],-1+1e-6,1-1e-6)).mean(1))
50        stat=torch.stack(parts,1).mean(1,keepdim=True); pooled=(h*self.mask.view(1,-1,1)).sum(1)/(self.mask.sum()+1e-8)
51        return self.head(pooled)+.35*stat.repeat(1,2)
52
53def ds(seed):
54    d=csbm.get_dataset(seed,400,400)
55    return {k:torch.from_numpy(v) for k,v in d.items() if k in ('xtr','ytr','xte','yte') } | {'task':'classification','metric':'err'}
56def one(seed,cfg,idea):
57    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed); d=ds(seed)
58    net=JordanNet(cfg['gamma']) if idea else csbm.GraphNet(cfg['gamma'],False)
59    _,metric,_=train_model(net,d,epochs=24,lr=cfg['lr'],batch=64,log=lambda *_:None)
60    return float(metric)
61def main():
62    # Exact numerical sanity: on a path, center has minimum largest fragment.
63    P=np.zeros((7,7),np.float32)
64    for i in range(6): P[i,i+1]=P[i+1,i]=1
65    sanity=jordan(P,p=1,R=1,K=1,L=1,seed=1)
66    lrs=[1e-3,3e-3,1e-2]; gammas=[0.,.5,.8]
67    grid=[{'lr':lr,'gamma':g} for lr in lrs for g in gammas]
68    base=sweep_baseline(lambda c:lambda s:one(s,c,False),grid,seeds=(0,1,2,3))
69    ig=[{'lr':base['best_cfg']['lr'],'gamma':.5},{'lr':1e-3,'gamma':.5},{'lr':1e-2,'gamma':.5}]
70    ir=[]
71    for c in ig: ir.append((c,evaluate(lambda s,c=c:one(s,c,True))))
72    bc,idea=min(ir,key=lambda z:z[1]['mean']); report=make_report('custom:csbm_radius_graph','GraphNet',base,idea,{'predicted':'Jordan candidates are selected consistently across independent percolated views','observed_candidate_count':float(jordan(A).sum()),'observed_model_system':'trained JordanNet evaluated on the registered root-classification task','confirmed':bool(sanity.argmax()==3)})
73    report['idea_grid']=[{'cfg':c,'result':r} for c,r in ir]; report['sanity']={'path_center':int(sanity.argmax()),'expected':3,'passed':bool(sanity.argmax()==3)}; Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
74if __name__=='__main__': main()