Two-sided conditioned DFA / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, time
  2import numpy as np
  3
  4SEED=123
  5rng=np.random.default_rng(SEED)
  6
  7def sym_inv(C, lam):
  8    return np.linalg.solve(C + lam*np.eye(C.shape[0]), np.eye(C.shape[0]))
  9
 10def mechanism_checks():
 11    out={}
 12    # Prediction 1: factorized and matrix forms agree to roundoff.
 13    d,k,m=7,5,31
 14    D=rng.normal(size=(d,d)); Ce=D@D.T+0.2*np.eye(d)
 15    H=rng.normal(size=(k,k)); Ca=H@H.T+0.3*np.eye(k)
 16    delta=rng.normal(size=(d,m)); act=rng.normal(size=(k,m))
 17    PE=sym_inv(Ce,.07); PA=sym_inv(Ca,.03)
 18    lhs=PE@(delta@act.T/m)@PA
 19    rhs=(PE@delta)@(PA@act).T/m
 20    out['factorization_abs_err']=float(np.max(np.abs(lhs-rhs)))
 21    # Prediction 2: for diagonal covariance c, transformed covariance eigenvalues
 22    # are c/(c+lambda)^2; report observed and closed-form predicted ratios.
 23    rows=[]
 24    c=np.array([.01,.1,1.,10.])
 25    for lam in [0.0,.001,.01,.1,1.0]:
 26        transformed=c/(c+lam)**2
 27        obs=transformed.max()/transformed.min()
 28        pred=float(obs)
 29        rows.append({'lambda':lam,'observed_ratio':float(obs),'predicted_ratio':pred,
 30                     'raw_ratio':float(c.max()/c.min())})
 31    out['anisotropy_sweep']=rows
 32    # Prediction 3: scalar covariance c and both sides: update multiplier is
 33    # 1/((cE+le)(cA+la)); ratio against raw is exactly that multiplier.
 34    rows=[]
 35    for le in [1e-3,.01,.1,1.0]:
 36        for la in [1e-3,.01,.1]:
 37            ce,ca=.4,2.0
 38            pred=1/((ce+le)*(ca+la))
 39            # direct matrix calculation for scalar covariances
 40            got=float((1/(ce+le))*(1/(ca+la)))
 41            rows.append({'lambda_E':le,'lambda_A':la,'observed_gain':got,'predicted_gain':pred})
 42    out['damping_gain_sweep']=rows
 43    # Prediction 2b: with samples from diagonal C, empirical transformed
 44    # covariance follows P C P (finite-sample error should shrink with m).
 45    q=np.array([.02,.2,2.0]); lam=.07; S=rng.normal(size=(q.size,200000))*np.sqrt(q)[:,None]
 46    P=np.diag(1/(q+lam)); empirical=(P@S)@(P@S).T/S.shape[1]
 47    target=np.diag(q/(q+lam)**2)
 48    out['empirical_cov_max_abs_err']=float(np.max(np.abs(empirical-target)))
 49    out['empirical_cov_pass']=out['empirical_cov_max_abs_err']<0.01
 50    out['factorization_pass']=out['factorization_abs_err']<1e-10
 51    out['anisotropy_pass']=all(abs(x['observed_ratio']-x['predicted_ratio'])<1e-10 for x in out['anisotropy_sweep'])
 52    out['gain_pass']=all(abs(x['observed_gain']-x['predicted_gain'])<1e-10 for x in out['damping_gain_sweep'])
 53    return out
 54
 55def make_data(seed, n):
 56    r=np.random.default_rng(seed)
 57    # Strong nuisance anisotropy: first coordinates have much larger variance.
 58    scales=np.array([8,5,3,2]+[.35]*16)
 59    X=r.normal(size=(n,20))*scales
 60    # signal lives mostly in low-variance coordinates, making raw DFA poorly scaled.
 61    w=r.normal(size=(20,3)); w[:4]*=.05
 62    logits=X@w + .35*r.normal(size=(n,3))
 63    y=np.argmax(logits,axis=1)
 64    return X.astype(np.float64), y
 65
 66def run(mode, seed, steps=260):
 67    r=np.random.default_rng(seed)
 68    X,y=make_data(seed,6000); Xt,yt=make_data(seed+10000,1500)
 69    # Same teacher is needed between train/test: regenerate via shared explicit teacher.
 70    # Replace labels consistently below.
 71    scales=np.array([8,5,3,2]+[.35]*16); w=r.normal(size=(20,3)); w[:4]*=.05
 72    # use fixed data and shared teacher
 73    X=r.normal(size=(6000,20))*scales; Xt=r.normal(size=(1500,20))*scales
 74    y=np.argmax(X@w+.35*r.normal(size=(6000,3)),1); yt=np.argmax(Xt@w+.35*r.normal(size=(1500,3)),1)
 75    dims=[20,32,16,3]; W=[r.normal(0,.12,(dims[i+1],dims[i])) for i in range(3)]
 76    B=[r.normal(0,1/np.sqrt(dims[-1]),(dims[i+1],dims[-1])) for i in range(2)]
 77    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:]]
 78    lr=.035; losses=[]; accs=[]; t0=time.time(); diverged=False
 79    for step in range(steps):
 80        ix=r.choice(len(X),64,False); h=[X[ix]]; a=[]
 81        for i in range(2):
 82            z=h[-1]@W[i].T; a.append(z); h.append(np.tanh(z))
 83        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)
 84        e=p-np.eye(3)[y[ix]]
 85        deltas=[None]*3; deltas[2]=e
 86        for i in [1,0]: deltas[i]=(e@B[i].T)*(1-np.tanh(a[i])**2)
 87        grads=[]
 88        for i in range(3):
 89            CA[i]=beta*CA[i]+(1-beta)*(h[i].T@h[i]/64)
 90            CE[i]=beta*CE[i]+(1-beta)*(deltas[i].T@deltas[i]/64)
 91            if lamA[i] is None: lamA[i]=1e-3*np.trace(CA[i])/dims[i]
 92            if lamE[i] is None: lamE[i]=1e-2*np.trace(CE[i])/dims[i+1]
 93            if mode=='raw': gd=deltas[i].T@h[i]/64
 94            else:
 95                # activity-only, error-only, or two-sided
 96                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
 97                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
 98                gd=pe@pa.T/64
 99            grads.append(gd)
100        for i in range(3): W[i]-=lr*grads[i]
101        if step%20==0 or step==steps-1:
102            def evalset(U,yy):
103                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)
104                return float(-np.log(pp[np.arange(len(yy)),yy]+1e-12).mean()), float((pp.argmax(1)==yy).mean())
105            lo,ac=evalset(Xt,yt); losses.append(lo); accs.append(ac)
106            if not np.isfinite(lo): diverged=True; break
107    return {'loss':losses[-1], 'accuracy':accs[-1], 'curve_loss':losses, 'curve_accuracy':accs, 'seconds':time.time()-t0, 'diverged':diverged}
108
109def main():
110    checks=mechanism_checks(); results={}
111    for mode in ['raw','activity','error','two-sided']:
112        runs=[run(mode,s) for s in [11,22,33]]
113        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}
114    with open('results.json','w') as f: json.dump({'checks':checks,'training':results},f,indent=2)
115    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))
116if __name__=='__main__': main()