Critical-Gain Covariance Controller / experiment.py
Mechanism failed
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 1628
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9
10def jacobian(kE, kI):
11 return np.array([[2*kE-2, 0, 2*kI],
12 [0, 2*kI-2, 2*kE],
13 [kE, kI, kE+kI-2]], dtype=float)
14
15def predicted_eigs(kE, kI):
16 d = kE-kI
17 return np.array([-2., 2*(d-1), d-2])
18
19def fixed_point(kE, kI, D):
20 return -np.linalg.solve(jacobian(kE,kI), np.asarray(D, dtype=float))
21
22def toy_checks():
23 # The supplied J is checked directly. Its characteristic polynomial is
24 # (lambda+2)(lambda-2(kE+kI-1))(lambda-(kE+kI-2)).
25 max_stated_err = 0.; max_corrected_err = 0.; boundary = []
26 for kI in np.linspace(.05,.8,8):
27 for d in np.linspace(-.8,1.25,18):
28 kE=d+kI; vals=np.sort(np.linalg.eigvals(jacobian(kE,kI)).real)
29 stated=np.sort(predicted_eigs(kE,kI))
30 s=kE+kI
31 corrected=np.sort(np.array([-2.,2*(s-1),s-2]))
32 max_stated_err=max(max_stated_err,float(np.max(np.abs(vals-stated))))
33 max_corrected_err=max(max_corrected_err,float(np.max(np.abs(vals-corrected))))
34 boundary.append((s,float(np.max(vals))))
35 # Continuous-time stability transition predicted at kE+kI=1, tested by sweep.
36 stable_s=sorted(set(round(s,10) for s,m in boundary if m < -1e-8))
37 unstable_s=sorted(set(round(s,10) for s,m in boundary if m > 1e-8))
38 # Measure slow decay from a pure dominant eigenmode, avoiding mixtures.
39 decay=[]; dt=.001
40 for s in [.20,.40,.60,.75,.85,.92,.97]:
41 kE,kI=s*.6,s*.4; J=jacobian(kE,kI)
42 vals,vecs=np.linalg.eig(J); ix=np.argmax(vals.real)
43 x=vecs[:,ix].real; x=x/np.linalg.norm(x); norms=[]
44 for _ in range(3000):
45 norms.append(np.linalg.norm(x)); x=x+dt*J@x
46 slope=np.polyfit(np.arange(500,2500)*dt,np.log(norms[500:2500]),1)[0]
47 decay.append([s,-float(slope),2*(1-s)])
48 # Fixed point scales as (1-s)^-1 when D has projection on the critical mode.
49 fp=[]; D=np.array([.002,.001,0.])
50 for s in [.50,.60,.70,.78,.84,.88,.91,.93,.95,.96]:
51 kE,kI=.6*s,.4*s; c=fixed_point(kE,kI,D); fp.append([1-s,float(np.linalg.norm(c))])
52 slope=float(np.polyfit(np.log([q[0] for q in fp]),np.log([q[1] for q in fp]),1)[0])
53 return {'stated_delta_eigenvalue_max_abs_error':max_stated_err,
54 'corrected_sum_eigenvalue_max_abs_error':max_corrected_err,
55 'claimed_boundary_delta':1.0,'actual_boundary_sum_kE_plus_kI':1.0,
56 'stable_sweep_max_sum':max(stable_s),'unstable_sweep_min_sum':min(unstable_s),
57 'decay_rate_rows_sum_measured_predicted':decay,
58 'fixed_point_rows_margin_norm':fp,
59 'fixed_point_loglog_slope_measured_predicted':[slope,-1.0]}
60
61class TinyRNN(nn.Module):
62 def __init__(self, vocab=8, h=24):
63 super().__init__(); self.h=h; self.emb=nn.Embedding(vocab,h)
64 self.W=nn.Parameter(torch.empty(h,h)); self.b=nn.Parameter(torch.zeros(h)); self.out=nn.Linear(h,vocab)
65 nn.init.normal_(self.W,0,.55); nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)
66 def forward(self,x, controller=False, alpha=3., eps=.05, tau=.8):
67 B,T=x.shape; h=torch.zeros(B,self.h,device=x.device); hs=[]
68 for t in range(T):
69 z=self.emb(x[:,t])+h@self.W.T+self.b; h=torch.tanh(z); hs.append(h)
70 H=torch.stack(hs,1); logits=self.out(H)
71 d=self.h//2
72 # block spectral gains; singular values are differentiable and inexpensive at this size
73 kE=torch.linalg.matrix_norm(self.W[:d,:d],ord=2); kI=torch.linalg.matrix_norm(self.W[d:,d:],ord=2)
74 delta=kE-kI
75 cov=H.reshape(-1,self.h)-H.reshape(-1,self.h).mean(0,keepdim=True)
76 tr=(cov.square().sum()/max(cov.shape[0]-1,1))
77 penalty=alpha*torch.relu(delta-(1-eps))**2 + .03*torch.log1p(tr/tau) if controller else 0.*tr
78 return logits,H,delta,tr,penalty
79
80def train_variant(controller, steps=220):
81 device='cuda' if torch.cuda.is_available() else 'cpu'
82 try:
83 torch.manual_seed(SEED); model=TinyRNN().to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
84 # deterministic synthetic Markov sequence: target is next symbol, long enough to expose state growth
85 g=torch.Generator(device=device); g.manual_seed(SEED)
86 losses=[]; traces=[]; deltas=[]
87 for step in range(steps):
88 B,T=64,32; x=torch.randint(0,8,(B,T),generator=g,device=device); y=(x+1)%8
89 logits,H,d,tr,p=model(x,controller); loss=nn.functional.cross_entropy(logits.reshape(-1,8),y.reshape(-1))+p
90 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.0); opt.step()
91 if step%20==0 or step==steps-1:
92 losses.append(float(loss.detach().cpu())); traces.append(float(tr.detach().cpu())); deltas.append(float(d.detach().cpu()))
93 # long horizon evaluation without update
94 x=torch.randint(0,8,(64,128),generator=g,device=device); y=(x+1)%8
95 with torch.no_grad():
96 logits,H,d,tr,p=model(x,controller)
97 longloss=float(nn.functional.cross_entropy(logits.reshape(-1,8),y.reshape(-1)).cpu())
98 longtr=float(tr.cpu()); longd=float(d.cpu())
99 return {'train_loss_checkpoints':losses,'trace_checkpoints':traces,'delta_checkpoints':deltas,
100 'long_loss':longloss,'long_trace':longtr,'long_delta':longd,'device':device}
101 except Exception as e:
102 if device=='cuda':
103 torch.cuda.empty_cache(); torch.manual_seed(SEED)
104 # retry on CPU by temporarily disabling CUDA availability is awkward; call an explicit CPU helper
105 old=torch.cuda.is_available
106 torch.cuda.is_available=lambda: False
107 try: return train_variant(controller,steps)
108 finally: torch.cuda.is_available=old
109 raise
110
111def main():
112 out={'toy_checks':toy_checks(),'baseline':train_variant(False),'controller':train_variant(True)}
113 Path('results.json').write_text(json.dumps(out,indent=2))
114 print(json.dumps(out,indent=2))
115if __name__=='__main__': main()