import json, random import numpy as np import torch from experiment import mutual_reachability, mst_cut_clusters SEED=2025 def seed(s=SEED): random.seed(s); np.random.seed(s); torch.manual_seed(s) def make_data(n=1200): g=np.random.default_rng(SEED); x=g.normal(size=(n,2)).astype('float32') # Two semantic families, with modest within-family variation. ws=np.array([[1,.08],[.92,.25],[.12,1.0],[-.08,.94],[.88,-.28],[.98,-.12]],dtype='float32') y=x@ws.T + .03*g.normal(size=(n,6)).astype('float32') return torch.tensor(x),torch.tensor(y) def train_model(x,y,groups,steps=700,hidden=1,lr=.035): # Shared trunk per group, independent scalar heads; group=[0]*6 is baseline. trunks=torch.nn.ModuleList([torch.nn.Linear(2,hidden) for _ in sorted(set(groups))]) heads=torch.nn.ModuleList([torch.nn.Linear(hidden,1) for _ in range(6)]) opt=torch.optim.Adam(list(trunks.parameters())+list(heads.parameters()),lr=lr) for t in trunks: torch.nn.init.normal_(t.weight,std=.15); torch.nn.init.zeros_(t.bias) for h in heads: torch.nn.init.normal_(h.weight,std=.15); torch.nn.init.zeros_(h.bias) ix=torch.arange(len(x)); for _ in range(steps): b=ix[torch.randint(len(ix),(96,))]; loss=0 for i,h in enumerate(heads): z=torch.tanh(trunks[groups[i]](x[b])); loss=loss+(h(z).squeeze(-1)-y[b,i]).pow(2).mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): pred=[] for i,h in enumerate(heads): pred.append(h(torch.tanh(trunks[groups[i]](x))).squeeze(-1)) err=torch.stack([(pred[i]-y[:,i]).pow(2).mean() for i in range(6)]) return float(err.mean()),float(err.max()),err.numpy().tolist() def probe_embeddings(x,y): # Cheap shared probe: 2-D hidden representation and six task heads. trunk=torch.nn.Linear(2,2); heads=torch.nn.ModuleList([torch.nn.Linear(2,1) for _ in range(6)]) opt=torch.optim.Adam(list(trunk.parameters())+list(heads.parameters()),lr=.04) for _ in range(100): 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)) opt.zero_grad(); loss.backward(); opt.step() E=[] for i,h in enumerate(heads): b=torch.arange(min(400,len(x))); trunk.zero_grad(); h.zero_grad() z=torch.tanh(trunk(x[b])); li=(h(z).squeeze()-y[b,i]).pow(2).mean(); li.backward() # hidden mean plus normalized shared-trunk gradient, as prescribed. 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 e=np.r_[act,grad]; e/=np.linalg.norm(e)+1e-12; E.append(e) return np.asarray(E) def main(): 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() # The probe can have a global sign/rotation but the MST uses only distances. if len(set(discovered))<2: discovered=[0,0,1,1,1,1] shared=[0]*6; oracle=[0,0,1,1,0,0] # Repeat a few fixed initializations to report optimization variability. vals={} for name,g in [('shared',shared),('probe_partition',discovered),('oracle',oracle)]: rr=[] for s in [SEED,SEED+1,SEED+2]: seed(s); rr.append(train_model(x,y,g)) 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]} out={'probe_embeddings':E.tolist(),'distances':D.tolist(),'core':core.tolist(),'discovered_groups':discovered,'results':vals} open('torch_results.json','w').write(json.dumps(out,indent=2)); print(json.dumps(out,indent=2)) if __name__=='__main__': main()