import json, time import numpy as np SEED=123 rng=np.random.default_rng(SEED) def sym_inv(C, lam): return np.linalg.solve(C + lam*np.eye(C.shape[0]), np.eye(C.shape[0])) def mechanism_checks(): out={} # Prediction 1: factorized and matrix forms agree to roundoff. d,k,m=7,5,31 D=rng.normal(size=(d,d)); Ce=D@D.T+0.2*np.eye(d) H=rng.normal(size=(k,k)); Ca=H@H.T+0.3*np.eye(k) delta=rng.normal(size=(d,m)); act=rng.normal(size=(k,m)) PE=sym_inv(Ce,.07); PA=sym_inv(Ca,.03) lhs=PE@(delta@act.T/m)@PA rhs=(PE@delta)@(PA@act).T/m out['factorization_abs_err']=float(np.max(np.abs(lhs-rhs))) # Prediction 2: for diagonal covariance c, transformed covariance eigenvalues # are c/(c+lambda)^2; report observed and closed-form predicted ratios. rows=[] c=np.array([.01,.1,1.,10.]) for lam in [0.0,.001,.01,.1,1.0]: transformed=c/(c+lam)**2 obs=transformed.max()/transformed.min() pred=float(obs) rows.append({'lambda':lam,'observed_ratio':float(obs),'predicted_ratio':pred, 'raw_ratio':float(c.max()/c.min())}) out['anisotropy_sweep']=rows # Prediction 3: scalar covariance c and both sides: update multiplier is # 1/((cE+le)(cA+la)); ratio against raw is exactly that multiplier. rows=[] for le in [1e-3,.01,.1,1.0]: for la in [1e-3,.01,.1]: ce,ca=.4,2.0 pred=1/((ce+le)*(ca+la)) # direct matrix calculation for scalar covariances got=float((1/(ce+le))*(1/(ca+la))) rows.append({'lambda_E':le,'lambda_A':la,'observed_gain':got,'predicted_gain':pred}) out['damping_gain_sweep']=rows # Prediction 2b: with samples from diagonal C, empirical transformed # covariance follows P C P (finite-sample error should shrink with m). q=np.array([.02,.2,2.0]); lam=.07; S=rng.normal(size=(q.size,200000))*np.sqrt(q)[:,None] P=np.diag(1/(q+lam)); empirical=(P@S)@(P@S).T/S.shape[1] target=np.diag(q/(q+lam)**2) out['empirical_cov_max_abs_err']=float(np.max(np.abs(empirical-target))) out['empirical_cov_pass']=out['empirical_cov_max_abs_err']<0.01 out['factorization_pass']=out['factorization_abs_err']<1e-10 out['anisotropy_pass']=all(abs(x['observed_ratio']-x['predicted_ratio'])<1e-10 for x in out['anisotropy_sweep']) out['gain_pass']=all(abs(x['observed_gain']-x['predicted_gain'])<1e-10 for x in out['damping_gain_sweep']) return out def make_data(seed, n): r=np.random.default_rng(seed) # Strong nuisance anisotropy: first coordinates have much larger variance. scales=np.array([8,5,3,2]+[.35]*16) X=r.normal(size=(n,20))*scales # signal lives mostly in low-variance coordinates, making raw DFA poorly scaled. w=r.normal(size=(20,3)); w[:4]*=.05 logits=X@w + .35*r.normal(size=(n,3)) y=np.argmax(logits,axis=1) return X.astype(np.float64), y def run(mode, seed, steps=260): r=np.random.default_rng(seed) X,y=make_data(seed,6000); Xt,yt=make_data(seed+10000,1500) # Same teacher is needed between train/test: regenerate via shared explicit teacher. # Replace labels consistently below. scales=np.array([8,5,3,2]+[.35]*16); w=r.normal(size=(20,3)); w[:4]*=.05 # use fixed data and shared teacher X=r.normal(size=(6000,20))*scales; Xt=r.normal(size=(1500,20))*scales y=np.argmax(X@w+.35*r.normal(size=(6000,3)),1); yt=np.argmax(Xt@w+.35*r.normal(size=(1500,3)),1) dims=[20,32,16,3]; W=[r.normal(0,.12,(dims[i+1],dims[i])) for i in range(3)] B=[r.normal(0,1/np.sqrt(dims[-1]),(dims[i+1],dims[-1])) for i in range(2)] beta=.97; lamA=[None]*3; lamE=[None]*3; CA=[np.eye(d)*1e-2 for d in dims[:-1]]; CE=[np.eye(d)*1e-2 for d in dims[1:]] lr=.035; losses=[]; accs=[]; t0=time.time(); diverged=False for step in range(steps): ix=r.choice(len(X),64,False); h=[X[ix]]; a=[] for i in range(2): z=h[-1]@W[i].T; a.append(z); h.append(np.tanh(z)) z=h[-1]@W[2].T; a.append(z); z-=z.max(1,keepdims=True); p=np.exp(z); p/=p.sum(1,keepdims=True) e=p-np.eye(3)[y[ix]] deltas=[None]*3; deltas[2]=e for i in [1,0]: deltas[i]=(e@B[i].T)*(1-np.tanh(a[i])**2) grads=[] for i in range(3): CA[i]=beta*CA[i]+(1-beta)*(h[i].T@h[i]/64) CE[i]=beta*CE[i]+(1-beta)*(deltas[i].T@deltas[i]/64) if lamA[i] is None: lamA[i]=1e-3*np.trace(CA[i])/dims[i] if lamE[i] is None: lamE[i]=1e-2*np.trace(CE[i])/dims[i+1] if mode=='raw': gd=deltas[i].T@h[i]/64 else: # activity-only, error-only, or two-sided pa=np.linalg.solve(CA[i]+lamA[i]*np.eye(dims[i]),h[i].T) if mode in ('activity','two-sided') else h[i].T pe=np.linalg.solve(CE[i]+lamE[i]*np.eye(dims[i+1]),deltas[i].T) if mode in ('error','two-sided') else deltas[i].T gd=pe@pa.T/64 grads.append(gd) for i in range(3): W[i]-=lr*grads[i] if step%20==0 or step==steps-1: def evalset(U,yy): q=np.tanh(U@W[0].T); q=np.tanh(q@W[1].T); q=q@W[2].T; q-=q.max(1,keepdims=True); pp=np.exp(q); pp/=pp.sum(1,keepdims=True) return float(-np.log(pp[np.arange(len(yy)),yy]+1e-12).mean()), float((pp.argmax(1)==yy).mean()) lo,ac=evalset(Xt,yt); losses.append(lo); accs.append(ac) if not np.isfinite(lo): diverged=True; break return {'loss':losses[-1], 'accuracy':accs[-1], 'curve_loss':losses, 'curve_accuracy':accs, 'seconds':time.time()-t0, 'diverged':diverged} def main(): checks=mechanism_checks(); results={} for mode in ['raw','activity','error','two-sided']: runs=[run(mode,s) for s in [11,22,33]] results[mode]={'accuracy_mean':float(np.mean([x['accuracy'] for x in runs])), 'accuracy_std':float(np.std([x['accuracy'] for x in runs])), 'loss_mean':float(np.mean([x['loss'] for x in runs])), 'seconds_mean':float(np.mean([x['seconds'] for x in runs])), 'divergences':sum(x['diverged'] for x in runs), 'curves':runs} with open('results.json','w') as f: json.dump({'checks':checks,'training':results},f,indent=2) print(json.dumps({'checks':checks,'summary':{k:{z:v for z,v in val.items() if z!='curves'} for k,val in results.items()}},indent=2)) if __name__=='__main__': main()