Shared-Private Matrix-Weighted Expert Layers / bench_experiment.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, sweep_baseline, evaluate, make_report
9
10DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
11OUT = Path('bench_report.json')
12
13class ExpertNet(nn.Module):
14 def __init__(self, hidden=32, out_dim=1):
15 super().__init__()
16 self.rnn = nn.GRU(3, hidden, batch_first=True)
17 self.proj = nn.Linear(hidden, hidden)
18 self.head = nn.Linear(hidden, out_dim)
19 def forward(self, x):
20 _, h = self.rnn(x.view(x.shape[0], -1, 3))
21 z = torch.tanh(self.proj(h[-1]))
22 return self.head(z), z
23
24
25def seed_all(seed):
26 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
27 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
28
29
30def tensors(ds):
31 return tuple(torch.tensor(ds[k], dtype=torch.float32) for k in ('xtr','ytr','xte','yte'))
32
33
34def run_system(seed, cfg, idea, collect=False):
35 seed_all(seed)
36 ds = get_dataset('dynamics', seed, n_train=400, n_test=120)
37 xtr, ytr, xte, yte = tensors(ds)
38 # Four task/expert views: deterministic perturbations represent related tasks.
39 task_offsets = torch.tensor([-0.12, -0.04, 0.04, 0.12])
40 models = [ExpertNet().to(DEVICE) for _ in range(4)]
41 params = [p for m in models for p in m.parameters()]
42 opt = torch.optim.AdamW(params, lr=cfg['lr'], weight_decay=cfg['wd'])
43 U = torch.zeros(32, cfg['rank'], device=DEVICE)
44 U[:cfg['rank'], :] = torch.eye(cfg['rank'], device=DEVICE)
45 # Fixed orthonormal shared feature directions; coupling is exactly U C U^T.
46 n = len(xtr); batch = min(128, cfg['batch'])
47 for _ in range(cfg['epochs']):
48 order = torch.randperm(n)
49 for start in range(0, n, batch):
50 ix = order[start:start+batch]
51 xb, yb = xtr[ix].to(DEVICE), ytr[ix].to(DEVICE).reshape(-1,1)
52 opt.zero_grad(); zs=[]; losses=[]
53 for k,m in enumerate(models):
54 pred,z=m(xb); zs.append(z)
55 target = yb + task_offsets[k].to(DEVICE)
56 losses.append(((pred-target)**2).mean())
57 loss = torch.stack(losses).mean()
58 if idea:
59 # Complete graph, C=I: exact shared-coordinate regularizer.
60 zstack=torch.stack(zs)
61 q=torch.einsum('dr,bd->br', U, zstack.reshape(-1,32)).reshape(4,-1,cfg['rank'])
62 reg=0.0
63 for i in range(4):
64 for j in range(i+1,4): reg += ((q[i]-q[j])**2).mean()
65 loss = loss + cfg['coupling'] * reg / 6.0
66 loss.backward(); opt.step()
67 with torch.no_grad():
68 preds=[]; features=[]
69 for k,m in enumerate(models):
70 p,z=m(xte.to(DEVICE)); preds.append(p.cpu()); features.append(z.cpu())
71 pred=torch.stack(preds); target=yte.reshape(1,-1,1)+task_offsets.reshape(4,1,1)
72 mse=((pred-target)**2).mean(dim=(1,2)).numpy()
73 feat=torch.stack(features)
74 q=feat[:,:,:cfg['rank']]
75 priv=feat[:,:,cfg['rank']:]
76 signature={'shared_disagreement':float(((q-q.mean(0))**2).mean()),'private_variance':float(priv.var()),'test_mse':float(mse.mean())}
77 return float(mse.mean()), signature
78
79
80def run_cfg(cfg, idea):
81 def train(seed): return run_system(seed,cfg,idea)[0]
82 return train
83
84
85def _main():
86 # Union parity: baseline and idea both run all three lr values and same weight decay.
87 grid=[{'lr':v,'wd':1e-4,'epochs':12,'batch':128,'rank':8,'coupling':0.0} for v in (1e-3,3e-3,1e-2)]
88 base=sweep_baseline(lambda c: run_cfg(c,False), grid)
89 best=base['best_cfg']
90 idea_grid=[dict(best,coupling=c) for c in (0.0,0.03,0.10)]
91 # Include the baseline best setting and two nearby intervention settings.
92 idea_runs=[]
93 for c in idea_grid:
94 idea_runs.append((c,evaluate(run_cfg(c,True))))
95 idea_cfg, idea_res=min(idea_runs,key=lambda t:t[1]['mean'])
96 sigs=[run_system(s,idea_cfg,True,True)[1] for s in range(8)]
97 bsigs=[run_system(s,best,False,True)[1] for s in range(8)]
98 signature={
99 'predicted': 'shared disagreement decreases while private variance remains nonzero',
100 'observed_idea_shared_disagreement_mean':float(np.mean([x['shared_disagreement'] for x in sigs])),
101 'observed_baseline_shared_disagreement_mean':float(np.mean([x['shared_disagreement'] for x in bsigs])),
102 'observed_idea_private_variance_mean':float(np.mean([x['private_variance'] for x in sigs])),
103 'observed_baseline_private_variance_mean':float(np.mean([x['private_variance'] for x in bsigs])),
104 '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)
105 }
106 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})
107 rep['custom_track']=None
108 OUT.write_text(json.dumps(rep,indent=2))
109 print(json.dumps(rep,indent=2))
110
111def main():
112 global DEVICE
113 try:
114 return _main()
115 except RuntimeError as exc:
116 if DEVICE == 'cuda':
117 print('[bench] CUDA failure; retrying entire protocol on CPU:', str(exc)[:160])
118 DEVICE = 'cpu'
119 return _main()
120 raise
121
122if __name__=='__main__': main()