PSD-plus-low-rank curvature optimizer / psd_lowrank_experiment.py

Mechanism failed

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