Submetry-Lifted Relational Alignment / bench_run.py

Mechanism confirmed, baseline not beaten

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