Mean-field rainbow relation router / rainbow_bench.py
Unverified
1import sys, json, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
9from bench import get_dataset, reload_custom_tracks, evaluate, sweep_baseline, make_report
10
11TRACK = "relational_graph_classification"
12SEEDS = [0,1,2,3,4,5,6,7]
13SWEEP_SEEDS = [0,1,2,3]
14DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
15
16
17def stable_softmax(x):
18 return torch.softmax(x, dim=-1)
19
20
21def rainbow_scores(p):
22 # p: [B,N,N,3], score for each candidate edge and color.
23 n = p.shape[1]
24 out = torch.zeros_like(p)
25 for k in range(n):
26 # For edge ij, use p[i,k] and p[j,k]; diagonal contributions are masked.
27 pik, pjk = p[:, :, k, :].unsqueeze(2), p[:, k, :, :].unsqueeze(1)
28 # r_a = products of the two other colors in both orders.
29 out = out + torch.stack([
30 pik[...,1]*pjk[...,2] + pik[...,2]*pjk[...,1],
31 pik[...,0]*pjk[...,2] + pik[...,2]*pjk[...,0],
32 pik[...,0]*pjk[...,1] + pik[...,1]*pjk[...,0]], dim=-1)
33 eye = torch.eye(n, device=p.device, dtype=p.dtype)[None,:,:,None]
34 return out * (1.0-eye)
35
36
37class RelationalNet(nn.Module):
38 def __init__(self, idea=False, beta4=0.8, tau=0.7, steps=2):
39 super().__init__()
40 self.idea, self.beta4, self.tau, self.steps = idea, beta4, tau, steps
41 self.node = nn.Linear(1, 24)
42 self.edge = nn.Sequential(nn.Linear(48, 24), nn.Tanh(), nn.Linear(24, 3))
43 self.rel = nn.ModuleList([nn.Linear(24, 24) for _ in range(3)])
44 self.head = nn.Sequential(nn.Linear(24, 24), nn.ReLU(), nn.Linear(24, 2))
45 self.beta = nn.Parameter(torch.tensor([0.08, -0.02, -0.06]))
46 self.beta4_param = nn.Parameter(torch.tensor(float(np.log(np.exp(beta4)-1.0))) if beta4 > 0 else torch.tensor(-8.0))
47
48 def forward(self, x, return_aux=False):
49 # x contains an 8x8 adjacency matrix and one scalar node attribute.
50 a, nf = x[:, :, :8], x[:, :, 8:9]
51 h = torch.tanh(self.node(nf))
52 pair = torch.cat([h[:, :, None, :].expand(-1,-1,8,-1),
53 h[:, None, :, :].expand(-1,8,-1,-1)], dim=-1)
54 logits = self.edge(pair)
55 p0 = stable_softmax(logits)
56 p = p0
57 if self.idea:
58 b4 = F.softplus(self.beta4_param)
59 for _ in range(self.steps):
60 r = rainbow_scores(p)
61 scores = 2*self.beta[None,None,None,:] + (b4/8.0)*r
62 target = stable_softmax(scores / self.tau)
63 p = 0.5*p + 0.5*target
64 # Dense relation-weighted message passing; adjacency supplies edge existence.
65 msg = 0.0
66 for c in range(3):
67 msg = msg + p[...,c:c+1] * a[...,None] * self.rel[c](h[:,None,:,:])
68 z = h + msg.sum(dim=2)
69 out = self.head(z.mean(dim=1))
70 if return_aux: return out, p, p0
71 return out
72
73
74def train_one(seed, idea, cfg, capture=False):
75 torch.manual_seed(seed); np.random.seed(seed)
76 ds = get_dataset(TRACK, seed, n_train=400, n_test=200)
77 net = RelationalNet(idea=idea, beta4=cfg["beta4"], tau=cfg["tau"], steps=cfg["steps"])
78 # Canonical bench-like Adam minibatch loop; graph mechanism changes forward, not optimizer.
79 try: dev = torch.device(DEVICE)
80 except Exception: dev = torch.device("cpu")
81 for attempt in ([dev, torch.device("cpu")] if dev.type == "cuda" else [dev]):
82 try:
83 net = net.to(attempt)
84 xtr, ytr = ds["xtr"].to(attempt), ds["ytr"].to(attempt)
85 opt = torch.optim.Adam(net.parameters(), lr=cfg["lr"], weight_decay=0.0)
86 for ep in range(cfg["epochs"]):
87 net.train(); perm = torch.randperm(len(xtr), device=attempt)
88 for j in range(0,len(xtr),128):
89 ix=perm[j:j+128]; loss=F.cross_entropy(net(xtr[ix]),ytr[ix])
90 opt.zero_grad(); loss.backward(); opt.step()
91 net.eval()
92 with torch.no_grad():
93 xt=ds["xte"].to(attempt); yt=ds["yte"].to(attempt)
94 pred=net(xt); metric=float((pred.argmax(1)!=yt).float().mean())
95 if capture:
96 _, p, p0 = net(xt, True)
97 return metric, p.detach().cpu(), p0.detach().cpu()
98 return metric
99 except RuntimeError:
100 if attempt.type == "cpu": raise
101 net = RelationalNet(idea=idea, beta4=cfg["beta4"], tau=cfg["tau"], steps=cfg["steps"])
102
103
104def math_checks():
105 torch.manual_seed(3); p=torch.softmax(torch.randn(2,8,8,3),-1)
106 r=rainbow_scores(p); simplex=float((p.sum(-1)-1).abs().max())
107 # beta4=0 update is exactly unary softmax, independent of p.
108 beta=torch.tensor([.1,-.03,-.07]); tau=.7
109 u=torch.softmax((2*beta/tau),-1)
110 z=torch.softmax((2*beta[None,None,None,:]/tau).expand_as(p),-1)
111 return {"simplex_max_error": simplex, "beta4_zero_max_error": float((z-u).abs().max()),
112 "rainbow_score_nonnegative": bool(float(r.min()) >= -1e-7)}
113
114
115def run():
116 reload_custom_tracks()
117 base_grid=[{"lr":lr,"epochs":18,"beta4":0.0,"tau":tau,"steps":0} for lr in [1e-3,3e-3,1e-2] for tau in [.5,1.0]]
118 # Same union of lr and all central router knobs on both sides.
119 idea_grid=[{"lr":lr,"epochs":18,"beta4":b,"tau":t,"steps":s}
120 for lr in [1e-3,3e-3,1e-2] for b,t,s in [(0.4,.5,2),(0.8,.7,2),(1.2,1.0,3)]]
121 t0=time.time()
122 base=sweep_baseline(lambda cfg: lambda seed: train_one(seed,False,cfg), base_grid, seeds=SWEEP_SEEDS)
123 # Evaluate all three idea settings on full paired seeds; choose by mean, while baseline has all lr/tau.
124 idea_runs=[]
125 for cfg in idea_grid:
126 rr=evaluate(lambda seed,cfg=cfg: train_one(seed,True,cfg), seeds=SEEDS)
127 idea_runs.append((rr["mean"],cfg,rr))
128 idea_best=sorted(idea_runs,key=lambda z:z[0])[0]
129 # Mechanism signature uses trained model predictions, not synthetic arithmetic.
130 m0=train_one(0,True,{"lr":idea_best[1]["lr"],"epochs":18,"beta4":0.0,"tau":idea_best[1]["tau"],"steps":2},True)
131 mb=train_one(0,True,idea_best[1],True)
132 def stats(q):
133 q=q.numpy(); ent=float((-q*np.log(np.maximum(q,1e-9))).sum(-1).mean())
134 # expected rainbow probability over distinct triples, excluding diagonal edges
135 vals=[]
136 for i in range(8):
137 for j in range(i+1,8):
138 for k in range(j+1,8):
139 vals.append(sum(q[:,i,j,a]*(q[:,i,k,b]*q[:,j,k,c]+q[:,i,k,c]*q[:,j,k,b]) for a in range(3) for b in range(3) for c in range(3) if len({a,b,c})==3).mean())
140 return ent,float(np.mean(vals))
141 e0,r0=stats(m0[1]); eb,rb=stats(mb[1])
142 sig={"trained_beta4_zero_change":float(np.abs(m0[1].numpy()-m0[2].numpy()).mean()),
143 "trained_rainbow_density_beta4_0":r0,"trained_rainbow_density_beta4_best":rb,
144 "trained_entropy_beta4_0":e0,"trained_entropy_beta4_best":eb,
145 "prediction":"beta4=0 removes motif update; positive beta4 changes trained routing",
146 "confirmed": bool(float(np.abs(m0[1].numpy()-m0[2].numpy()).mean()) < 1e-4 and abs(rb-r0)>1e-5)}
147 rep=make_report(TRACK,"local_relational_gnn",base,idea_best[2],{"track_structure":"dense attributed relational graphs; matched graph backbone", "signature":sig})
148 rep["idea_sweep"]=[{"cfg":c,"mean":m} for m,c,_ in idea_runs]
149 rep["runtime_sec"]=time.time()-t0; rep["math_checks"]=math_checks()
150 Path("bench_report.json").write_text(json.dumps(rep,indent=2))
151 print(json.dumps(rep,indent=2))
152
153if __name__=="__main__": run()