KL-TopK Activation Bottleneck / kl_topk_experiment.py

Failed on benchmark

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