Mean-field rainbow relation router / rainbow_bench.py

Unverified

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