PSD-plus-low-rank curvature optimizer / psd_lowrank_experiment.py
Mechanism failed
1import json, random
2import numpy as np
3
4
5def toy_check(seed=7):
6 rng=np.random.default_rng(seed); d=24; r=2
7 P=np.diag(np.linspace(.7,2.0,d))
8 Q,_=np.linalg.qr(rng.normal(size=(d,r)))
9 # Signed correction, but total H remains positive definite as required by
10 # the local positive-quadratic stability claim.
11 C=np.diag([-.30, .45])
12 H=P+Q@[email protected]
13 Ps=np.sqrt(P)
14 M=np.linalg.solve(Ps,H) @ np.linalg.inv(Ps)
15 lam=float(np.max(np.linalg.eigvalsh(M))); eta_c=2.0/lam
16 A=np.eye(d)-eta_c*.95*np.linalg.solve(P,H)
17 B=np.eye(d)-eta_c*1.05*np.linalg.solve(P,H)
18 # Spectral contraction is the direct numerical stability test.
19 below_rho=float(max(abs(np.linalg.eigvals(A))))
20 above_rho=float(max(abs(np.linalg.eigvals(B))))
21 def run(eta,n=100):
22 x=rng.normal(size=d); norms=[]
23 for _ in range(n):
24 norms.append(float(np.linalg.norm(x)))
25 x=x-eta*np.linalg.solve(P,H@x)
26 return norms
27 below,above=run(.95*eta_c),run(1.05*eta_c)
28 neg_h=int(np.sum(np.linalg.eigvalsh(H)<0)); neg_c=int(np.sum(np.linalg.eigvalsh(C)<0))
29 stable=below_rho<1 and below[-1]<below[0]
30 divergent=above_rho>1 and above[-1]>above[0]
31 return {'eta_critical':eta_c,'lambda_max_M':lam,'below_spectral_radius':below_rho,
32 'above_spectral_radius':above_rho,'below_boundary_decay':stable,
33 'above_boundary_growth':divergent,'negative_H_eigenvalues':neg_h,
34 'negative_C_eigenvalues':neg_c,'passed':stable and divergent and neg_h==0}
35
36
37def run_digits(seed=11, steps=90):
38 import torch
39 from sklearn.datasets import load_digits
40 from sklearn.model_selection import train_test_split
41 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
42 device='cuda' if torch.cuda.is_available() else 'cpu'
43 try:
44 X,y=load_digits(return_X_y=True); X=X.astype('float32')/16.0
45 Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=seed,stratify=y)
46 def make():
47 return torch.nn.Sequential(torch.nn.Linear(64,48),torch.nn.Tanh(),
48 torch.nn.Linear(48,32),torch.nn.Tanh(),torch.nn.Linear(32,16),
49 torch.nn.Tanh(),torch.nn.Linear(16,10)).to(device)
50 class PSDLowRank(torch.optim.Optimizer):
51 def __init__(self,params,lr=.12,rank=2,every=15,damping=.03):
52 super().__init__(params,dict(lr=lr)); self.rank=rank; self.every=every; self.damping=damping
53 self.flat_params=[p for g in self.param_groups for p in g['params']]
54 self.v=[torch.zeros_like(p) for p in self.flat_params]; self.U=None; self.Cmat=None; self.t=0
55 def _flat(self,xs): return torch.cat([x.reshape(-1) for x in xs])
56 def _unflat(self,z):
57 out=[]; k=0
58 for p in self.flat_params:
59 n=p.numel(); out.append(z[k:k+n].view_as(p)); k+=n
60 return out
61 def step(self,closure=None):
62 if closure is None: raise RuntimeError('closure required')
63 loss=closure(); self.t+=1
64 gs=[p.grad.detach().clone() for p in self.flat_params]
65 for i,g in enumerate(gs): self.v[i]=.95*self.v[i]+.05*g*g
66 pdiag=self._flat([torch.sqrt(v)+self.damping for v in self.v])
67 gflat=self._flat(gs); pinvg=gflat/pdiag
68 if self.t==1 or self.t%self.every==0:
69 # Hessian-vector products of the current minibatch loss.
70 grads=torch.autograd.grad(loss,self.flat_params,create_graph=True,retain_graph=True)
71 def hv(v):
72 dot=sum((a*b).sum() for a,b in zip(grads,self._unflat(v)))
73 out=torch.autograd.grad(dot,self.flat_params,retain_graph=True,allow_unused=True)
74 return self._flat([o if o is not None else torch.zeros_like(p)
75 for o,p in zip(out,self.flat_params)])
76 W=torch.randn((gflat.numel(),self.rank),device=device); W,_=torch.linalg.qr(W)
77 for _ in range(2):
78 Z=torch.stack([hv(W[:,j])-pdiag*W[:,j] for j in range(self.rank)],1); W,_=torch.linalg.qr(Z)
79 R=torch.stack([hv(W[:,j])-pdiag*W[:,j] for j in range(self.rank)],1)
80 self.U=W.detach(); self.Cmat=((W.T@R)+(R.T@W)).detach()/2
81 if self.U is not None:
82 ev,V=torch.linalg.eigh(self.Cmat); z=self.U.T@pinvg
83 den=1.0+0.10*torch.minimum(ev,torch.zeros_like(ev))
84 corr=self.U@(V@((V.T@z)/den)); delta=pinvg+corr
85 # Conservative small-matrix spectral estimate and clipping.
86 invp_u=self.U/torch.sqrt(pdiag[:,None])
87 K=invp_u.T@invp_u
88 lmax=float(1+torch.linalg.norm(self.Cmat,2)*torch.linalg.norm(K,2))
89 eta=min(self.param_groups[0]['lr'],.9*2/max(lmax,1e-6))
90 else: delta=pinvg; eta=self.param_groups[0]['lr']
91 with torch.no_grad():
92 for p,zp in zip(self.flat_params,self._unflat(delta)): p.add_(zp,alpha=-eta)
93 return loss.detach()
94 initial=make().state_dict()
95 initial={k:v.detach().clone() for k,v in initial.items()}
96 def train(kind,lr):
97 net=make(); net.load_state_dict(initial); lossfn=torch.nn.CrossEntropyLoss()
98 if kind=='sgd': opt=torch.optim.SGD(net.parameters(),lr=lr)
99 else: opt=PSDLowRank(net.parameters(),lr=lr)
100 losses=[]
101 for t in range(steps):
102 ix=np.random.default_rng(seed+t).choice(len(Xtr),128,replace=False)
103 xb=torch.tensor(Xtr[ix],device=device); yb=torch.tensor(ytr[ix],device=device)
104 if kind=='sgd':
105 opt.zero_grad(); z=lossfn(net(xb),yb); z.backward(); losses.append(float(z)); opt.step()
106 else:
107 def closure():
108 opt.zero_grad(set_to_none=True); z=lossfn(net(xb),yb); z.backward(create_graph=True); return z
109 losses.append(float(opt.step(closure)))
110 with torch.no_grad():
111 acc=float((net(torch.tensor(Xte,device=device)).argmax(1).cpu().numpy()==yte).mean())
112 return {'final_loss':losses[-1],'min_loss':min(losses),'accuracy':acc,'diverged':not np.isfinite(losses).all()}
113 return {'device':device,'sgd':train('sgd',.12),'psd_lowrank':train('curvature',.12)}
114 except Exception as e:
115 if device=='cuda':
116 torch.cuda.empty_cache(); return {'device':'cpu-fallback-failed','error':repr(e)}
117 return {'device':device,'error':repr(e)}
118
119if __name__=='__main__':
120 out={'toy':toy_check(),'digits':run_digits()}
121 with open('results.json','w') as f: json.dump(out,f,indent=2)
122 print(json.dumps(out,indent=2))