Probe-Then-Partitioned Multi-Task Trunk / experiment.py
Failed on benchmark
1import json, math, random
2from pathlib import Path
3import numpy as np
4
5SEED=1234
6rng=np.random.default_rng(SEED)
7
8
9def cosine_dist(a,b):
10 a=np.asarray(a,float); b=np.asarray(b,float)
11 return 1-float(a@b)/(np.linalg.norm(a)*np.linalg.norm(b)+1e-12)
12
13
14def mutual_reachability(D,m=2):
15 n=len(D); cores=np.zeros(n)
16 for i in range(n):
17 vals=np.sort(np.delete(D[i],i)); cores[i]=vals[min(m-1,len(vals)-1)]
18 MR=np.maximum(D,np.maximum(cores[:,None],cores[None,:])); np.fill_diagonal(MR,0)
19 return cores,MR
20
21
22def mst_cut_clusters(MR):
23 # Kruskal MST, cut the largest unusually separated edge (density hierarchy toy).
24 n=len(MR); edges=sorted((MR[i,j],i,j) for i in range(n) for j in range(i))
25 par=list(range(n))
26 def find(x):
27 while par[x]!=x: par[x]=par[par[x]]; x=par[x]
28 return x
29 mst=[]
30 for w,i,j in edges:
31 a,b=find(i),find(j)
32 if a!=b: par[a]=b; mst.append((w,i,j))
33 ws=np.array([x[0] for x in mst])
34 if len(ws)<2: return np.arange(n)
35 gaps=np.diff(np.sort(ws)); k=int(np.argmax(gaps))
36 cut=float(np.sort(ws)[k+1])
37 # only split if gap is meaningful; this is the singleton/noise-safe hierarchy cut
38 if cut <= (np.median(ws)+1e-12)*1.15: cut=float('inf')
39 par=list(range(n))
40 for w,i,j in mst:
41 if w<cut:
42 a,b=find(i),find(j)
43 if a!=b: par[a]=b
44 roots={}; out=[]
45 for i in range(n):
46 r=find(i); roots.setdefault(r,len(roots)); out.append(roots[r])
47 return np.array(out)
48
49
50def math_checks():
51 u=np.array([1.,0.,0.]); v=np.array([0.,1.,0.]);
52 exact={"same":cosine_dist(u,u),"orthogonal":cosine_dist(u,v),"antipodal":cosine_dist(u,-u)}
53 # For e=normalize(mu + sigma*z), small-noise prediction E[d] ~= (d-1)sigma^2.
54 d=8; mu=np.zeros(d); mu[0]=1
55 rows=[]
56 for s in [0.01,0.03,0.06,0.10,0.18,0.30]:
57 vals=[]
58 for _ in range(30000):
59 a=mu+s*rng.normal(size=d); b=mu+s*rng.normal(size=d)
60 vals.append(cosine_dist(a,b))
61 obs=float(np.mean(vals)); pred=(d-1)*s*s
62 rows.append({"sigma":s,"observed":obs,"predicted":pred,"ratio":obs/pred})
63 return exact,rows
64
65
66def clustering_check():
67 d=8; centers=np.zeros((2,d)); centers[0,0]=1; centers[1,1]=1
68 rows=[]
69 for s in [0.01,0.05,0.10,0.20,0.35,0.60]:
70 E=[]; labels=[]
71 for c in range(2):
72 for _ in range(8):
73 E.append(centers[c]+s*rng.normal(size=d)); labels.append(c)
74 E=np.asarray(E); E/=np.linalg.norm(E,axis=1,keepdims=True)
75 D=1-E@E.T; np.fill_diagonal(D,0)
76 core,MR=mutual_reachability(D,2); pred=float(np.mean([D[i,j] for i in range(8) for j in range(8,16)]))
77 within=float(np.mean([D[i,j] for i in range(16) for j in range(i) if labels[i]==labels[j]]))
78 # density hierarchy output; match up to cluster label permutation
79 got=mst_cut_clusters(MR)
80 same=(got[:8,None]==got[None,:8]).mean() # diagnostic only
81 # pairwise purity, invariant to numeric cluster labels
82 purity=np.mean([max(np.mean(got[np.array(labels)==c]==q) for q in set(got)) for c in [0,1]])
83 rows.append({"sigma":s,"within_distance":within,"between_distance":pred,"between_minus_within":pred-within,"cluster_purity":float(purity),"mean_core":float(core.mean())})
84 return rows
85
86
87def gradient_embeddings(angles):
88 # At the zero-output probe point, average trunk gradient is proportional to target vector.
89 E=np.array([[math.cos(a),math.sin(a)] for a in angles],float)
90 E/=np.linalg.norm(E,axis=1,keepdims=True)
91 return E
92
93
94def rank1_loss(angles, groups):
95 # Exact population optimum for a scalar shared trunk and scalar task heads.
96 # For each group, residual is the smaller eigenvalue of sum_i w_i w_i^T.
97 total=0.
98 for g in sorted(set(groups)):
99 W=np.array([[math.cos(angles[i]),math.sin(angles[i])] for i in range(len(angles)) if groups[i]==g])
100 ev=np.linalg.eigvalsh(W.T@W)
101 total += float(ev[0])
102 return total/len(angles)
103
104
105def training_check():
106 # Two task families: discovery should recover them from probe gradient embeddings.
107 angles=np.array([0.00,0.10,0.18,0.25,1.30,1.40,1.48,1.58])
108 E=gradient_embeddings(angles); D=1-E@E.T; np.fill_diagonal(D,0)
109 core,MR=mutual_reachability(D,m=2); discovered=mst_cut_clusters(MR)
110 shared=np.zeros(len(angles),dtype=int)
111 oracle=np.array([0]*4+[1]*4)
112 # Compare equal analytic population objective; add finite-sample SGD confirmation.
113 return {"angles":angles.tolist(),"discovered_groups":discovered.tolist(),"probe_core":core.tolist(),
114 "shared_population_loss":rank1_loss(angles,shared),
115 "discovered_population_loss":rank1_loss(angles,discovered),
116 "oracle_population_loss":rank1_loss(angles,oracle),
117 "conflict_sweep":conflict_sweep()}
118
119
120def conflict_sweep():
121 rows=[]
122 for theta in [0,.2,.4,.6,.8,1.0,1.2,math.pi/2]:
123 a=np.array([0.,theta]); shared=rank1_loss(a,[0,0]); part=rank1_loss(a,[0,1]);
124 pred=(1-abs(math.cos(theta)))/2 # per-task average, two unit tasks
125 rows.append({"angle":theta,"shared_loss":shared,"partition_loss":part,"observed_gain":shared-part,"predicted_gain":pred})
126 return rows
127
128
129def main():
130 exact,noise=math_checks(); clusters=clustering_check(); train=training_check()
131 out={"seed":SEED,"math_exact":exact,"noise_scaling":noise,"density_sweep":clusters,"training":train,
132 "predictions":{"P1":"cosine distance is 0/1/2 for same/orthogonal/antipodal normalized embeddings",
133 "P2":"for d=8 and small sigma, E[distance] ~= 7 sigma^2",
134 "P3":"for two unit linear tasks at angle theta, partition gain per task = (1-|cos(theta)|)/2"}}
135 Path('results.json').write_text(json.dumps(out,indent=2))
136 print(json.dumps(out,indent=2))
137
138if __name__=='__main__': main()