Probe-Then-Partitioned Multi-Task Trunk / bench_run.py
Failed on benchmark
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()