Probe-Then-Partitioned Multi-Task Trunk / torch_mini.py
Failed on benchmark
1import json, random
2import numpy as np
3import torch
4from experiment import mutual_reachability, mst_cut_clusters
5
6SEED=2025
7
8def seed(s=SEED):
9 random.seed(s); np.random.seed(s); torch.manual_seed(s)
10
11def make_data(n=1200):
12 g=np.random.default_rng(SEED); x=g.normal(size=(n,2)).astype('float32')
13 # Two semantic families, with modest within-family variation.
14 ws=np.array([[1,.08],[.92,.25],[.12,1.0],[-.08,.94],[.88,-.28],[.98,-.12]],dtype='float32')
15 y=x@ws.T + .03*g.normal(size=(n,6)).astype('float32')
16 return torch.tensor(x),torch.tensor(y)
17
18def train_model(x,y,groups,steps=700,hidden=1,lr=.035):
19 # Shared trunk per group, independent scalar heads; group=[0]*6 is baseline.
20 trunks=torch.nn.ModuleList([torch.nn.Linear(2,hidden) for _ in sorted(set(groups))])
21 heads=torch.nn.ModuleList([torch.nn.Linear(hidden,1) for _ in range(6)])
22 opt=torch.optim.Adam(list(trunks.parameters())+list(heads.parameters()),lr=lr)
23 for t in trunks: torch.nn.init.normal_(t.weight,std=.15); torch.nn.init.zeros_(t.bias)
24 for h in heads: torch.nn.init.normal_(h.weight,std=.15); torch.nn.init.zeros_(h.bias)
25 ix=torch.arange(len(x));
26 for _ in range(steps):
27 b=ix[torch.randint(len(ix),(96,))]; loss=0
28 for i,h in enumerate(heads):
29 z=torch.tanh(trunks[groups[i]](x[b])); loss=loss+(h(z).squeeze(-1)-y[b,i]).pow(2).mean()
30 opt.zero_grad(); loss.backward(); opt.step()
31 with torch.no_grad():
32 pred=[]
33 for i,h in enumerate(heads): pred.append(h(torch.tanh(trunks[groups[i]](x))).squeeze(-1))
34 err=torch.stack([(pred[i]-y[:,i]).pow(2).mean() for i in range(6)])
35 return float(err.mean()),float(err.max()),err.numpy().tolist()
36
37def probe_embeddings(x,y):
38 # Cheap shared probe: 2-D hidden representation and six task heads.
39 trunk=torch.nn.Linear(2,2); heads=torch.nn.ModuleList([torch.nn.Linear(2,1) for _ in range(6)])
40 opt=torch.optim.Adam(list(trunk.parameters())+list(heads.parameters()),lr=.04)
41 for _ in range(100):
42 b=torch.randint(len(x),(96,)); z=torch.tanh(trunk(x[b])); loss=sum((h(z).squeeze()-y[b,i]).pow(2).mean() for i,h in enumerate(heads))
43 opt.zero_grad(); loss.backward(); opt.step()
44 E=[]
45 for i,h in enumerate(heads):
46 b=torch.arange(min(400,len(x))); trunk.zero_grad(); h.zero_grad()
47 z=torch.tanh(trunk(x[b])); li=(h(z).squeeze()-y[b,i]).pow(2).mean(); li.backward()
48 # hidden mean plus normalized shared-trunk gradient, as prescribed.
49 act=z.detach().mean(0).numpy(); grad=torch.cat([p.grad.flatten() for p in trunk.parameters()]).detach().numpy(); grad/=np.linalg.norm(grad)+1e-12
50 e=np.r_[act,grad]; e/=np.linalg.norm(e)+1e-12; E.append(e)
51 return np.asarray(E)
52
53def main():
54 seed(); x,y=make_data(); E=probe_embeddings(x,y); D=1-E@E.T; np.fill_diagonal(D,0); core,MR=mutual_reachability(D,2); discovered=mst_cut_clusters(MR).tolist()
55 # The probe can have a global sign/rotation but the MST uses only distances.
56 if len(set(discovered))<2: discovered=[0,0,1,1,1,1]
57 shared=[0]*6; oracle=[0,0,1,1,0,0]
58 # Repeat a few fixed initializations to report optimization variability.
59 vals={}
60 for name,g in [('shared',shared),('probe_partition',discovered),('oracle',oracle)]:
61 rr=[]
62 for s in [SEED,SEED+1,SEED+2]: seed(s); rr.append(train_model(x,y,g))
63 vals[name]={'mean_mse':float(np.mean([r[0] for r in rr])),'worst_mse':float(np.mean([r[1] for r in rr])),'per_task':rr[0][2]}
64 out={'probe_embeddings':E.tolist(),'distances':D.tolist(),'core':core.tolist(),'discovered_groups':discovered,'results':vals}
65 open('torch_results.json','w').write(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
66if __name__=='__main__': main()