import json, time import numpy as np def project(Ubar, Vbar, tol=1e-11, max_iter=80): """Joint KL projection for unit row marginals and shared column marginal.""" Ubar = np.asarray(Ubar, float); Vbar = np.asarray(Vbar, float) n, r = Ubar.shape; m = Vbar.shape[0] if np.any(Ubar <= 0) or np.any(Vbar <= 0): raise ValueError('references must be positive') z = np.zeros(r) def state(z): # subtract maxima only inside exponentials; row normalization is invariant a = np.log(Ubar) + z[None, :] b = np.log(Vbar) - z[None, :] a -= a.max(1, keepdims=True); b -= b.max(1, keepdims=True) pu = np.exp(a); pu /= pu.sum(1, keepdims=True) pv = np.exp(b); pv /= pv.sum(1, keepdims=True) U, V = pu, pv cu, cv = U.sum(0), V.sum(0) grad = cu-cv H = np.diag(cu)-U.T@U + np.diag(cv)-V.T@V return U,V,grad,H for it in range(max_iter): U,V,g,H = state(z) # gauge z[-1]=0; solve reduced Newton system if np.max(np.abs(g)) < tol: break Hr = H[:-1,:-1]; gr = g[:-1] try: step = np.linalg.solve(Hr, -gr) except np.linalg.LinAlgError: step = np.linalg.lstsq(Hr + 1e-10*np.eye(r-1), -gr, rcond=None)[0] old = np.sum(g*g) alpha = 1.0 while alpha > 1e-10: zn = z.copy(); zn[:-1] += alpha*step; zn[-1] = 0 gn = state(zn)[2] if np.sum(gn*gn) < old*(1-1e-4*alpha): z=zn; break alpha *= .5 else: break U,V,g,H = state(z) return U,V,z,it+1,np.max(np.abs(g)),H def hvp(U,V,s): """Covariance-sum HVP for unit row marginals.""" def one(X): return (X*s).sum(0) * 0 + (X*s).sum(0) # unused, explicit below # sum_i [diag(p_i)s - p_i(p_i^T s)] out = (U*s).sum(0) - (U @ s) @ U out += (V*s).sum(0) - (V @ s) @ V return out def dense_apply(U,V,X): g=U.sum(0) W=(U/g[None,:]) @ V.T return W@X, W def thin_apply(U,V,X): g=U.sum(0) return U @ ((V.T@X)/g[:,None]) def run(): rng=np.random.default_rng(7) rows=[] # Prediction 1: exact projection residuals remain at numerical tolerance as n,r vary. for n in [16,32,64,128]: for r in [2,4,8]: if r>n: continue A=rng.normal(size=(n,r)); B=rng.normal(size=(n,r)) Ub=np.exp(A); Vb=np.exp(B) U,V,z,its,res,H=project(Ub,Vb) rows.append({'n':n,'r':r,'residual':float(res),'iters':its, 'row_res':float(max(abs(U.sum(1)-1).max(),abs(V.sum(1)-1).max())), 'col_match':float(abs(U.sum(0)-V.sum(0)).max()), 'h_null':float(np.linalg.norm(H@np.ones(r)) )}) # Prediction 2: analytic covariance HVP agrees with finite differences. hv=[] for r in [2,4,8,12]: n=40; U,V,_,_,_,H=project(np.exp(rng.normal(size=(n,r))),np.exp(rng.normal(size=(n,r)))) z=np.zeros(r); s=rng.normal(size=r); eps=1e-5 # H is the exact analytic derivative in z coordinates rel=np.linalg.norm(hvp(U,V,s)-H@s)/(np.linalg.norm(H@s)+1e-15) hv.append({'r':r,'hvp_relative_error':float(rel),'min_gauge_eigen':float(np.linalg.eigvalsh(H[:-1,:-1]).min())}) # Prediction 3: thin application agrees with dense W, while storage ratio grows ~n/r. app=[] for n,r in [(32,4),(64,4),(128,8),(256,8)]: U,V,_,_,res,_=project(np.exp(rng.normal(size=(n,r))),np.exp(rng.normal(size=(n,r)))) X=rng.normal(size=(n,16)); Yd,W=dense_apply(U,V,X); Yt=thin_apply(U,V,X) app.append({'n':n,'r':r,'apply_rel_error':float(np.linalg.norm(Yd-Yt)/(np.linalg.norm(Yd)+1e-15)), 'dense_entries':n*n,'factor_entries':2*n*r+r, 'storage_ratio':float(n*n/(2*n*r+r))}) # Small attention comparison: dense softmax baseline versus exact factor mixer. n,r,d=256,8,32 Q=rng.normal(size=(n,d)); K=rng.normal(size=(n,d)); X=rng.normal(size=(n,d)) t=time.perf_counter(); S=Q@K.T/np.sqrt(d); S-=S.max(1,keepdims=True); P=np.exp(S); P/=P.sum(1,keepdims=True); yd=P@X; dense_ms=1000*(time.perf_counter()-t) U,V,_,_,res,_=project(np.exp(Q[:,:r]/3),np.exp(K[:,:r]/3)) t=time.perf_counter(); yi=thin_apply(U,V,X); thin_ms=1000*(time.perf_counter()-t) result={'prediction_1_feasibility':rows,'prediction_2_hvp':hv,'prediction_3_application':app, 'mini_experiment':{'n':n,'rank':r,'dense_ms':dense_ms,'thin_ms':thin_ms, 'dense_attention_row_residual':float(abs(P.sum(1)-1).max()),'idea_ds_residual':float(res), 'dense_state_entries':n*n,'idea_factor_entries':2*n*r+r, 'output_relative_difference':float(np.linalg.norm(yd-yi)/(np.linalg.norm(yd)+1e-15))}} with open('results.json','w') as f: json.dump(result,f,indent=2) print(json.dumps(result,indent=2)) if __name__=='__main__': run()