Shared-Private Matrix-Weighted Expert Layers / bench_experiment.py

Mechanism confirmed, baseline not beaten

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