Mean-field rainbow relation router / bench_rainbow.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, time
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import train_model, sweep_baseline, make_report
8from bench.protocol import evaluate
9from bench.custom_tracks.relational_graph_classification import get_dataset
10
11SEEDS = tuple(range(8))
12SWEEP_SEEDS = (0, 1, 2, 3)
13
14
15def seed_all(seed):
16 np.random.seed(seed)
17 torch.manual_seed(seed)
18 if torch.cuda.is_available():
19 torch.cuda.manual_seed_all(seed)
20
21
22def sm(x):
23 return torch.softmax(x, dim=-1)
24
25
26class RelationalGNN(nn.Module):
27 """Shared small GNN; only route() differs between baseline and idea."""
28 def __init__(self, mode='baseline', tau=1.0, beta4=0.0, steps=2):
29 super().__init__()
30 self.mode, self.tau, self.beta4, self.steps = mode, float(tau), float(beta4), int(steps)
31 self.node = nn.Sequential(nn.Linear(1, 24), nn.ReLU(), nn.Linear(24, 24))
32 self.edge = nn.Sequential(nn.Linear(2 * 24 + 1, 32), nn.ReLU(), nn.Linear(32, 3))
33 self.rel = nn.ModuleList([nn.Linear(24, 24, bias=False) for _ in range(3)])
34 self.head = nn.Sequential(nn.Linear(24, 24), nn.ReLU(), nn.Linear(24, 2))
35 self.beta = nn.Parameter(torch.zeros(3))
36
37 def unary(self, x):
38 h = self.node(x[..., 8:9])
39 hi, hj = h.unsqueeze(2).expand(-1,-1,8,-1), h.unsqueeze(1).expand(-1,8,-1,-1)
40 ef = torch.cat([hi, hj, x[..., :8].unsqueeze(-1)], dim=-1)
41 return self.edge(ef), h
42
43 def route(self, logits):
44 p = sm(logits / self.tau)
45 if self.mode == 'baseline' or self.beta4 == 0:
46 return p
47 n = p.shape[1]
48 q = p
49 for _ in range(self.steps):
50 r = torch.zeros_like(q)
51 for a in range(3):
52 other = [z for z in range(3) if z != a]
53 r[..., a] = (q[:, :, :, other[0]].unsqueeze(2) * q[:, :, :, other[1]].unsqueeze(1) +
54 q[:, :, :, other[1]].unsqueeze(2) * q[:, :, :, other[0]].unsqueeze(1)).sum(dim=3)
55 score = (2.0 * self.beta.view(1,1,1,3) + (self.beta4 / n) * r) / self.tau
56 q = sm(score)
57 return q
58
59 def forward(self, x):
60 logits, h = self.unary(x)
61 p = self.route(logits)
62 a = x[..., :8]
63 msgs = []
64 for z in range(3):
65 mz = self.rel[z](h)
66 msgs.append(torch.einsum('bij,bjd->bid', a * p[..., z], mz))
67 pooled = h + sum(msgs) / 8.0
68 return self.head(pooled.mean(dim=1))
69
70 @torch.no_grad()
71 def route_stats(self, x, beta4=None):
72 old = self.beta4
73 if beta4 is not None: self.beta4 = float(beta4)
74 logits, _ = self.unary(x)
75 p = self.route(logits)
76 self.beta4 = old
77 n = p.shape[1]
78 tri = []
79 for i in range(n):
80 for j in range(i+1,n):
81 for k in range(j+1,n):
82 v = 0.
83 for a in range(3):
84 o = [z for z in range(3) if z != a]
85 v = v + p[:,i,j,a] * (p[:,i,k,o[0]]*p[:,j,k,o[1]] + p[:,i,k,o[1]]*p[:,j,k,o[0]])
86 tri.append(v)
87 density = torch.stack(tri, 1).mean().item()
88 return {'rainbow_density': density, 'entropy': float((-p*torch.log(p.clamp_min(1e-8))).sum(-1).mean().item()),
89 'mean_motif_logit': float((self.beta4/n * torch.abs(torch.stack(tri,1))).mean().item())}
90
91
92def make_ds(seed, ntr=400, nte=160):
93 d = get_dataset(seed, ntr, nte)
94 for k in ('xtr','ytr','xte','yte'):
95 d[k] = torch.from_numpy(d[k])
96 return d
97
98
99def run_cfg(cfg, seed, mode):
100 seed_all(seed)
101 ds = make_ds(seed)
102 net = RelationalGNN(mode=mode, tau=cfg['tau'], beta4=cfg.get('beta4', 0), steps=2)
103 _, metric, _ = train_model(net, ds, epochs=18, lr=cfg['lr'], batch=64, log=lambda *_: None)
104 return metric
105
106
107def main():
108 grid = [{'lr': lr, 'tau': tau} for lr in (0.0015, 0.003, 0.006) for tau in (0.5, 1.0, 2.0)]
109 baseline = sweep_baseline(lambda c: lambda s: run_cfg(c, s, 'baseline'), grid, seeds=SWEEP_SEEDS)
110 bc = baseline['best_cfg']
111 idea_grid = [dict(bc, beta4=0.75), dict(bc, beta4=1.5), dict(bc, beta4=3.0)]
112 idea_runs = []
113 best_cfg, best_mean = None, float('inf')
114 for c in idea_grid:
115 r = evaluate(lambda s, c=c: run_cfg(c, s, 'idea'), seeds=SEEDS)
116 idea_runs.append({'cfg': c, 'result': r})
117 if r['mean'] < best_mean: best_mean, best_cfg = r['mean'], c
118 idea = next(z['result'] for z in idea_runs if z['cfg'] == best_cfg)
119 sig_rows = []
120 for s in SEEDS:
121 seed_all(s); ds = make_ds(s)
122 net = RelationalGNN(mode='idea', tau=best_cfg['tau'], beta4=best_cfg['beta4'], steps=2)
123 net, _, _ = train_model(net, ds, epochs=18, lr=best_cfg['lr'], batch=64, log=lambda *_: None)
124 net.eval(); x = ds['xte'].to(next(net.parameters()).device)
125 z0, z1 = net.route_stats(x, 0.0), net.route_stats(x, best_cfg['beta4'])
126 sig_rows.append({'seed': s, 'density0': z0['rainbow_density'], 'density_beta4': z1['rainbow_density'],
127 'delta_density': z1['rainbow_density']-z0['rainbow_density'],
128 'entropy_beta4': z1['entropy']})
129 dmean = float(np.mean([z['delta_density'] for z in sig_rows]))
130 signature = {'prediction': 'trained positive beta4 increases rainbow density while beta4=0 is unary',
131 'observed_mean_density_change': dmean, 'per_seed': sig_rows,
132 'beta4_zero_max_change': 0.0,
133 'confirmed': bool(dmean > 0.001)}
134 rep = make_report('relational_graph_classification', 'local_relational_gnn', baseline, idea,
135 {'mechanism_signature': signature, 'idea_sweep': idea_runs,
136 'protocol_note': 'Built-in relational custom track used; no structurally matching built-in track exists.'})
137 rep['wall_clock_note'] = 'canonical train_model; equal 18 epochs, batch 64; timing not primary metric'
138 with open('bench_report.json', 'w') as f: json.dump(rep, f, indent=2)
139 print(json.dumps(rep, indent=2))
140
141if __name__ == '__main__': main()