import sys, json, random 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 get_dataset, sweep_baseline, evaluate, make_report DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' OUT = Path('bench_report.json') class ExpertNet(nn.Module): def __init__(self, hidden=32, out_dim=1): super().__init__() self.rnn = nn.GRU(3, hidden, batch_first=True) self.proj = nn.Linear(hidden, hidden) self.head = nn.Linear(hidden, out_dim) def forward(self, x): _, h = self.rnn(x.view(x.shape[0], -1, 3)) z = torch.tanh(self.proj(h[-1])) return self.head(z), z def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def tensors(ds): return tuple(torch.tensor(ds[k], dtype=torch.float32) for k in ('xtr','ytr','xte','yte')) def run_system(seed, cfg, idea, collect=False): seed_all(seed) ds = get_dataset('dynamics', seed, n_train=400, n_test=120) xtr, ytr, xte, yte = tensors(ds) # Four task/expert views: deterministic perturbations represent related tasks. task_offsets = torch.tensor([-0.12, -0.04, 0.04, 0.12]) models = [ExpertNet().to(DEVICE) for _ in range(4)] params = [p for m in models for p in m.parameters()] opt = torch.optim.AdamW(params, lr=cfg['lr'], weight_decay=cfg['wd']) U = torch.zeros(32, cfg['rank'], device=DEVICE) U[:cfg['rank'], :] = torch.eye(cfg['rank'], device=DEVICE) # Fixed orthonormal shared feature directions; coupling is exactly U C U^T. n = len(xtr); batch = min(128, cfg['batch']) for _ in range(cfg['epochs']): order = torch.randperm(n) for start in range(0, n, batch): ix = order[start:start+batch] xb, yb = xtr[ix].to(DEVICE), ytr[ix].to(DEVICE).reshape(-1,1) opt.zero_grad(); zs=[]; losses=[] for k,m in enumerate(models): pred,z=m(xb); zs.append(z) target = yb + task_offsets[k].to(DEVICE) losses.append(((pred-target)**2).mean()) loss = torch.stack(losses).mean() if idea: # Complete graph, C=I: exact shared-coordinate regularizer. zstack=torch.stack(zs) q=torch.einsum('dr,bd->br', U, zstack.reshape(-1,32)).reshape(4,-1,cfg['rank']) reg=0.0 for i in range(4): for j in range(i+1,4): reg += ((q[i]-q[j])**2).mean() loss = loss + cfg['coupling'] * reg / 6.0 loss.backward(); opt.step() with torch.no_grad(): preds=[]; features=[] for k,m in enumerate(models): p,z=m(xte.to(DEVICE)); preds.append(p.cpu()); features.append(z.cpu()) pred=torch.stack(preds); target=yte.reshape(1,-1,1)+task_offsets.reshape(4,1,1) mse=((pred-target)**2).mean(dim=(1,2)).numpy() feat=torch.stack(features) q=feat[:,:,:cfg['rank']] priv=feat[:,:,cfg['rank']:] signature={'shared_disagreement':float(((q-q.mean(0))**2).mean()),'private_variance':float(priv.var()),'test_mse':float(mse.mean())} return float(mse.mean()), signature def run_cfg(cfg, idea): def train(seed): return run_system(seed,cfg,idea)[0] return train def _main(): # Union parity: baseline and idea both run all three lr values and same weight decay. grid=[{'lr':v,'wd':1e-4,'epochs':12,'batch':128,'rank':8,'coupling':0.0} for v in (1e-3,3e-3,1e-2)] base=sweep_baseline(lambda c: run_cfg(c,False), grid) best=base['best_cfg'] idea_grid=[dict(best,coupling=c) for c in (0.0,0.03,0.10)] # Include the baseline best setting and two nearby intervention settings. idea_runs=[] for c in idea_grid: idea_runs.append((c,evaluate(run_cfg(c,True)))) idea_cfg, idea_res=min(idea_runs,key=lambda t:t[1]['mean']) sigs=[run_system(s,idea_cfg,True,True)[1] for s in range(8)] bsigs=[run_system(s,best,False,True)[1] for s in range(8)] signature={ 'predicted': 'shared disagreement decreases while private variance remains nonzero', 'observed_idea_shared_disagreement_mean':float(np.mean([x['shared_disagreement'] for x in sigs])), 'observed_baseline_shared_disagreement_mean':float(np.mean([x['shared_disagreement'] for x in bsigs])), 'observed_idea_private_variance_mean':float(np.mean([x['private_variance'] for x in sigs])), 'observed_baseline_private_variance_mean':float(np.mean([x['private_variance'] for x in bsigs])), 'confirmed': bool(np.mean([x['shared_disagreement'] for x in sigs]) < np.mean([x['shared_disagreement'] for x in bsigs]) and np.mean([x['private_variance'] for x in sigs]) > 1e-8) } rep=make_report('dynamics','rnn_small',base,idea_res,{'mechanism_signature':signature,'idea_sweep':[{'cfg':c,'mean':r['mean']} for c,r in idea_runs], 'device':DEVICE}) rep['custom_track']=None OUT.write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) def main(): global DEVICE try: return _main() except RuntimeError as exc: if DEVICE == 'cuda': print('[bench] CUDA failure; retrying entire protocol on CPU:', str(exc)[:160]) DEVICE = 'cpu' return _main() raise if __name__=='__main__': main()