import json, math import numpy as np from scipy.linalg import eigh from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score SEED = 2340 rng = np.random.default_rng(SEED) def equicorr(p, c): return (1-c)*np.eye(p) + c*np.ones((p,p)) def invsqrt(R, eps=1e-10): w, U = eigh(R) return (U * (1.0/np.sqrt(np.maximum(w, eps))) ) @ U.T def sqrtm(R): w, U = eigh(R) return (U * np.sqrt(np.maximum(w, 0))) @ U.T def anchored_whitener(R, Q=None, eps=1e-10): if Q is None: Q = np.eye(R.shape[0]) return invsqrt(R, eps) @ Q def metrics(R, T): C = T.T @ R @ T fidelity = np.diag(R @ T) return float(np.max(np.abs(C-np.eye(R.shape[0])))), fidelity, float(np.linalg.norm(C-np.eye(R.shape[0]), 'fro')) def analytic_identity_fidelity(p, c): # diag(sqrt(R)) for equicorrelation R return (math.sqrt(1+(p-1)*c) + (p-1)*math.sqrt(1-c))/p def run_math_sweeps(): p = 6 # Prediction 1: exact covariance residual is numerical precision and does not depend on c. decor = [] for c in [0.0, .1, .3, .6, .9]: R = equicorr(p,c); T = anchored_whitener(R) d, f, fn = metrics(R,T) decor.append({'c':c, 'max_cov_error':d, 'fidelity_min':float(f.min())}) # Prediction 2: Q=I fidelity follows the closed form, measured over correlation sweep. fidelity = [] for c in [0.0, .1, .3, .6, .9]: R = equicorr(p,c); T = anchored_whitener(R) observed = float(metrics(R,T)[1].min()) predicted = analytic_identity_fidelity(p,c) fidelity.append({'c':c, 'predicted':predicted, 'observed':observed, 'abs_error':abs(predicted-observed)}) # Prediction 3: for an identity-initialized anchored layer, the declared threshold # is satisfied exactly up to rho0(c)=diag(sqrt(R)); above it, the initial anchor violates it. transition=[] for c in [0.0, .1, .3, .6, .9]: rho0=analytic_identity_fidelity(p,c) for rho in [rho0-.02, rho0+.02]: R=equicorr(p,c); T=anchored_whitener(R) minf=float(metrics(R,T)[1].min()) transition.append({'c':c,'rho':rho,'predicted_feasible':bool(rho<=rho0), 'observed_feasible':bool(minf+1e-8>=rho), 'rho0_predicted':rho0,'min_fidelity':minf}) return decor, fidelity, transition def run_toy_classification(): # Same correlated standardized input for all transforms; labels depend on a rotated # signal. This compares a standard per-channel normalization to exact ZCA and anchored ZCA. n, p = 5000, 6 R = equicorr(p, .6) X = rng.multivariate_normal(np.zeros(p), R, size=n) w = np.array([1.0, -0.8, .5, 0, 0, 0]) y = (X @ w + .35*rng.normal(size=n) > 0).astype(int) Xtr, Xte, ytr, yte = train_test_split(X,y,test_size=.35,random_state=SEED,stratify=y) # BN-like feature standardization (using training statistics only). mu=Xtr.mean(0); sd=Xtr.std(0); Xbn=(Xtr-mu)/sd; Xbnte=(Xte-mu)/sd Rhat=np.cov(Xbn,rowvar=False,bias=True) T=anchored_whitener(Rhat) Xzca=Xbn@T; Xzcat=Xbnte@T # Anchor Q=I is the identity-preserving member of the exact whitening family. def fit_acc(A, At): clf=LogisticRegression(C=1e3,max_iter=300,random_state=SEED) clf.fit(A,ytr); return accuracy_score(yte,clf.predict(At)) _, fbn, offbn = metrics(Rhat,np.eye(p)) err_zca, fz, offz = metrics(Rhat,T) return { 'accuracy_bn_like':fit_acc(Xbn,Xbnte), 'accuracy_zca_q_identity':fit_acc(Xzca,Xzcat), 'bn_like_offdiag_fro':float(np.linalg.norm(Rhat-np.diag(np.diag(Rhat)),'fro')), 'zca_offdiag_fro':float(np.linalg.norm(T.T@Rhat@T-np.diag(np.diag(T.T@Rhat@T)),'fro')), 'zca_max_cov_error':err_zca, 'zca_min_input_fidelity':float(fz.min()), 'rho_identity_prediction':analytic_identity_fidelity(p,.6), 'n_train':len(Xtr) } if __name__ == '__main__': decor, fidelity, transition = run_math_sweeps() toy = run_toy_classification() out={'seed':SEED,'predictions':{ 'P1_exact_decorrelation':'max |T^T R T-I| should remain at floating point precision for all c and Q=I', 'P2_identity_fidelity':'min fidelity should equal [sqrt(1+(p-1)c)+(p-1)sqrt(1-c)]/p', 'P3_threshold_transition':'identity anchor is feasible iff rho <= rho0(c), where rho0 is the P2 curve'}, 'decorrelation_sweep':decor,'fidelity_sweep':fidelity,'threshold_sweep':transition,'toy':toy} print(json.dumps(out,indent=2))