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

Failed on benchmark

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