Connectivity-Preserving Wedge Token Pooling / graph_wedge_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5ROOT='/home/maxwelhelp/all/math2nn'
 6if ROOT not in sys.path: sys.path.insert(0,ROOT)
 7from bench import train_model, make_report, sweep_baseline
 8from bench.protocol import DEFAULT_SEEDS, SWEEP_SEEDS
 9META={'name':'wedge_graph_regression','domain':'graph-nn','description':'Graph-level regression from node signals on connected cycle-plus-chord graphs; adaptive shortest-path wedge pooling preserves connected regions.'}
10N,D,M=18,4,6
11ADJ=[set() for _ in range(N)]
12for i in range(N):
13    ADJ[i].update(((i-1)%N,(i+1)%N))
14for a,b in [(0,9),(3,12),(6,15)]: ADJ[a].add(b); ADJ[b].add(a)
15ADJ=[sorted(s) for s in ADJ]
16
17def get_dataset(seed,n_train=400,n_test=200):
18    rng=np.random.default_rng(seed); n=n_train+n_test; phase=rng.uniform(0,2*np.pi,n); amp=rng.uniform(.5,1.5,n); x=np.zeros((n,N,D),np.float32); idx=np.arange(N)
19    for k in range(n):
20        x[k,:,0]=amp[k]*np.sin(2*np.pi*idx/N+phase[k])+.08*rng.normal(size=N)
21        x[k,:,1]=np.where(idx<9,rng.uniform(-1,1),rng.uniform(-1,1))+.08*rng.normal(size=N)
22        x[k,:,2]=np.cos(4*np.pi*idx/N+phase[k])+.08*rng.normal(size=N); x[k,:,3]=rng.normal(size=N)
23    edges=[(a,b) for a in range(N) for b in ADJ[a] if b>a]; e=np.zeros(n)
24    for a,b in edges: e+=(x[:,a,0]-x[:,b,0])**2+.5*(x[:,a,1]-x[:,b,1])**2
25    y=(e/len(edges)+.25*x[:,:,0].mean(1)+.1*x[:,:,2].mean(1)).astype(np.float32)[:,None]
26    return {'xtr':torch.from_numpy(x[:n_train]),'ytr':torch.from_numpy(y[:n_train]),'xte':torch.from_numpy(x[n_train:]),'yte':torch.from_numpy(y[n_train:]),'task':'regression','metric':'mse','input_shape':(M,D),'out_dim':1}
27
28def bfs(region):
29    R=set(region); allD={}
30    for s in R:
31        ds={s:0}; q=[s]
32        for v in q:
33            for z in ADJ[v]:
34                if z in R and z not in ds: ds[z]=ds[v]+1; q.append(z)
35        allD[s]=ds
36    return allD
37
38def wedge_regions(X):
39    regions=[list(range(N))]
40    while len(regions)<M:
41        best=None
42        for ri,R in enumerate(regions):
43            if len(R)<2: continue
44            ds=bfs(R); mu=X[R].mean(0); base=((X[R]-mu)**2).sum()
45            for u in R:
46                for w in R:
47                    if u>=w: continue
48                    A=[v for v in R if ds[u][v]<=ds[w][v]]; B=[v for v in R if ds[u][v]>ds[w][v]]
49                    if not A or not B: continue
50                    gain=base-((X[A]-X[A].mean(0))**2).sum()-((X[B]-X[B].mean(0))**2).sum()
51                    if best is None or gain>best[0]: best=(gain,ri,A,B)
52        if best is None or best[0]<=1e-10: break
53        _,ri,A,B=best; regions.pop(ri); regions.extend((A,B))
54    return sorted(regions,key=min)
55
56def pool_array(x,mode,seed):
57    out=np.empty((len(x),M,D),np.float32)
58    if mode=='random':
59        rng=np.random.default_rng(seed); labels=np.repeat(np.arange(M),int(np.ceil(N/M)))[:N]; rng.shuffle(labels); regs=[np.flatnonzero(labels==i) for i in range(M)]
60        for k in range(len(x)):
61            for i,r in enumerate(regs): out[k,i]=x[k,r].mean(0)
62    else:
63        for k in range(len(x)):
64            regs=wedge_regions(x[k])
65            for i,r in enumerate(regs): out[k,i]=x[k,r].mean(0)
66            if len(regs)<M: out[k,len(regs):]=out[k,len(regs)-1]
67    return torch.from_numpy(out)
68
69class PoolNet(nn.Module):
70    def __init__(self):
71        super().__init__(); self.net=nn.Sequential(nn.Linear(M*D,64),nn.ReLU(),nn.Linear(64,1))
72    def forward(self,x): return self.net(x.reshape(x.shape[0],-1))
73
74def train_one(mode,lr,seed,epochs=18):
75    torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset(seed,400,200)
76    ds['xtr']=pool_array(ds['xtr'].numpy(),mode,seed); ds['xte']=pool_array(ds['xte'].numpy(),mode,seed+10000)
77    net,metric,_=train_model(PoolNet(),ds,epochs=epochs,lr=lr,batch=128,log=lambda *a:None)
78    net=net.cpu()
79    with torch.no_grad(): pred=net(ds['xte']).numpy().ravel(); obs=ds['yte'].numpy().ravel()
80    return float(metric),{'pred_std':float(pred.std()),'obs_std':float(obs.std()),'pred_obs_corr':float(np.corrcoef(pred,obs)[0,1])}
81
82def eval_cfg(mode,cfg,seeds=DEFAULT_SEEDS):
83    scores=[]; sig=[]
84    for s in seeds:
85        v,z=train_one(mode,float(cfg['lr']),int(s),int(cfg.get('epochs',18))); scores.append(v); sig.append(z)
86    return {'mean':float(np.mean(scores)),'std':float(np.std(scores)),'per_seed':scores,'n':len(scores),'signature':sig}
87
88def main():
89    grid=[{'lr':1e-3,'epochs':18},{'lr':3e-3,'epochs':18},{'lr':1e-2,'epochs':18}]
90    base=sweep_baseline(lambda c: (lambda s: train_one('random',c['lr'],s,c['epochs'])[0]),grid,seeds=SWEEP_SEEDS)
91    # Evaluate the same union of configs for the idea, then select its best.
92    idea_grid={str(c['lr']):eval_cfg('wedge',c) for c in grid}
93    best=min(idea_grid,key=lambda k:idea_grid[k]['mean']); idea=idea_grid[best]
94    # Mechanism signature is measured from trained models, not an identity.
95    sig={'predicted':'graph-aware connected pooling should retain graph-smooth signal (higher prediction/observation correlation than random pooling)','observed':{'baseline_best_corr_mean':float(np.mean([train_one('random',base['best_cfg']['lr'],s,base['best_cfg']['epochs'])[1]['pred_obs_corr'] for s in DEFAULT_SEEDS])),'idea_corr_mean':float(np.mean([z['pred_obs_corr'] for z in idea['signature']]))},'confirmed':False}
96    report=make_report('wedge_graph_regression','mlp_tiny',base,idea,{'mechanism_signature':sig,'custom_track':{'name':META['name'],'file':'graph_wedge_bench.py','domain':META['domain']},'idea_grid':idea_grid})
97    print(json.dumps(report,indent=2))
98if __name__=='__main__': main()