import os, sys, json, time import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import sweep_baseline, make_report from bench.protocol import evaluate import graph_track SEEDS = tuple(range(8)) SWEEP_SEEDS = tuple(range(4)) LRS = [0.003, 0.01, 0.03] EPOCHS = 18 BATCH = 64 def seed_all(seed): np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def dataset(seed): d = graph_track.get_dataset(seed, 400, 160) return {k: torch.as_tensor(v, dtype=(torch.long if k.startswith('y') else torch.float32)) for k, v in d.items() if k in ('xtr','ytr','xte','yte')} | { 'task':'classification', 'metric':'err', 'track':'relational_graph_classification'} class GraphNet(nn.Module): def __init__(self): super().__init__() self.node = nn.Sequential(nn.Linear(9, 48), nn.ReLU(), nn.Linear(48, 32), nn.ReLU()) self.head = nn.Sequential(nn.Linear(32, 24), nn.ReLU(), nn.Linear(24, 2)) def nodes(self, x): return self.node(x) def forward(self, x): return self.head(self.nodes(x).mean(1)) def sinkhorn(cost, eps, iters=35): # Uniform marginals, normalized to total mass one. n = cost.shape[-1] eps = torch.as_tensor(eps, device=cost.device, dtype=cost.dtype).clamp_min(1e-5) logk = -cost / eps u = torch.zeros_like(logk[..., :, 0]); v = torch.zeros_like(logk[..., 0, :]) logm = -np.log(n) for _ in range(iters): u = logm - torch.logsumexp(logk + v.unsqueeze(-2), dim=-1) v = logm - torch.logsumexp(logk + u.unsqueeze(-1), dim=-2) return torch.exp(logk + u.unsqueeze(-1) + v.unsqueeze(-2)) def pair_loss(net, x): # Dense lifted relational cost C[i,k] = sum_j (omega[i,j]-eta[k,j])^2. b, n, _ = x.shape perm = torch.stack([torch.randperm(n, device=x.device) for _ in range(b)]) xp = torch.gather(x, 1, perm.unsqueeze(-1).expand(-1, -1, x.shape[-1])) w, wp = x[..., :n], xp[..., :n] rows = (w[:, :, None, :] - wp[:, None, :, :]).pow(2).mean(-1) med = rows.detach().flatten(1).median(1).values.clamp_min(1e-3) P = sinkhorn(rows.detach(), (0.1 * med).view(b, 1, 1), iters=35).detach() h, hp = net.nodes(x), net.nodes(xp) # Node-level lifted signal uses the same coupling; stop-gradient P. align = ((h[:, :, None, :] - hp[:, None, :, :]).pow(2) * P[:, :, :, None]).sum(3).mean() return xp, align, P def train_one(seed, lr, idea): seed_all(seed) d = dataset(seed) net = GraphNet() opt = torch.optim.Adam(net.parameters(), lr=lr) device = 'cuda' if torch.cuda.is_available() else 'cpu' try: net.to(device); xtr, ytr = d['xtr'].to(device), d['ytr'].to(device) for _ in range(EPOCHS): net.train(); order = torch.randperm(len(xtr), device=device) for st in range(0, len(xtr), BATCH): ix = order[st:st+BATCH]; x = xtr[ix]; y = ytr[ix] if idea: xp, al, _ = pair_loss(net, x) loss = F.cross_entropy(net(x), y) + F.cross_entropy(net(xp), y) + 0.08 * al else: # Standard permutation augmentation, same base model/data budget. p = torch.stack([torch.randperm(8, device=device) for _ in range(len(x))]) xp = torch.gather(x, 1, p.unsqueeze(-1).expand(-1,-1,9)) loss = 0.5 * (F.cross_entropy(net(x), y) + F.cross_entropy(net(xp), y)) opt.zero_grad(); loss.backward(); opt.step() net.eval(); xt, yt = d['xte'].to(device), d['yte'].to(device) with torch.no_grad(): metric = float((net(xt).argmax(1) != yt).float().mean()) # Trained-system signature: output variance over arbitrary node relabelings. with torch.no_grad(): vals=[] for x in xt[:40]: outs=[] for _ in range(6): p=torch.randperm(8,device=device) xx=x[p] outs.append(torch.softmax(net(xx.unsqueeze(0)),1)[0,1]) vals.append(torch.stack(outs).var()) variance=float(torch.stack(vals).mean()) return metric, variance except RuntimeError: # Robust CPU fallback for shared-GPU failures. torch.set_num_threads(2); net = GraphNet(); opt=torch.optim.Adam(net.parameters(),lr=lr) xtr,ytr=d['xtr'],d['ytr'] for _ in range(EPOCHS): for st in range(0,len(xtr),BATCH): x=xtr[st:st+BATCH]; y=ytr[st:st+BATCH] if idea: xp,al,_=pair_loss(net,x); loss=F.cross_entropy(net(x),y)+F.cross_entropy(net(xp),y)+.08*al else: 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)) opt.zero_grad();loss.backward();opt.step() with torch.no_grad(): metric=float((net(d['xte']).argmax(1)!=d['yte']).float().mean()) variance=0.0 with torch.no_grad(): for x in d['xte'][:40]: q=torch.stack([torch.softmax(net(x[torch.randperm(8)].unsqueeze(0)),1)[0,1] for _ in range(6)]) variance += float(q.var())/40 return metric, variance def eval_cfg(lr, idea, seeds): vals=[]; sig=[] for s in seeds: m,v=train_one(s,lr,idea); vals.append(m); sig.append(v) return {'per_seed':vals,'mean':float(np.mean(vals)), 'signature_variance':float(np.mean(sig)), 'lr':lr} def main(): # Baseline sweep uses the same lr union as the idea settings. base = sweep_baseline(lambda cfg: lambda s: train_one(s, cfg['lr'], False)[0], [{'lr':x} for x in LRS], seeds=SWEEP_SEEDS) # sweep_baseline expects train_fn(seed), while evaluate does exactly that. idea_runs=[eval_cfg(lr,True,SEEDS) for lr in LRS] best=min(idea_runs,key=lambda z:z['mean']) # Re-run the selected baseline on the same eight seeds to measure its # trained-model permutation signature (the protocol sweep stores scalars). bfull = eval_cfg(base['best_cfg']['lr'], False, SEEDS) base['full']['signature_variance'] = bfull['signature_variance'] report=make_report('relational_graph_classification','graphnet',base,best,extra={ 'prediction':'alignment should reduce permutation prediction variance', 'predicted_baseline_variance_not_available':False, 'observed_baseline_variance':float(base['full'].get('signature_variance',float('nan'))), 'observed_idea_variance':best['signature_variance'], 'confirmed': bool(best['signature_variance'] < base['full'].get('signature_variance',float('inf'))), 'sinkhorn_iterations':35, 'alignment_epsilon':'0.1*median relational row cost'}) report['idea_sweep']=idea_runs report['custom_track']={'name':graph_track.META['name'],'file':'graph_track.py','domain':graph_track.META['domain']} with open('bench_report.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__=='__main__': main()