Critical-Gain Covariance Controller / experiment.py

Mechanism failed

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