KL-TopK Activation Bottleneck / kl_topk_experiment.py
Failed on benchmark
1import json, math, os
2import numpy as np
3
4SEED = 945
5rng = np.random.default_rng(SEED)
6
7def topk_residual(z, d):
8 p = z.shape[-1]
9 # Keep exactly d largest magnitudes, without sorting all coordinates.
10 idx = np.argpartition(z*z, p-d, axis=-1)[..., :p-d]
11 # residual is the sum of the p-d smallest squared coordinates
12 return np.take_along_axis(z*z, idx, axis=-1).sum(axis=-1)
13
14def orthogonal(p, r):
15 a = r.normal(size=(p,p))
16 q, rr = np.linalg.qr(a)
17 q *= np.sign(np.diag(rr))[None, :]
18 return q
19
20def sample_diag(lam, n, r):
21 return r.normal(size=(n,len(lam))) * np.sqrt(lam)[None,:]
22
23def estimate(lam, d, U, n=250000, seed=0):
24 r = np.random.default_rng(seed)
25 x = sample_diag(np.asarray(lam), n, r)
26 # U acts on column vectors; rows therefore multiply U.T.
27 return float(topk_residual(x @ U.T, d).mean()), float(topk_residual(x, d).mean())
28
29def ci_se(lam, d, U, n=100000, seed=1):
30 r = np.random.default_rng(seed)
31 x = sample_diag(np.asarray(lam), n, r)
32 a = topk_residual(x @ U.T, d)
33 return float(a.mean()), float(a.std(ddof=1)/math.sqrt(n))
34
35def main():
36 out = {'seed': SEED, 'predictions': {}, 'sweeps': {}}
37 # Prediction 1: multiplying covariance by c multiplies expected residual by c.
38 lam = np.array([9., 4., 1., .25, .09, .01])
39 d = 2
40 I = np.eye(len(lam))
41 vals = []
42 base = None
43 for c in [0.25, 0.5, 1., 2., 4.]:
44 v,se = ci_se(c*lam, d, I, 160000, 10)
45 if c == 1.0: base = v
46 vals.append({'scale':c, 'residual':v, 'ratio_to_unit_scale':v/base if base is not None else None, 'predicted_ratio':c})
47 unit = next(a['residual'] for a in vals if a['scale'] == 1.0)
48 for a in vals: a['ratio_to_unit_scale'] = a['residual'] / unit
49 out['predictions']['homogeneity_in_covariance_scale'] = vals
50 # Prediction 2: isotropic Gaussian is rotation invariant.
51 p=8; di=3; iso=np.ones(p)
52 rots=[np.eye(p)] + [orthogonal(p, np.random.default_rng(100+i)) for i in range(5)]
53 iso_vals=[]
54 for i,u in enumerate(rots):
55 v,se=ci_se(iso,di,u,220000,20+i)
56 iso_vals.append({'rotation':i,'residual':v,'se':se})
57 out['predictions']['isotropic_rotation_invariance'] = iso_vals
58 # Prediction 3: PCA basis should be no worse than tested rotations; inspect gap vs d.
59 lam=np.array([16.,9.,4.,2.,1.,.5,.2,.05])
60 rotation_bank=[orthogonal(len(lam), np.random.default_rng(500+i)) for i in range(12)]
61 gap=[]
62 for dd in [1,2,3,4,5,6,7]:
63 pca,se=ci_se(lam,dd,np.eye(len(lam)),220000,300+dd)
64 rs=[]
65 for j,u in enumerate(rotation_bank):
66 v,_=ci_se(lam,dd,u,70000,700+10*dd+j)
67 rs.append(v)
68 gap.append({'d':dd,'pca_residual':pca,'random_rotation_mean':float(np.mean(rs)),
69 'random_rotation_min':float(np.min(rs)), 'gap_mean':float(np.mean(rs)-pca),
70 'gap_fraction_of_pca':float((np.mean(rs)-pca)/pca)})
71 out['predictions']['anisotropic_rotation_gap'] = gap
72 # Mini experiment: finite calibration PCA versus random rotation on held-out activations.
73 p=32; dlist=[3,8,16]; ncal=3000; ntest=100000
74 lam=np.geomspace(12,.08,p)
75 rr=np.random.default_rng(808)
76 cal=sample_diag(lam,ncal,rr); test=sample_diag(lam,ntest,rr)
77 # Covariance eigensystem from centered calibration; rows use H Q.
78 C=np.cov(cal,rowvar=False,bias=True)
79 ev,Q=np.linalg.eigh(C); Q=Q[:,np.argsort(ev)[::-1]]
80 # Centering is zero in this controlled model, but use estimated mean as implementation does.
81 mu=cal.mean(0); R=orthogonal(p,np.random.default_rng(809))
82 mini=[]
83 for dd in dlist:
84 z_p=(test-mu)@Q
85 z_r=(test-mu)@R.T
86 vp=float(topk_residual(z_p,dd).mean()); vr=float(topk_residual(z_r,dd).mean())
87 total=float(np.sum(lam)); mini.append({'d':dd,'d_over_p':dd/p,'pca_mse':vp,'random_rotation_mse':vr,
88 'pca_retained_fraction':1-vp/total,'random_retained_fraction':1-vr/total,
89 'stored_value_fraction':dd/p})
90 out['mini_experiment']={'covariance_eigenbasis_vs_random_rotation':mini,
91 'note':'MSE is exact top-d reconstruction residual; stored value fraction is d/p (indices/metadata excluded).'}
92 # Numeric implementation sanity: reconstruct and compare direct residual.
93 z=sample_diag(lam,1000,np.random.default_rng(901))@Q
94 direct=topk_residual(z,8); mask=np.zeros_like(z); ii=np.argpartition(z*z,p-8,axis=1)[:,p-8:]
95 # retain the d largest coordinates
96 mask[np.arange(len(z))[:,None],ii]=z[np.arange(len(z))[:,None],ii]
97 impl=np.sum((z-mask)**2,axis=1)
98 out['math_sanity']={'max_abs_residual_formula_error':float(np.max(np.abs(direct-impl))),
99 'orthogonality_error':float(np.max(np.abs(Q.T@Q-np.eye(p))))}
100 with open('results.json','w') as f: json.dump(out,f,indent=2)
101 print(json.dumps(out,indent=2))
102
103if __name__=='__main__': main()