Residual-Redundancy Adapter Clustering / bench_run.py

Mechanism confirmed, baseline not beaten

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