Exact doubly stochastic low-rank attention / verify.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, time
 2import numpy as np
 3from exact_ds import project, thin_apply, dense_apply
 4
 5
 6def kl(X, X0):
 7    return float(np.sum(X*np.log(X/X0)-X+X0))
 8
 9
10def main():
11    rng = np.random.default_rng(7)
12    feasibility=[]
13    for n in [16, 32, 64, 128, 256]:
14        for r in [2, 4, 8, 16]:
15            if r > n: continue
16            ub=np.exp(rng.normal(size=(n,r))); vb=np.exp(rng.normal(size=(n,r)))
17            U,V,z,it,res,H=project(ub,vb)
18            feasibility.append({
19                'n':n,'r':r,'iterations':it,
20                'max_row_residual':float(max(np.max(abs(U.sum(1)-1)),np.max(abs(V.sum(1)-1)))),
21                'max_shared_column_residual':float(np.max(abs(U.sum(0)-V.sum(0)))),
22                'projected_KL':kl(U,ub)+kl(V,vb),
23                'reference_row_normalized_KL':kl(ub/ub.sum(1,keepdims=True),ub)+kl(vb/vb.sum(1,keepdims=True),vb),
24                'hessian_null_norm':float(np.linalg.norm(H@np.ones(r)))})
25    hvp=[]
26    for r in [2,4,8,12,16]:
27        n=60; ub=np.exp(rng.normal(size=(n,r))); vb=np.exp(rng.normal(size=(n,r)))
28        U,V,_,_,_,H=project(ub,vb); s=rng.normal(size=r)
29        exact=H@s
30        got=(U*s).sum(0)-(U@s)@U + (V*s).sum(0)-(V@s)@V
31        hvp.append({'r':r,'predicted_relative_error':0.0,'observed_relative_error':float(np.linalg.norm(got-exact)/(np.linalg.norm(exact)+1e-15)),
32                    'predicted_positive_gauge_curvature':'> 0','observed_min_reduced_eigenvalue':float(np.linalg.eigvalsh(H[:-1,:-1]).min())})
33    application=[]
34    for n,r in [(32,4),(64,4),(128,8),(256,8),(512,16)]:
35        U,V,_,_,_,_=project(np.exp(rng.normal(size=(n,r))),np.exp(rng.normal(size=(n,r))))
36        X=rng.normal(size=(n,32)); yd,W=dense_apply(U,V,X); yt=thin_apply(U,V,X)
37        expected=n*n/(2*n*r+r)
38        application.append({'n':n,'r':r,'predicted_output_relative_error':0.0,
39          'observed_output_relative_error':float(np.linalg.norm(yd-yt)/(np.linalg.norm(yd)+1e-15)),
40          'predicted_storage_ratio':expected,'observed_storage_ratio':float(n*n/(2*n*r+r)),
41          'dense_entries':n*n,'factor_entries':2*n*r+r})
42    # Application-only timing, repeated to reduce timer noise.
43    n,r,d=512,16,32
44    U,V,_,_,res,_=project(np.exp(rng.normal(size=(n,r))),np.exp(rng.normal(size=(n,r))))
45    X=rng.normal(size=(n,d))
46    dense_apply(U,V,X); thin_apply(U,V,X)
47    def measure(fn):
48        ts=[]
49        for _ in range(20):
50            t=time.perf_counter(); fn(); ts.append((time.perf_counter()-t)*1000)
51        return float(np.median(ts))
52    dense_ms=measure(lambda: dense_apply(U,V,X)[0]); thin_ms=measure(lambda: thin_apply(U,V,X))
53    out={'prediction_1':{'claim':'residuals are numerical-zero independent of n,r','sweep':feasibility},
54         'prediction_2':{'claim':'covariance-sum HVP is exact and reduced curvature is positive','sweep':hvp},
55         'prediction_3':{'claim':'factor storage scales as n^2/(2nr+r), and thin equals dense factor application','sweep':application},
56         'mini_experiment':{'n':n,'r':r,'d':d,'dense_apply_median_ms':dense_ms,'thin_apply_median_ms':thin_ms,
57           'speedup':dense_ms/thin_ms,'projection_shared_column_residual':float(res),
58           'dense_attention_state_entries':n*n,'idea_factor_entries':2*n*r+r}}
59    with open('results.json','w') as f: json.dump(out,f,indent=2)
60    print(json.dumps(out,indent=2))
61
62if __name__=='__main__': main()