import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import train_model, evaluate, sweep_baseline, make_report, reload_custom_tracks import importlib.util _spec=importlib.util.spec_from_file_location('matched_track','/home/maxwelhelp/all/math2nn/bench/custom_tracks/correlated_multitask_regression.py') _track=importlib.util.module_from_spec(_spec); _spec.loader.exec_module(_track) def get_ds(seed): d0=_track.get_dataset(seed=seed,n_train=400,n_test=400) return {'track':TRACK,'task':d0.get('task','regression'),'metric':d0.get('metric','mse'), 'xtr':torch.as_tensor(np.asarray(d0['xtr']),dtype=torch.float32), 'ytr':torch.as_tensor(np.asarray(d0['ytr']),dtype=torch.float32), 'xte':torch.as_tensor(np.asarray(d0['xte']),dtype=torch.float32), 'yte':torch.as_tensor(np.asarray(d0['yte']),dtype=torch.float32), 'input_shape':tuple(np.asarray(d0['xtr']).shape[1:]),'out_dim':int(np.asarray(d0['ytr']).shape[1])} TRACK = 'correlated_multitask_regression' MODEL = 'mlp_tiny' SEEDS = [11,22,33,44,55,66,77,88] SWEEP_SEEDS = [11,22,33] LR_GRID = [1e-3, 3e-3, 1e-2] EPOCHS = 22 PROBE_EPOCHS = 4 BATCH = 128 WIDTH = 64 class SharedMTL(nn.Module): def __init__(self, p, width=WIDTH): super().__init__() self.trunk = nn.Sequential(nn.Linear(p,width), nn.ReLU(), nn.Linear(width,width), nn.ReLU()) self.heads = nn.ModuleList([nn.Linear(width,1) for _ in range(6)]) def features(self, x): return self.trunk(x) def forward(self, x): z = self.features(x) return torch.cat([h(z) for h in self.heads], 1) class PartitionMTL(nn.Module): def __init__(self, p, groups, width=WIDTH): super().__init__() self.groups = list(map(int, groups)); self.ng = max(self.groups)+1 self.trunks = nn.ModuleList([nn.Sequential(nn.Linear(p,width), nn.ReLU(), nn.Linear(width,width), nn.ReLU()) for _ in range(self.ng)]) self.heads = nn.ModuleList([nn.Linear(width,1) for _ in range(6)]) def forward(self, x): return torch.cat([self.heads[i](self.trunks[self.groups[i]](x)) for i in range(6)], 1) def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def loss_for(out, y): return ((out-y)**2).mean() def probe_groups(ds, seed): seed_all(seed) p = ds['xtr'].shape[1] trunk = nn.Sequential(nn.Linear(p, WIDTH), nn.ReLU(), nn.Linear(WIDTH, WIDTH), nn.ReLU()) heads = nn.ModuleList([nn.Linear(WIDTH,1) for _ in range(6)]) opt = torch.optim.Adam(list(trunk.parameters())+list(heads.parameters()), lr=3e-3) x,y=ds['xtr'],ds['ytr'] for _ in range(PROBE_EPOCHS): for j in range(0,len(x),BATCH): z=trunk(x[j:j+BATCH]); out=torch.cat([h(z) for h in heads],1) loss=loss_for(out,y[j:j+BATCH]); opt.zero_grad(); loss.backward(); opt.step() emb=[] idx=torch.arange(min(256,len(x))) for i,h in enumerate(heads): trunk.zero_grad(set_to_none=True); h.zero_grad(set_to_none=True) z=trunk(x[idx]); li=((h(z)-y[idx,i:i+1])**2).mean(); li.backward() act=z.detach().mean(0).numpy() grad=np.concatenate([p.grad.detach().numpy().ravel() for p in trunk.parameters() if p.grad is not None]) grad=grad/(np.linalg.norm(grad)+1e-12) e=np.r_[act,grad]; e=e/(np.linalg.norm(e)+1e-12); emb.append(e) E=np.asarray(emb); D=1-E@E.T; np.fill_diagonal(D,0) m=2; core=np.zeros(6) for i in range(6): core[i]=np.sort(np.delete(D[i],i))[m-1] MR=np.maximum(D,np.maximum(core[:,None],core[None,:])) edges=sorted((MR[i,j],i,j) for i in range(6) for j in range(i)) par=list(range(6)) def find(a): while par[a]!=a: par[a]=par[par[a]]; a=par[a] return a mst=[] for w,i,j in edges: a,b=find(i),find(j) if a!=b: par[a]=b; mst.append((w,i,j)) ws=np.array([z[0] for z in mst]); cut=float('inf') if len(ws)>=2: s=np.sort(ws); k=int(np.argmax(np.diff(s))) if s[k+1] > 1.15*np.median(s): cut=float(s[k+1]) par=list(range(6)) for w,i,j in mst: if w