Schäffer-Covariant Isometric Recurrent Layer / schaffer_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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))