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

Failed on benchmark

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