Mean-field rainbow relation router / bench_rainbow.py

Mechanism confirmed, baseline not beaten

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