GECC-Gated Loop-Aware Message Passing / gecc_demo.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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))