Probe-Then-Partitioned Multi-Task Trunk / bench_run.py

Failed on benchmark

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 train_model, evaluate, sweep_baseline, make_report, reload_custom_tracks
  9import importlib.util
 10_spec=importlib.util.spec_from_file_location('matched_track','/home/maxwelhelp/all/math2nn/bench/custom_tracks/correlated_multitask_regression.py')
 11_track=importlib.util.module_from_spec(_spec); _spec.loader.exec_module(_track)
 12
 13def get_ds(seed):
 14    d0=_track.get_dataset(seed=seed,n_train=400,n_test=400)
 15    return {'track':TRACK,'task':d0.get('task','regression'),'metric':d0.get('metric','mse'),
 16            'xtr':torch.as_tensor(np.asarray(d0['xtr']),dtype=torch.float32),
 17            'ytr':torch.as_tensor(np.asarray(d0['ytr']),dtype=torch.float32),
 18            'xte':torch.as_tensor(np.asarray(d0['xte']),dtype=torch.float32),
 19            'yte':torch.as_tensor(np.asarray(d0['yte']),dtype=torch.float32),
 20            'input_shape':tuple(np.asarray(d0['xtr']).shape[1:]),'out_dim':int(np.asarray(d0['ytr']).shape[1])}
 21
 22
 23TRACK = 'correlated_multitask_regression'
 24MODEL = 'mlp_tiny'
 25SEEDS = [11,22,33,44,55,66,77,88]
 26SWEEP_SEEDS = [11,22,33]
 27LR_GRID = [1e-3, 3e-3, 1e-2]
 28EPOCHS = 22
 29PROBE_EPOCHS = 4
 30BATCH = 128
 31WIDTH = 64
 32
 33class SharedMTL(nn.Module):
 34    def __init__(self, p, width=WIDTH):
 35        super().__init__()
 36        self.trunk = nn.Sequential(nn.Linear(p,width), nn.ReLU(), nn.Linear(width,width), nn.ReLU())
 37        self.heads = nn.ModuleList([nn.Linear(width,1) for _ in range(6)])
 38    def features(self, x): return self.trunk(x)
 39    def forward(self, x):
 40        z = self.features(x)
 41        return torch.cat([h(z) for h in self.heads], 1)
 42
 43class PartitionMTL(nn.Module):
 44    def __init__(self, p, groups, width=WIDTH):
 45        super().__init__()
 46        self.groups = list(map(int, groups)); self.ng = max(self.groups)+1
 47        self.trunks = nn.ModuleList([nn.Sequential(nn.Linear(p,width), nn.ReLU(), nn.Linear(width,width), nn.ReLU()) for _ in range(self.ng)])
 48        self.heads = nn.ModuleList([nn.Linear(width,1) for _ in range(6)])
 49    def forward(self, x):
 50        return torch.cat([self.heads[i](self.trunks[self.groups[i]](x)) for i in range(6)], 1)
 51
 52def seed_all(s):
 53    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 54    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 55
 56def loss_for(out, y): return ((out-y)**2).mean()
 57
 58def probe_groups(ds, seed):
 59    seed_all(seed)
 60    p = ds['xtr'].shape[1]
 61    trunk = nn.Sequential(nn.Linear(p, WIDTH), nn.ReLU(), nn.Linear(WIDTH, WIDTH), nn.ReLU())
 62    heads = nn.ModuleList([nn.Linear(WIDTH,1) for _ in range(6)])
 63    opt = torch.optim.Adam(list(trunk.parameters())+list(heads.parameters()), lr=3e-3)
 64    x,y=ds['xtr'],ds['ytr']
 65    for _ in range(PROBE_EPOCHS):
 66        for j in range(0,len(x),BATCH):
 67            z=trunk(x[j:j+BATCH]); out=torch.cat([h(z) for h in heads],1)
 68            loss=loss_for(out,y[j:j+BATCH]); opt.zero_grad(); loss.backward(); opt.step()
 69    emb=[]
 70    idx=torch.arange(min(256,len(x)))
 71    for i,h in enumerate(heads):
 72        trunk.zero_grad(set_to_none=True); h.zero_grad(set_to_none=True)
 73        z=trunk(x[idx]); li=((h(z)-y[idx,i:i+1])**2).mean(); li.backward()
 74        act=z.detach().mean(0).numpy()
 75        grad=np.concatenate([p.grad.detach().numpy().ravel() for p in trunk.parameters() if p.grad is not None])
 76        grad=grad/(np.linalg.norm(grad)+1e-12)
 77        e=np.r_[act,grad]; e=e/(np.linalg.norm(e)+1e-12); emb.append(e)
 78    E=np.asarray(emb); D=1-E@E.T; np.fill_diagonal(D,0)
 79    m=2; core=np.zeros(6)
 80    for i in range(6): core[i]=np.sort(np.delete(D[i],i))[m-1]
 81    MR=np.maximum(D,np.maximum(core[:,None],core[None,:]))
 82    edges=sorted((MR[i,j],i,j) for i in range(6) for j in range(i))
 83    par=list(range(6))
 84    def find(a):
 85        while par[a]!=a:
 86            par[a]=par[par[a]]; a=par[a]
 87        return a
 88    mst=[]
 89    for w,i,j in edges:
 90        a,b=find(i),find(j)
 91        if a!=b: par[a]=b; mst.append((w,i,j))
 92    ws=np.array([z[0] for z in mst]); cut=float('inf')
 93    if len(ws)>=2:
 94        s=np.sort(ws); k=int(np.argmax(np.diff(s)))
 95        if s[k+1] > 1.15*np.median(s): cut=float(s[k+1])
 96    par=list(range(6))
 97    for w,i,j in mst:
 98        if w<cut:
 99            a,b=find(i),find(j)
100            if a!=b: par[a]=b
101    roots={}; groups=[]
102    for i in range(6):
103        r=find(i); roots.setdefault(r,len(roots)); groups.append(roots[r])
104    return groups, float(np.mean(D)), float(np.mean(core))
105
106def train_baseline(cfg, seed, return_aux=False):
107    seed_all(seed); ds=get_ds(seed)
108    net=SharedMTL(ds['xtr'].shape[1]); net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
109    if return_aux: return metric, {'groups':[0]*6,'probe':None,'model':net,'ds':ds}
110    return metric
111
112def train_idea(cfg, seed, return_aux=False):
113    seed_all(seed); ds=get_ds(seed)
114    groups,md,mc=probe_groups(ds,seed)
115    seed_all(seed)
116    net=PartitionMTL(ds['xtr'].shape[1],groups)
117    net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
118    if return_aux: return metric, {'groups':groups,'probe_mean_distance':md,'probe_mean_core':mc,'model':net,'ds':ds}
119    return metric
120
121def signature(seed, base_metric, idea_metric, groups):
122    # Re-test the mechanism on trained systems: task-gradient cosine conflict.
123    def conflicts(kind):
124        ds=get_ds(seed); seed_all(seed)
125        _,aux=(train_baseline({'lr':3e-3},seed,True) if kind=='base' else train_idea({'lr':3e-3},seed,True))
126        net,ds=aux['model'],aux['ds']; device=next(net.parameters()).device; x,y=ds['xtr'][:128].to(device),ds['ytr'][:128].to(device)
127        gs=[]
128        for i in range(6):
129            net.zero_grad(set_to_none=True); out=net(x); ((out[:,i:i+1]-y[:128,i:i+1])**2).mean().backward()
130            ps=[]
131            for p in net.parameters():
132                if p.grad is not None: ps.append(p.grad.detach().flatten())
133            gs.append(torch.cat(ps))
134        C=torch.stack(gs); cs=[]
135        for i in range(6):
136            for j in range(i): cs.append(float(torch.dot(C[i],C[j])/(C[i].norm()*C[j].norm()+1e-12)))
137        return float(np.mean(np.asarray(cs)<0))
138    cb=conflicts('base'); ci=conflicts('idea')
139    return {'predicted':'partitioning reduces negative cross-task gradient cosine','baseline_conflict_rate':cb,'idea_conflict_rate':ci,'observed_reduction':cb-ci,'confirmed':bool(ci < cb)}
140
141def main():
142    reload_custom_tracks()
143    grid=[{'lr':x} for x in LR_GRID]
144    base=sweep_baseline(lambda cfg: lambda s: train_baseline(cfg,s),grid,seeds=SWEEP_SEEDS)
145    idea_grid={x: evaluate(lambda s,c={'lr':x}: train_idea(c,s),seeds=SWEEP_SEEDS) for x in LR_GRID}
146    best_lr=min(idea_grid,key=lambda x: idea_grid[x]['mean'])
147    idea=evaluate(lambda s: train_idea({'lr':best_lr},s),seeds=SEEDS)
148    # Also run two nearby settings on the same sweep union, already evaluated above.
149    sig=signature(SEEDS[0],base['full']['mean'],idea['mean'],train_idea({'lr':best_lr},SEEDS[0],True)[1]['groups'])
150    rep=make_report(TRACK,MODEL,base,idea,{'track_rationale':'multi-task regression with six task heads and latent task groups; custom track is structurally matched','idea_grid':idea_grid,'selected_idea_lr':best_lr,'probe_epochs':PROBE_EPOCHS,'groups_seed0':train_idea({'lr':best_lr},SEEDS[0],True)[1]['groups'],'mechanism_signature':sig})
151    rep['parameter_note']='Partitioned trunks replicate the trunk per discovered cluster; heads remain separate. Equal optimizer/data/epoch budget, but inference parameter count can differ.'
152    Path('bench_report.json').write_text(json.dumps(rep,indent=2))
153    print(json.dumps(rep,indent=2))
154if __name__=='__main__': main()