GECC-Gated Loop-Aware Message Passing / gecc_demo.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 2043
8
9def seed_all(seed=SEED):
10 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
12
13
14def ego_intersections(n, edges):
15 """Exact n=1 closed ego sets S[u]={u and its neighbors}; undirected edges."""
16 nbr = [set([i]) for i in range(n)]
17 for u,v in edges:
18 nbr[u].add(v); nbr[v].add(u)
19 stats = {}
20 for u,v in edges:
21 I = nbr[u] & nbr[v]
22 # Formula in prompt, with epsilon only in denominator.
23 C = 2.0*len(I)/(len(nbr[u])+len(nbr[v])+1e-12)
24 r = len(I)/(len(nbr[u])*len(nbr[v])+1e-12)
25 stats[(u,v)] = (sorted(I), C, r, len(nbr[u]), len(nbr[v]))
26 stats[(v,u)] = stats[(u,v)]
27 return nbr, stats
28
29
30def math_verification():
31 # Prediction 1: for a regular edge, adding t common neighbors gives
32 # C=(2(t+2))/(2(k+1))=(t+2)/(k+1), hence slope 1/(k+1).
33 k=6
34 rows=[]
35 for t in range(0, k-1):
36 # endpoints have k-1 other neighbors, t shared among them, plus edge.
37 # Build only sets directly to avoid irrelevant graph edges.
38 su=set(range(k+1)); sv={k+1}; sv.update(range(0,t+1)); sv.add(k+2)
39 # normalize desired sizes: endpoint sets each k+1 and intersection is edge endpoints + t
40 su={0}|set(range(2,k+1)); sv={1}|set(range(2, t+2))|set(range(k+2, 2*k-t+1))
41 # u=0, v=1; force both sizes k+1 and common {0,1,2..t+1}
42 # easier use abstract sets with exact sizes
43 su={0,1}|set(range(2,t+2))|set(range(100,100+(k-1-t)))
44 sv={0,1}|set(range(2,t+2))|set(range(200,200+(k-1-t)))
45 C=2*len(su&sv)/(len(su)+len(sv)+1e-12)
46 pred=(t+2)/(k+1)
47 rows.append((t,C,pred))
48 max_err=max(abs(x-y) for _,x,y in rows)
49 slope=np.polyfit([t for t,_,_ in rows],[c for _,c,_ in rows],1)[0]
50 slope_pred=1/(k+1)
51
52 # Prediction 2: alpha=0.5 at C*=-(a0+a2*r)/a1 and is increasing in C.
53 a0,a1,a2=-1.2,4.0,0.7
54 r=0.08
55 cstar=-(a0+a2*r)/a1
56 Cs=np.linspace(0,1,101)
57 alphas=1/(1+np.exp(-(a0+a1*Cs+a2*r)))
58 crossing=float(Cs[np.argmin(abs(alphas-.5))])
59 monotone=bool(np.all(np.diff(alphas)>0))
60
61 # Prediction 3: for fixed ordinary o and overlap q, update is affine in alpha,
62 # with endpoint values o and q and slope q-o.
63 o=np.array([1.3,-0.4,0.7]); q=np.array([-0.2,0.8,0.1])
64 aa=np.linspace(0,1,11)
65 updates=np.array([(1-x)*o+x*q for x in aa])
66 fit=np.polyfit(aa, updates[:,0], 1)[0]
67 predicted_slope=float(q[0]-o[0])
68 return {
69 'closure_sweep': {'k':k,'rows':rows,'max_abs_error':float(max_err),
70 'observed_slope':float(slope),'predicted_slope':float(slope_pred)},
71 'gate_sweep': {'a0':a0,'a1':a1,'a2':a2,'r':r,'predicted_C_at_alpha_half':float(cstar),
72 'observed_C_at_alpha_half':crossing,'monotone_in_C':monotone,
73 'alpha_at_0':float(alphas[0]),'alpha_at_1':float(alphas[-1])},
74 'update_sweep': {'predicted_slope_coordinate0':predicted_slope,'observed_slope_coordinate0':float(fit),
75 'endpoint_error':float(np.max(np.abs(updates[0]-o))+np.max(np.abs(updates[-1]-q)))}
76 }
77
78
79def make_graph(n=48, p_in=.30, p_out=.045, triangle_bias=0.55):
80 y=np.array([i < n//2 for i in range(n)],dtype=np.int64)
81 edges=[]
82 rng=np.random.RandomState(SEED+int(1000*p_in))
83 for i in range(n):
84 for j in range(i+1,n):
85 p=p_in if y[i]==y[j] else p_out
86 if rng.rand()<p: edges.append((i,j))
87 # Add triangles within each class, making high closure but modest graph size.
88 for i in range(0,n,3):
89 group=list(range(i,min(i+3,n)))
90 for a in range(len(group)):
91 for b in range(a+1,len(group)):
92 e=(group[a],group[b])
93 if e not in edges: edges.append(e)
94 return edges,y
95
96
97def features(n,y,edges):
98 rng=np.random.RandomState(99)
99 x=rng.randn(n,8).astype(np.float32)*0.9
100 x[:,0] += (2*y-1)*0.35
101 return torch.tensor(x), torch.tensor(y,dtype=torch.long)
102
103class GeccLayer(nn.Module):
104 def __init__(self, d):
105 super().__init__(); self.w0=nn.Linear(d,d,bias=False); self.w1=nn.Linear(d,d,bias=False); self.w2=nn.Linear(d,d,bias=False)
106 self.a=nn.Parameter(torch.tensor([-1.,2.,.2]))
107 def forward(self,h, edges, stats, force_gate=None):
108 out=self.w0(h); ordinary=torch.zeros_like(h); corr=torch.zeros_like(h)
109 for u,v in edges:
110 for dst,src in ((u,v),(v,u)):
111 I,C,r,_,_=stats[(dst,src)]
112 alpha=torch.sigmoid(self.a[0]+self.a[1]*C+self.a[2]*r)
113 if force_gate is not None: alpha=torch.tensor(force_gate,device=h.device)
114 ordinary[dst]+=self.w1(h[src])
115 q=h[I].mean(0) if I else torch.zeros(h.shape[1],device=h.device)
116 corr[dst]+=self.w2(q)
117 out[dst] += (1-alpha)*self.w1(h[src]) + alpha*self.w2(q) - self.w0(h[dst])*0
118 return torch.relu(out)
119
120class GCN(nn.Module):
121 def __init__(self,d=8,hid=16):
122 super().__init__(); self.l1=nn.Linear(d,hid); self.l2=nn.Linear(hid,2)
123 def forward(self,x,edges):
124 n=x.shape[0]; z=torch.zeros_like(x); deg=torch.zeros(n,device=x.device)
125 for u,v in edges: z[u]+=x[v]; z[v]+=x[u]; deg[u]+=1; deg[v]+=1
126 z=z/(deg[:,None]+1); return self.l2(torch.relu(self.l1(x+z)))
127
128class GECC(nn.Module):
129 def __init__(self,d=8,hid=16):
130 super().__init__(); self.inp=nn.Linear(d,hid); self.msg=GeccLayer(hid); self.out=nn.Linear(hid,2)
131 def forward(self,x,edges,stats): return self.out(self.msg(torch.relu(self.inp(x)),edges,stats))
132
133
134def train_one(model,x,y,edges,stats,train_idx,steps=60):
135 opt=torch.optim.Adam(model.parameters(),lr=.025,weight_decay=2e-3)
136 for _ in range(steps):
137 opt.zero_grad(); logits=model(x,edges,stats) if isinstance(model,GECC) else model(x,edges)
138 loss=nn.functional.cross_entropy(logits[train_idx],y[train_idx]); loss.backward(); opt.step()
139 with torch.no_grad():
140 logits=model(x,edges,stats) if isinstance(model,GECC) else model(x,edges)
141 pred=logits.argmax(1); return float((pred!=y).float().mean()), float((pred[~train_idx]!=y[~train_idx]).float().mean())
142
143
144def mini_experiment():
145 seed_all(); edges,y_np=make_graph(); n=len(y_np); _,stats=ego_intersections(n,edges); x,y=features(n,y_np,edges)
146 train_idx=torch.zeros(n,dtype=torch.bool); train_idx[:n//2:2]=True; train_idx[n//2::2]=True
147 base=train_one(GCN(),x,y,edges,stats,train_idx); idea=train_one(GECC(),x,y,edges,stats,train_idx)
148 Cs=[v[1] for v in stats.values()]; return {'n':n,'edges':len(edges),'mean_C_directed':float(np.mean(Cs)), 'train_nodes':int(train_idx.sum()),'baseline_error_all_and_test':base,'idea_error_all_and_test':idea}
149
150if __name__=='__main__':
151 seed_all(); result={'math_verification':math_verification(),'mini_experiment':mini_experiment()}
152 Path('results.json').write_text(json.dumps(result,indent=2))
153 print(json.dumps(result,indent=2))