Følner-Gated Message Passing / run_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7
  8SEED = 2062
  9random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 10try:
 11    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 12except Exception:
 13    device = 'cpu'
 14
 15
 16def cycle(n):
 17    g=[set() for _ in range(n)]
 18    for i in range(n):
 19        for j in ((i-1)%n,(i+1)%n): g[i].add(j)
 20    return g
 21
 22def grid(side):
 23    n=side*side; g=[set() for _ in range(n)]
 24    for r in range(side):
 25      for c in range(side):
 26        i=r*side+c
 27        for rr,cc in ((r-1,c),(r+1,c),(r,c-1),(r,c+1)):
 28          if 0<=rr<side and 0<=cc<side: g[i].add(rr*side+cc)
 29    return g
 30
 31def random_regular(n, d=3):
 32    # deterministic retry via networkx if available; fallback ring chords
 33    try:
 34      import networkx as nx
 35      G=nx.random_regular_graph(d,n,seed=SEED+n)
 36      return [set(G.neighbors(i)) for i in range(n)]
 37    except Exception:
 38      g=[set() for _ in range(n)]
 39      for i in range(n):
 40        for k in range(1,d//2+1): g[i].update(((i-k)%n,(i+k)%n))
 41      return g
 42
 43def tree(d, depth):
 44    g=[]; levels=[[0]]; g.append(set())
 45    next_id=1
 46    for dep in range(depth):
 47      nxt=[]
 48      for p in levels[-1]:
 49        count=d if dep==0 else d-1
 50        for _ in range(count):
 51          x=next_id; next_id+=1
 52          while len(g)<=x: g.append(set())
 53          g[p].add(x); g[x].add(p); nxt.append(x)
 54      levels.append(nxt)
 55    return g, levels
 56
 57def expand(g, Fset):
 58    out=set(Fset)
 59    for x in Fset: out.update(g[x])
 60    return out
 61
 62def ratio(g,Fset): return len(expand(g,Fset))/max(1,len(Fset))
 63
 64def graph_matrix(g):
 65    n=len(g); A=torch.zeros((n,n),dtype=torch.float32)
 66    for i,ns in enumerate(g):
 67      A[i,i]=1
 68      for j in ns: A[i,j]=1
 69    A=A/A.sum(1,keepdim=True).clamp_min(1)
 70    return A
 71
 72class GNN(nn.Module):
 73    def __init__(self,din,hidden,layers, gated=False, delta=.35):
 74      super().__init__(); self.gated=gated; self.delta=delta
 75      self.inp=nn.Linear(din,hidden); self.ws=nn.ModuleList([nn.Linear(hidden,hidden) for _ in range(layers)])
 76      self.out=nn.Linear(hidden,2)
 77    def forward(self,x,A,rbar):
 78      h=F.relu(self.inp(x)); gs=[]
 79      for w in self.ws:
 80        msg=A@h
 81        if self.gated:
 82          g=torch.sigmoid(8*((1+self.delta)-rbar))
 83          # pooled skip is a stable global context, weighted by non-message fraction
 84          pooled=h.mean(0,keepdim=True).expand_as(h)
 85          h=h + g*w(msg) + (1-g)*0.25*w(pooled)
 86          gs.append(float(g.detach().cpu()))
 87        else:
 88          h=h+w(msg); gs.append(1.0)
 89        h=F.relu(h)
 90      return self.out(h),gs
 91
 92def mechanism_checks():
 93    # Prediction 1: closed one-hop interval on a cycle has r=1+2/m.
 94    cyc=cycle(200); cycle_rows=[]
 95    for m in [2,4,8,16,32,64,100]:
 96      Fset=set(range(m)); obs=ratio(cyc,Fset); pred=1+2/m
 97      cycle_rows.append({'m':m,'predicted':pred,'observed':obs,'abs_error':abs(pred-obs)})
 98    # Prediction 2: d=3 tree balls have ratios tending to 2; exact ball sweep.
 99    tr,levels=tree(3,7); ball=set()
100    tree_rows=[]
101    for k in range(0,6):
102      ball.update(levels[k]); obs=ratio(tr,ball)
103      pred=(1+3*(2**(k+1)-1)/2 if False else None)
104      # exact finite tree ratio from current ball and next level
105      exact=(len(ball)+len(levels[k+1]) if k+1<len(levels) else len(expand(tr,ball)))/len(ball)
106      tree_rows.append({'radius':k,'size':len(ball),'predicted_exact':exact,'observed':obs})
107    # Prediction 3: sigmoid gate boundary is r=1+delta, independent of slope.
108    delta=.35; boundary_rows=[]
109    for slope in [2,4,8,16]:
110      rs=np.linspace(1.0,2.0,1001); gs=1/(1+np.exp(-slope*((1+delta)-rs)))
111      boundary_rows.append({'slope':slope,'predicted_r':1+delta,'observed_r_at_g=.5':float(rs[np.argmin(abs(gs-.5))])})
112    # Persistence prediction: EMA crosses threshold iff sustained r exceeds it.
113    persistence=[]
114    for r in [1.1,1.3,1.6,2.0]:
115      beta=.8; rb=1.; crossing=None
116      for t in range(1,31):
117        rb=beta*rb+(1-beta)*r
118        if crossing is None and rb>1+delta: crossing=t
119      predicted = (math.log((1+delta-1)/(r-1))/math.log(beta) if r>1+delta else None)
120      persistence.append({'r':r,'cross_step_observed':crossing,'cross_step_continuous_prediction':None if predicted is None else predicted})
121    return {'cycle_boundary_prediction':cycle_rows,'tree_expansion_prediction':tree_rows,'gate_boundary_prediction':boundary_rows,'ema_persistence_prediction':persistence}
122
123def train_case(name,g,epochs=60):
124    n=len(g); A=graph_matrix(g).to(device)
125    # fixed random features; label is a local structural signal (degree parity plus index hash)
126    deg=np.array([len(x) for x in g]); x=np.random.default_rng(SEED+n).normal(size=(n,8)).astype('float32')
127    x[:,0]=deg/ max(1,deg.max())
128    y=torch.tensor(((deg + np.arange(n))%2),dtype=torch.long,device=device)
129    perm=np.random.default_rng(SEED).permutation(n); cut=max(2,int(.7*n)); tridx=torch.tensor(perm[:cut],device=device); vaidx=torch.tensor(perm[cut:],device=device)
130    # sampled seed frontier and a stable EMA estimate, as the controller would see it online
131    seeds=set(perm[:min(16,n)]); Fset=seeds; rb=1.; beta=.7; rs=[]; frontiers=[]
132    for _ in range(6):
133      Fset=expand(g,Fset); r=len(Fset)/max(1,len(seeds)); rb=beta*rb+(1-beta)*r; rs.append(rb); frontiers.append(len(Fset))
134    rbar=torch.tensor(float(rb),device=device)
135    result={}
136    for gated in [False,True]:
137      torch.manual_seed(SEED+int(gated)+n)
138      model=GNN(8,24,6,gated=gated).to(device); opt=torch.optim.Adam(model.parameters(),lr=.025)
139      t0=time.time()
140      for _ in range(epochs):
141        opt.zero_grad(); logits,_=model(torch.tensor(x,device=device),A,rbar)
142        loss=F.cross_entropy(logits[tridx],y[tridx]); loss.backward(); opt.step()
143      with torch.no_grad():
144        logits,gs=model(torch.tensor(x,device=device),A,rbar); acc=(logits[vaidx].argmax(1)==y[vaidx]).float().mean().item()
145        h=logits[vaidx]; cos=F.cosine_similarity(h[:,None,:],h[None,:,:],dim=-1).mean().item() if len(vaidx)>1 else 0
146      result['gated' if gated else 'baseline']={'val_accuracy':acc,'mean_pairwise_cosine':cos,'seconds':time.time()-t0,'gate_mean':float(np.mean(gs))}
147    result.update({'rbar':float(rb),'frontier_peak':max(frontiers),'raw_ratios':[float(x) for x in rs]})
148    return result
149
150def main():
151    checks=mechanism_checks(); cases={}
152    graphs={'cycle':cycle(100),'grid':grid(10),'random3':random_regular(100,3)}
153    tr,_=tree(3,5); graphs['tree3']=tr
154    for name,g in graphs.items():
155      try: cases[name]=train_case(name,g)
156      except Exception as e: cases[name]={'error':repr(e)}
157    out={'seed':SEED,'device':device,'checks':checks,'experiments':cases}
158    Path('results.json').write_text(json.dumps(out,indent=2))
159    print(json.dumps(out,indent=2))
160if __name__=='__main__': main()