Submetry-Lifted Relational Alignment / bench_run.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, time
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
8from bench import sweep_baseline, make_report
9from bench.protocol import evaluate
10import graph_track
11
12SEEDS = tuple(range(8))
13SWEEP_SEEDS = tuple(range(4))
14LRS = [0.003, 0.01, 0.03]
15EPOCHS = 18
16BATCH = 64
17
18
19def seed_all(seed):
20 np.random.seed(seed); torch.manual_seed(seed)
21 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
22
23
24def dataset(seed):
25 d = graph_track.get_dataset(seed, 400, 160)
26 return {k: torch.as_tensor(v, dtype=(torch.long if k.startswith('y') else torch.float32))
27 for k, v in d.items() if k in ('xtr','ytr','xte','yte')} | {
28 'task':'classification', 'metric':'err', 'track':'relational_graph_classification'}
29
30
31class GraphNet(nn.Module):
32 def __init__(self):
33 super().__init__()
34 self.node = nn.Sequential(nn.Linear(9, 48), nn.ReLU(), nn.Linear(48, 32), nn.ReLU())
35 self.head = nn.Sequential(nn.Linear(32, 24), nn.ReLU(), nn.Linear(24, 2))
36 def nodes(self, x): return self.node(x)
37 def forward(self, x): return self.head(self.nodes(x).mean(1))
38
39
40def sinkhorn(cost, eps, iters=35):
41 # Uniform marginals, normalized to total mass one.
42 n = cost.shape[-1]
43 eps = torch.as_tensor(eps, device=cost.device, dtype=cost.dtype).clamp_min(1e-5)
44 logk = -cost / eps
45 u = torch.zeros_like(logk[..., :, 0]); v = torch.zeros_like(logk[..., 0, :])
46 logm = -np.log(n)
47 for _ in range(iters):
48 u = logm - torch.logsumexp(logk + v.unsqueeze(-2), dim=-1)
49 v = logm - torch.logsumexp(logk + u.unsqueeze(-1), dim=-2)
50 return torch.exp(logk + u.unsqueeze(-1) + v.unsqueeze(-2))
51
52
53def pair_loss(net, x):
54 # Dense lifted relational cost C[i,k] = sum_j (omega[i,j]-eta[k,j])^2.
55 b, n, _ = x.shape
56 perm = torch.stack([torch.randperm(n, device=x.device) for _ in range(b)])
57 xp = torch.gather(x, 1, perm.unsqueeze(-1).expand(-1, -1, x.shape[-1]))
58 w, wp = x[..., :n], xp[..., :n]
59 rows = (w[:, :, None, :] - wp[:, None, :, :]).pow(2).mean(-1)
60 med = rows.detach().flatten(1).median(1).values.clamp_min(1e-3)
61 P = sinkhorn(rows.detach(), (0.1 * med).view(b, 1, 1), iters=35).detach()
62 h, hp = net.nodes(x), net.nodes(xp)
63 # Node-level lifted signal uses the same coupling; stop-gradient P.
64 align = ((h[:, :, None, :] - hp[:, None, :, :]).pow(2) * P[:, :, :, None]).sum(3).mean()
65 return xp, align, P
66
67
68def train_one(seed, lr, idea):
69 seed_all(seed)
70 d = dataset(seed)
71 net = GraphNet()
72 opt = torch.optim.Adam(net.parameters(), lr=lr)
73 device = 'cuda' if torch.cuda.is_available() else 'cpu'
74 try:
75 net.to(device); xtr, ytr = d['xtr'].to(device), d['ytr'].to(device)
76 for _ in range(EPOCHS):
77 net.train(); order = torch.randperm(len(xtr), device=device)
78 for st in range(0, len(xtr), BATCH):
79 ix = order[st:st+BATCH]; x = xtr[ix]; y = ytr[ix]
80 if idea:
81 xp, al, _ = pair_loss(net, x)
82 loss = F.cross_entropy(net(x), y) + F.cross_entropy(net(xp), y) + 0.08 * al
83 else:
84 # Standard permutation augmentation, same base model/data budget.
85 p = torch.stack([torch.randperm(8, device=device) for _ in range(len(x))])
86 xp = torch.gather(x, 1, p.unsqueeze(-1).expand(-1,-1,9))
87 loss = 0.5 * (F.cross_entropy(net(x), y) + F.cross_entropy(net(xp), y))
88 opt.zero_grad(); loss.backward(); opt.step()
89 net.eval(); xt, yt = d['xte'].to(device), d['yte'].to(device)
90 with torch.no_grad(): metric = float((net(xt).argmax(1) != yt).float().mean())
91 # Trained-system signature: output variance over arbitrary node relabelings.
92 with torch.no_grad():
93 vals=[]
94 for x in xt[:40]:
95 outs=[]
96 for _ in range(6):
97 p=torch.randperm(8,device=device)
98 xx=x[p]
99 outs.append(torch.softmax(net(xx.unsqueeze(0)),1)[0,1])
100 vals.append(torch.stack(outs).var())
101 variance=float(torch.stack(vals).mean())
102 return metric, variance
103 except RuntimeError:
104 # Robust CPU fallback for shared-GPU failures.
105 torch.set_num_threads(2); net = GraphNet(); opt=torch.optim.Adam(net.parameters(),lr=lr)
106 xtr,ytr=d['xtr'],d['ytr']
107 for _ in range(EPOCHS):
108 for st in range(0,len(xtr),BATCH):
109 x=xtr[st:st+BATCH]; y=ytr[st:st+BATCH]
110 if idea:
111 xp,al,_=pair_loss(net,x); loss=F.cross_entropy(net(x),y)+F.cross_entropy(net(xp),y)+.08*al
112 else:
113 p=torch.stack([torch.randperm(8) for _ in range(len(x))]); xp=torch.gather(x,1,p.unsqueeze(-1).expand(-1,-1,9)); loss=.5*(F.cross_entropy(net(x),y)+F.cross_entropy(net(xp),y))
114 opt.zero_grad();loss.backward();opt.step()
115 with torch.no_grad(): metric=float((net(d['xte']).argmax(1)!=d['yte']).float().mean())
116 variance=0.0
117 with torch.no_grad():
118 for x in d['xte'][:40]:
119 q=torch.stack([torch.softmax(net(x[torch.randperm(8)].unsqueeze(0)),1)[0,1] for _ in range(6)])
120 variance += float(q.var())/40
121 return metric, variance
122
123
124def eval_cfg(lr, idea, seeds):
125 vals=[]; sig=[]
126 for s in seeds:
127 m,v=train_one(s,lr,idea); vals.append(m); sig.append(v)
128 return {'per_seed':vals,'mean':float(np.mean(vals)), 'signature_variance':float(np.mean(sig)), 'lr':lr}
129
130
131def main():
132 # Baseline sweep uses the same lr union as the idea settings.
133 base = sweep_baseline(lambda cfg: lambda s: train_one(s, cfg['lr'], False)[0],
134 [{'lr':x} for x in LRS], seeds=SWEEP_SEEDS)
135 # sweep_baseline expects train_fn(seed), while evaluate does exactly that.
136 idea_runs=[eval_cfg(lr,True,SEEDS) for lr in LRS]
137 best=min(idea_runs,key=lambda z:z['mean'])
138 # Re-run the selected baseline on the same eight seeds to measure its
139 # trained-model permutation signature (the protocol sweep stores scalars).
140 bfull = eval_cfg(base['best_cfg']['lr'], False, SEEDS)
141 base['full']['signature_variance'] = bfull['signature_variance']
142 report=make_report('relational_graph_classification','graphnet',base,best,extra={
143 'prediction':'alignment should reduce permutation prediction variance',
144 'predicted_baseline_variance_not_available':False,
145 'observed_baseline_variance':float(base['full'].get('signature_variance',float('nan'))),
146 'observed_idea_variance':best['signature_variance'],
147 'confirmed': bool(best['signature_variance'] < base['full'].get('signature_variance',float('inf'))),
148 'sinkhorn_iterations':35, 'alignment_epsilon':'0.1*median relational row cost'})
149 report['idea_sweep']=idea_runs
150 report['custom_track']={'name':graph_track.META['name'],'file':'graph_track.py','domain':graph_track.META['domain']}
151 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
152 print(json.dumps(report,indent=2))
153
154if __name__=='__main__': main()