Residual-Redundancy Adapter Clustering / bench_run.py
Mechanism confirmed, baseline not beaten
1import sys, json
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import train_model, evaluate, sweep_baseline, make_report
8from custom_multitask_track import get_dataset
9
10EPOCHS = 14
11SEEDS = tuple(range(8))
12GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 6e-3}]
13
14class AdapterNet(nn.Module):
15 def __init__(self, groups, rank=3, seed=0):
16 super().__init__(); torch.manual_seed(seed)
17 self.groups = [list(g) for g in groups]
18 self.trunk = nn.Sequential(nn.Linear(12,48), nn.ReLU())
19 self.us = nn.ParameterList([nn.Parameter(torch.randn(48,rank)*.04) for _ in self.groups])
20 self.vs = nn.ParameterList([nn.Parameter(torch.randn(rank,48)*.04) for _ in self.groups])
21 self.head = nn.Linear(48,6)
22 def forward(self,x):
23 h=self.trunk(x); y=h.new_zeros((x.shape[0],6))
24 for g,u,v in zip(self.groups,self.us,self.vs):
25 y[:,g]=self.head(h+(h@u)@v)[:,g]
26 return y
27
28def as_tensors(seed):
29 d=get_dataset(seed,800,400)
30 return {k:(torch.from_numpy(v) if isinstance(v,np.ndarray) else v) for k,v in d.items()}
31
32def tc(cov,eps=1e-6):
33 cov=np.asarray(cov,float); diag=np.maximum(np.diag(cov),eps)
34 sign,ld=np.linalg.slogdet(cov+eps*np.eye(len(diag)))
35 return float(.5*(np.log(diag).sum()-ld)) if sign>0 else 0.0
36
37def cluster_residuals(residual,k=3,alpha=.10):
38 s=np.cov(residual,rowvar=False,ddof=1); s=(1-alpha)*s+alpha*np.diag(np.diag(s))
39 def merge(a,b):
40 ab=list(a)+list(b)
41 return tc(s[np.ix_(ab,ab)])-tc(s[np.ix_(a,a)])-tc(s[np.ix_(b,b)])
42 groups=[[i] for i in range(s.shape[0])]
43 while len(groups)>k:
44 _,i,j=max((merge(a,b),i,j) for i,a in enumerate(groups) for j,b in enumerate(groups) if j>i)
45 groups[i]=sorted(groups[i]+groups[j]); del groups[j]
46 return sorted(groups),s
47
48def fit_metric(seed,lr,mode,return_sig=False):
49 d=as_tensors(seed); n=len(d['xtr']); cut=n//5
50 # Held-out warm-up residual buffer: no gradient updates use this buffer.
51 warm=AdapterNet([list(range(6))],rank=9,seed=seed)
52 wd={'xtr':d['xtr'][:cut],'ytr':d['ytr'][:cut],'xte':d['xtr'][cut:],'yte':d['ytr'][cut:],'task':'regression'}
53 warm,_,_=train_model(warm,wd,epochs=5,lr=lr,batch=128)
54 with torch.no_grad():
55 dev=next(warm.parameters()).device
56 pred=warm(d['xtr'][cut:].to(dev)).cpu()
57 residual=(d['ytr'][cut:]-pred).numpy()
58 discovered,cov=cluster_residuals(residual,3)
59 groups=[list(range(6))] if mode=='baseline' else discovered
60 rank=9 if mode=='baseline' else 3
61 ds={'xtr':d['xtr'][cut:],'ytr':d['ytr'][cut:],'xte':d['xte'],'yte':d['yte'],'task':'regression'}
62 net=AdapterNet(groups,rank=rank,seed=seed)
63 _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=128)
64 if return_sig: return metric,{'groups':discovered,'all_tc':tc(cov),'within_tc':float(np.mean([tc(cov[np.ix_(g,g)]) for g in discovered]))}
65 return metric
66
67def baseline_fn(cfg): return lambda seed: fit_metric(seed,cfg['lr'],'baseline')
68def idea_fn(cfg): return lambda seed: fit_metric(seed,cfg['lr'],'idea')
69
70def main():
71 base=sweep_baseline(baseline_fn,GRID,seeds=(0,1,2,3))
72 candidates=[(evaluate(idea_fn(c),seeds=SEEDS),c) for c in GRID]
73 idea,cfg=min(candidates,key=lambda z:z[0]['mean'])
74 sig=[fit_metric(s,cfg['lr'],'idea',True)[1] for s in SEEDS]
75 rec=float(np.mean([x['groups']==[[0,1],[2,3],[4,5]] for x in sig]))
76 extra={'prediction':'residual-TC clustering discovers correlated task pairs and lowers within-cluster residual dependence','predicted_mean_all_tc':float(np.mean([x['all_tc'] for x in sig])),'observed_mean_within_cluster_tc':float(np.mean([x['within_tc'] for x in sig])),'pair_recovery_rate':rec,'confirmed':bool(rec>=0.5)}
77 rep=make_report('custom_correlated_multitask_regression','mlp_tiny',base,idea,extra)
78 rep['custom_track']={'name':'correlated_multitask_regression','file':'custom_multitask_track.py','domain':'multi_task_learning'}
79 rep['idea_sweep']=[{'cfg':c,'mean':r['mean']} for r,c in candidates]
80 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
81if __name__=='__main__': main()