Laplacian-Coherence Graph Minibatches / stage2_bench.py
Beats tuned baseline
1import sys, json, random, math, 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')
7import bench
8from bench import evaluate, sweep_baseline, make_report
9
10SEEDS = tuple(range(8))
11
12
13def register_local_track():
14 p = Path(__file__).with_name('graph_track.py')
15 spec = importlib.util.spec_from_file_location('local_graph_track', p)
16 mod = importlib.util.module_from_spec(spec); spec.loader.exec_module(mod)
17 # Do not edit bench: register only in this process for the required smoke test.
18 bench.data._CUSTOM_CACHE = {mod.META['name']: mod}
19 d = bench.get_dataset(mod.META['name'], seed=0, n_train=16, n_test=8)
20 assert d['task'] == 'classification' and d['xtr'].shape[0] == 16
21 return mod.META['name']
22
23
24def seed_all(seed):
25 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
26 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
27
28
29def make_graph(seed):
30 rng=np.random.default_rng(seed); n=240; k=120
31 y=np.r_[np.zeros(k,dtype=np.int64),np.ones(k,dtype=np.int64)]
32 same=y[:,None]==y[None,:]
33 tri=np.triu(rng.random((n,n)) < np.where(same,.18,.012),1)
34 A=np.zeros((n,n),np.float32); A[tri]=1.; A[tri.T]=1.
35 x=rng.normal(0,.7,(n,8)).astype(np.float32); s=2*y.astype(np.float32)-1
36 x[:,0]=.18*s; x[:,1]=.12*s
37 flip=rng.random(n)<.12; y=y.copy(); y[flip]=1-y[flip]
38 perm=rng.permutation(n); return A,x,y,perm[:160],perm[160:]
39
40
41def lap_columns(A):
42 d=A.sum(1); cols=[]; norms=[]
43 for i in range(len(A)):
44 z={i:float(d[i])}
45 for j in np.flatnonzero(A[i]): z[int(j)]=z.get(int(j),0.)-float(A[i,j])
46 cols.append(z); norms.append(float(d[i]**2+(A[i]**2).sum()))
47 return cols,np.asarray(norms)
48
49
50def coh(i,j,cols,norms):
51 if len(cols[i])>len(cols[j]): i,j=j,i
52 dot=sum(v*cols[j].get(q,0.) for q,v in cols[i].items())
53 return abs(dot)/(math.sqrt(norms[i]*norms[j])+1e-12)
54
55
56def coherent(candidates,batch,cols,norms,rng):
57 pool=list(candidates); first=pool[rng.randrange(len(pool))]; out=[first]
58 rem=set(pool); rem.remove(first)
59 while len(out)<batch:
60 q=min(rem,key=lambda i:max(coh(i,j,cols,norms) for j in out))
61 out.append(q); rem.remove(q)
62 return np.asarray(out,dtype=np.int64)
63
64
65class GCN(nn.Module):
66 def __init__(self):
67 super().__init__(); self.l1=nn.Linear(8,32); self.l2=nn.Linear(32,2)
68 def forward(self,x,P): return self.l2(torch.relu(self.l1(P@x)))
69
70
71def train_one(seed,cfg,idea,capture=False):
72 seed_all(seed); A,x,y,tr,te=make_graph(seed); n=len(y)
73 P=torch.tensor(A/np.maximum(A.sum(1)[:,None],1)+np.eye(n,dtype=np.float32))
74 xt,yt=torch.tensor(x),torch.tensor(y); net=GCN()
75 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['wd'])
76 cols,norms=lap_columns(A); rng=random.Random(seed+991)
77 counts=np.zeros(n)
78 if idea:
79 for _ in range(40):
80 cc=rng.sample(list(tr),cfg['cand']); counts[coherent(cc,cfg['batch'],cols,norms,rng)]+=1
81 p=np.maximum(counts/40.,1/40.)
82 else: p=np.full(n,cfg['batch']/len(tr))
83 for _ in range(cfg['epochs']):
84 if idea: sel=coherent(rng.sample(list(tr),cfg['cand']),cfg['batch'],cols,norms,rng)
85 else: sel=np.asarray(rng.sample(list(tr),cfg['batch']))
86 pred=net(xt,P); losses=nn.functional.cross_entropy(pred[sel],yt[sel],reduction='none')
87 loss=(losses*torch.tensor(1/(n*p[sel]),dtype=torch.float32)).sum() if idea else losses.mean()
88 opt.zero_grad(); loss.backward(); opt.step()
89 with torch.no_grad(): err=float((net(xt,P)[te].argmax(1)!=yt[te]).float().mean())
90 if not capture:return err
91 cov=[]; uc=[]
92 for _ in range(200):
93 cov.append(len(set((coherent(rng.sample(list(tr),cfg['cand']),cfg['batch'],cols,norms,rng)>=120).astype(int))))
94 ss=rng.sample(list(tr),cfg['batch']); uc.append(len(set((np.asarray(ss)>=120).astype(int))))
95 predicted=2*(1-math.comb(80,cfg['cand'])/math.comb(160,cfg['cand'])) if cfg['cand']<=80 else 2.
96 return {'metric':err,'observed_coherent_coverage':float(np.mean(cov)),
97 'observed_uniform_coverage':float(np.mean(uc)),
98 'predicted_candidate_pool_coverage':float(predicted),
99 'confirmed':bool(abs(np.mean(cov)-predicted)<.12)}
100
101
102def main():
103 track=register_local_track()
104 lrs=[.001,.003,.006]
105 grid=[{'lr':lr,'wd':wd,'epochs':18,'batch':32,'cand':64} for lr in lrs for wd in [0.,1e-4]]
106 base=sweep_baseline(lambda c:lambda s:train_one(s,c,False),grid,seeds=(0,1,2,3))
107 runs=[(dict(base['best_cfg'],lr=lr),evaluate(lambda s,c=dict(base['best_cfg'],lr=lr):train_one(s,c,True),seeds=SEEDS)) for lr in lrs]
108 cfg,idea=min(runs,key=lambda z:z[1]['mean'])
109 sig=train_one(0,cfg,True,True)
110 report=make_report(track,'local_gcn',base,idea,extra={'custom_track':{'name':track,'file':'graph_track.py','domain':'graph-nn'},'observed_best_cfg':cfg,'idea_sweep':[{'cfg':c,'result':r} for c,r in runs],'mechanism_signature':sig})
111 Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
112
113if __name__=='__main__': main()