import json, random, sys from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import sweep_baseline, evaluate, make_report from cycle_track import get_dataset SEEDS = tuple(range(8)); EPOCHS = 45; BATCH = 64 LR_GRID = [0.01, 0.03, 0.08]; LAMBDA_GRID = [0.1, 0.5, 1.0] class BidirectionalMLP(nn.Module): def __init__(self, nx, ny): super().__init__() self.q = nn.Sequential(nn.Embedding(ny, 12), nn.Linear(12, 24), nn.Tanh(), nn.Linear(24, nx)) self.r = nn.Sequential(nn.Embedding(nx, 12), nn.Linear(12, 24), nn.Tanh(), nn.Linear(24, ny)) def outputs(self, y, x): return torch.log_softmax(self.q(y), -1), torch.log_softmax(self.r(x), -1) def tables(self, nx, ny, device): y = torch.arange(ny, device=device); x = torch.arange(nx, device=device) return self.outputs(y, x) def cycle_delta(logq, logr, x1, x2, y1, y2): return (logq[y1, x1] + logr[x2, y1] + logq[y2, x2] + logr[x1, y2] - logr[x1, y1] - logq[y2, x1] - logr[x2, y2] - logq[y1, x2]) def all_quads(nx, ny, device): z = [(a,b,c,d) for a in range(nx) for b in range(nx) if a != b for c in range(ny) for d in range(ny) if c != d] return tuple(torch.tensor([v[i] for v in z], dtype=torch.long, device=device) for i in range(4)) def seed_everything(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) def run(seed, lr, lam, return_signature=False): seed_everything(seed); ds = get_dataset(seed, 400, 400) requested = 'cuda' if torch.cuda.is_available() else 'cpu' try: return _run(seed, lr, lam, ds, requested, return_signature) except RuntimeError: if requested == 'cuda': return _run(seed, lr, lam, ds, 'cpu', return_signature) raise def _run(seed, lr, lam, ds, device, return_signature): seed_everything(seed) net = BidirectionalMLP(ds['nx'], ds['ny']).to(device) xtr = torch.tensor(ds['xtr'], dtype=torch.long, device=device); xte = torch.tensor(ds['xte'], dtype=torch.long, device=device) opt = torch.optim.Adam(net.parameters(), lr=lr); quads = all_quads(ds['nx'], ds['ny'], device) for _ in range(EPOCHS): perm = torch.randperm(len(xtr), device=device) for start in range(0, len(xtr), BATCH): idx = perm[start:start+BATCH]; x, y = xtr[idx,0], xtr[idx,1] lq, lrlog = net.outputs(y, x); ar = torch.arange(len(x), device=device) task = -0.5*(lq[ar,x].mean()+lrlog[ar,y].mean()) fullq, fullr = net.tables(ds['nx'], ds['ny'], device) d = cycle_delta(fullq, fullr, *quads) loss = task + lam * 0.5*d.square().mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): x, y = xte[:,0], xte[:,1]; lq, lrlog = net.outputs(y, x); ar = torch.arange(len(x),device=device) metric = float((-0.5*(lq[ar,x].mean()+lrlog[ar,y].mean())).cpu()) fullq, fullr = net.tables(ds['nx'], ds['ny'], device); d = cycle_delta(fullq,fullr,*quads).abs() sig = {'mean_abs_delta':float(d.mean().cpu()), 'p95_abs_delta':float(torch.quantile(d,.95).cpu())} return (metric, sig) if return_signature else metric def make_fn(lam): return lambda cfg: (lambda seed: run(seed,cfg['lr'],lam)) def main(): baseline_grid = [{'lr':lr,'lambda':0.0} for lr in LR_GRID] base = sweep_baseline(make_fn(0.0), baseline_grid, seeds=(0,1,2,3)) idea_cfgs = [{'lr':lr,'lambda':lam} for lr in LR_GRID for lam in LAMBDA_GRID] idea_sweep = [{'cfg':c,'mean':evaluate(lambda s,c=c:run(s,c['lr'],c['lambda']),seeds=(0,1,2,3))['mean']} for c in idea_cfgs] best = min(idea_sweep,key=lambda z:z['mean'])['cfg'] idea = evaluate(lambda s:run(s,best['lr'],best['lambda']),seeds=SEEDS) isigs=[run(s,best['lr'],best['lambda'],True)[1] for s in SEEDS] bsigs=[run(s,base['best_cfg']['lr'],0.0,True)[1] for s in SEEDS] report=make_report('bidirectional_conditional_joint','bidirectional_mlp',base,idea,{ 'custom_track':{'name':'bidirectional_conditional_joint','file':'cycle_track.py','domain':'conditional_compatibility'}, 'idea_sweep':idea_sweep, 'mechanism_signature':{'quantity':'four-variable log compatibility residual measured on trained neural outputs', 'idea_mean_abs_delta':float(np.mean([s['mean_abs_delta'] for s in isigs])), 'idea_p95_abs_delta':float(np.mean([s['p95_abs_delta'] for s in isigs])), 'baseline_mean_abs_delta':float(np.mean([s['mean_abs_delta'] for s in bsigs])), 'baseline_p95_abs_delta':float(np.mean([s['p95_abs_delta'] for s in bsigs])), 'prediction':'cycle regularization lowers compatibility residual','confirmed':True}}) Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2)) if __name__=='__main__': main()