Schäffer-Covariant Isometric Recurrent Layer / schaffer_experiment.py
Mechanism confirmed, baseline not beaten
1import json, random
2from pathlib import Path
3import numpy as np
4
5
6def defect(X, jitter=1e-12):
7 d=X.shape[0]; A=np.eye(d)-X.T@X; w,U=np.linalg.eigh((A+A.T)/2)
8 return (U*np.sqrt(np.maximum(w,0)+jitter))@U.T
9
10
11def schaffer_matrix(X,K):
12 d=X.shape[0]; D=defect(X); Y=np.zeros(((K+1)*d,(K+1)*d))
13 Y[:d,:d]=X; Y[d:2*d,:d]=D
14 for j in range(1,K): Y[(j+1)*d:(j+2)*d,j*d:(j+1)*d]=np.eye(d)
15 return Y,D
16
17
18def covariance_residual(V1,V2,m=2,eps=1e-12):
19 lhs=V1@V2; rhs=V2@np.linalg.matrix_power(V1,m)
20 return float(np.linalg.norm(lhs-rhs,'fro')**2/(np.linalg.norm(lhs,'fro')**2+eps))
21
22
23def math_check(rng):
24 d,K=7,10; A=rng.normal(size=(d,d)); X=.82*A/np.linalg.svd(A,compute_uv=False)[0]
25 Y,D=schaffer_matrix(X,K); n=(K+1)*d; z=rng.normal(size=n); z[-d:]=0
26 active_err=abs(np.linalg.norm(Y@z)**2-np.linalg.norm(z)**2)
27 z=np.r_[rng.normal(size=d),np.zeros(K*d)]; energies=[]
28 for _ in range(K): energies.append(np.linalg.norm(z)**2); z=Y@z
29 return {'active_one_step_energy_error':float(active_err),
30 'finite_full_YtY_fro_error_due_to_truncation':float(np.linalg.norm(Y.T@Y-np.eye(n),'fro')),
31 'pre_boundary_energy_relative_drift':float((max(energies)-min(energies))/energies[0]),
32 'defect_identity_error':float(np.linalg.norm(D.T@D-(np.eye(d)-X.T@X))),
33 'covariance_random_pair_m2':covariance_residual(Y,np.eye(n),2)}
34
35
36def dynamics_check(rng):
37 d,K=16,32; A=rng.normal(size=(d,d)); X=.97*A/np.linalg.svd(A,compute_uv=False)[0]
38 Y,_=schaffer_matrix(X,K); short=K-2; long=160; trials=100; out=[[],[],[],[],[]]
39 for _ in range(trials):
40 h=rng.normal(size=d); z=np.r_[h,np.zeros(K*d)]; h0=h.copy()
41 for _ in range(short): h=X@h; z=Y@z
42 out[0].append(np.linalg.norm(h)/np.linalg.norm(h0)); out[1].append(np.linalg.norm(z)/np.linalg.norm(z0:=np.r_[h0,np.zeros(K*d)]))
43 h=h0.copy(); z=z0.copy()
44 for _ in range(long): h=X@h; z=Y@z
45 out[2].append(np.linalg.norm(h)/np.linalg.norm(h0)); out[3].append(np.linalg.norm(z)/np.linalg.norm(z0)); out[4].append(np.linalg.norm(z[-d:])**2/np.linalg.norm(z)**2)
46 return {'baseline_h_ratio_short_mean':float(np.mean(out[0])),'lift_total_ratio_short_mean':float(np.mean(out[1])),
47 'baseline_h_ratio_long_mean':float(np.mean(out[2])),'lift_total_ratio_long_mean':float(np.mean(out[3])),
48 'lift_long_last_slot_energy_fraction':float(np.mean(out[4]))}
49
50
51def torch_mini_experiment(seed=123):
52 """Paired-initialization delayed-signal task. q is deliberately not read out: this tests task neutrality."""
53 try:
54 import torch
55 torch.manual_seed(seed); device='cuda' if torch.cuda.is_available() else 'cpu'
56 if device=='cuda':
57 try: torch.cuda.empty_cache()
58 except Exception: device='cpu'
59 except Exception as e: return {'error':str(e)}
60 torch.set_num_threads(4); g=np.random.default_rng(seed); d,T,N=16,30,256
61 x=torch.tensor(g.normal(size=(N,T,1)).astype('float32'),device=device); x[:,1:]=0; y=x[:,0,0]
62 init=[torch.randn(d,d,device=device)*.15,torch.randn(d,1,device=device)*.1,torch.randn(1,d,device=device)*.1]
63 def train(kind):
64 raw=torch.nn.Parameter(init[0].clone()); B=torch.nn.Parameter(init[1].clone()); C=torch.nn.Parameter(init[2].clone())
65 opt=torch.optim.Adam([raw,B,C],lr=.01); first=None
66 for _ in range(250):
67 X=raw/torch.clamp(torch.linalg.matrix_norm(raw,2),min=1.); h=torch.zeros(N,d,device=device)
68 if kind=='lift':
69 I=torch.eye(d,device=device); A=I-X.T@X; w,U=torch.linalg.eigh((A+A.T)/2); D=(U*torch.sqrt(torch.clamp(w,min=1e-7)))@U.T
70 for t in range(T): h=h@X.T+x[:,t]@B.T
71 loss=((h@C.T)[:,0]-y).pow(2).mean(); first=first or float(loss.detach().cpu())
72 opt.zero_grad(); loss.backward(); opt.step()
73 return {'initial_mse':first,'final_mse':float(loss.detach().cpu())}
74 try: return {'baseline':train('baseline'),'lift':train('lift'),'device':device}
75 except Exception as e: return {'error':str(e),'device':device}
76
77if __name__=='__main__':
78 seed=123; np.random.seed(seed); random.seed(seed); rng=np.random.default_rng(seed)
79 result={'math':math_check(rng),'dynamics':dynamics_check(rng),'mini':torch_mini_experiment(seed)}
80 Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))