import json, random import numpy as np def toy_check(seed=7): rng=np.random.default_rng(seed); d=24; r=2 P=np.diag(np.linspace(.7,2.0,d)) Q,_=np.linalg.qr(rng.normal(size=(d,r))) # Signed correction, but total H remains positive definite as required by # the local positive-quadratic stability claim. C=np.diag([-.30, .45]) H=P+Q@C@Q.T Ps=np.sqrt(P) M=np.linalg.solve(Ps,H) @ np.linalg.inv(Ps) lam=float(np.max(np.linalg.eigvalsh(M))); eta_c=2.0/lam A=np.eye(d)-eta_c*.95*np.linalg.solve(P,H) B=np.eye(d)-eta_c*1.05*np.linalg.solve(P,H) # Spectral contraction is the direct numerical stability test. below_rho=float(max(abs(np.linalg.eigvals(A)))) above_rho=float(max(abs(np.linalg.eigvals(B)))) def run(eta,n=100): x=rng.normal(size=d); norms=[] for _ in range(n): norms.append(float(np.linalg.norm(x))) x=x-eta*np.linalg.solve(P,H@x) return norms below,above=run(.95*eta_c),run(1.05*eta_c) neg_h=int(np.sum(np.linalg.eigvalsh(H)<0)); neg_c=int(np.sum(np.linalg.eigvalsh(C)<0)) stable=below_rho<1 and below[-1]1 and above[-1]>above[0] return {'eta_critical':eta_c,'lambda_max_M':lam,'below_spectral_radius':below_rho, 'above_spectral_radius':above_rho,'below_boundary_decay':stable, 'above_boundary_growth':divergent,'negative_H_eigenvalues':neg_h, 'negative_C_eigenvalues':neg_c,'passed':stable and divergent and neg_h==0} def run_digits(seed=11, steps=90): import torch from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) device='cuda' if torch.cuda.is_available() else 'cpu' try: X,y=load_digits(return_X_y=True); X=X.astype('float32')/16.0 Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=seed,stratify=y) def make(): return torch.nn.Sequential(torch.nn.Linear(64,48),torch.nn.Tanh(), torch.nn.Linear(48,32),torch.nn.Tanh(),torch.nn.Linear(32,16), torch.nn.Tanh(),torch.nn.Linear(16,10)).to(device) class PSDLowRank(torch.optim.Optimizer): def __init__(self,params,lr=.12,rank=2,every=15,damping=.03): super().__init__(params,dict(lr=lr)); self.rank=rank; self.every=every; self.damping=damping self.flat_params=[p for g in self.param_groups for p in g['params']] self.v=[torch.zeros_like(p) for p in self.flat_params]; self.U=None; self.Cmat=None; self.t=0 def _flat(self,xs): return torch.cat([x.reshape(-1) for x in xs]) def _unflat(self,z): out=[]; k=0 for p in self.flat_params: n=p.numel(); out.append(z[k:k+n].view_as(p)); k+=n return out def step(self,closure=None): if closure is None: raise RuntimeError('closure required') loss=closure(); self.t+=1 gs=[p.grad.detach().clone() for p in self.flat_params] for i,g in enumerate(gs): self.v[i]=.95*self.v[i]+.05*g*g pdiag=self._flat([torch.sqrt(v)+self.damping for v in self.v]) gflat=self._flat(gs); pinvg=gflat/pdiag if self.t==1 or self.t%self.every==0: # Hessian-vector products of the current minibatch loss. grads=torch.autograd.grad(loss,self.flat_params,create_graph=True,retain_graph=True) def hv(v): dot=sum((a*b).sum() for a,b in zip(grads,self._unflat(v))) out=torch.autograd.grad(dot,self.flat_params,retain_graph=True,allow_unused=True) return self._flat([o if o is not None else torch.zeros_like(p) for o,p in zip(out,self.flat_params)]) W=torch.randn((gflat.numel(),self.rank),device=device); W,_=torch.linalg.qr(W) for _ in range(2): Z=torch.stack([hv(W[:,j])-pdiag*W[:,j] for j in range(self.rank)],1); W,_=torch.linalg.qr(Z) R=torch.stack([hv(W[:,j])-pdiag*W[:,j] for j in range(self.rank)],1) self.U=W.detach(); self.Cmat=((W.T@R)+(R.T@W)).detach()/2 if self.U is not None: ev,V=torch.linalg.eigh(self.Cmat); z=self.U.T@pinvg den=1.0+0.10*torch.minimum(ev,torch.zeros_like(ev)) corr=self.U@(V@((V.T@z)/den)); delta=pinvg+corr # Conservative small-matrix spectral estimate and clipping. invp_u=self.U/torch.sqrt(pdiag[:,None]) K=invp_u.T@invp_u lmax=float(1+torch.linalg.norm(self.Cmat,2)*torch.linalg.norm(K,2)) eta=min(self.param_groups[0]['lr'],.9*2/max(lmax,1e-6)) else: delta=pinvg; eta=self.param_groups[0]['lr'] with torch.no_grad(): for p,zp in zip(self.flat_params,self._unflat(delta)): p.add_(zp,alpha=-eta) return loss.detach() initial=make().state_dict() initial={k:v.detach().clone() for k,v in initial.items()} def train(kind,lr): net=make(); net.load_state_dict(initial); lossfn=torch.nn.CrossEntropyLoss() if kind=='sgd': opt=torch.optim.SGD(net.parameters(),lr=lr) else: opt=PSDLowRank(net.parameters(),lr=lr) losses=[] for t in range(steps): ix=np.random.default_rng(seed+t).choice(len(Xtr),128,replace=False) xb=torch.tensor(Xtr[ix],device=device); yb=torch.tensor(ytr[ix],device=device) if kind=='sgd': opt.zero_grad(); z=lossfn(net(xb),yb); z.backward(); losses.append(float(z)); opt.step() else: def closure(): opt.zero_grad(set_to_none=True); z=lossfn(net(xb),yb); z.backward(create_graph=True); return z losses.append(float(opt.step(closure))) with torch.no_grad(): acc=float((net(torch.tensor(Xte,device=device)).argmax(1).cpu().numpy()==yte).mean()) return {'final_loss':losses[-1],'min_loss':min(losses),'accuracy':acc,'diverged':not np.isfinite(losses).all()} return {'device':device,'sgd':train('sgd',.12),'psd_lowrank':train('curvature',.12)} except Exception as e: if device=='cuda': torch.cuda.empty_cache(); return {'device':'cpu-fallback-failed','error':repr(e)} return {'device':device,'error':repr(e)} if __name__=='__main__': out={'toy':toy_check(),'digits':run_digits()} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out,indent=2))