import random, time import numpy as np import torch import torch.nn as nn SEED = 47 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(12) torch.set_default_dtype(torch.float64) class MLP(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(2, 16), nn.Tanh(), nn.Linear(16, 1)) def forward(self, x): return self.net(x) def pvec(model): return torch.cat([p.detach().reshape(-1) for p in model.parameters()]) def setp(model, v): off = 0 with torch.no_grad(): for p in model.parameters(): n = p.numel(); p.copy_(v[off:off+n].reshape_as(p)); off += n def pieces(model, loss, create_graph=True): return torch.autograd.grad(loss, tuple(model.parameters()), create_graph=create_graph, retain_graph=create_graph) def flat(xs): return torch.cat([x.reshape(-1) for x in xs]) def split(model, v): out=[]; off=0 for p in model.parameters(): n=p.numel(); out.append(v[off:off+n].reshape_as(p)); off += n return out def loss_grad(model, x, y, graph=False): loss = ((model(x)-y)**2).mean() g = flat(pieces(model, loss, create_graph=graph)) return loss, g def hvp(model, loss, v): gs = pieces(model, loss, create_graph=True) dot = sum((a*b).sum() for a,b in zip(gs, split(model, v))) hs = torch.autograd.grad(dot, tuple(model.parameters()), retain_graph=True) return flat(hs) def lanczos(model, loss, g, k=6): ng = torch.linalg.vector_norm(g) if float(ng) < 1e-12: return torch.zeros((g.numel(),0)), torch.zeros((0,0)), 0 qs=[]; al=[]; bet=[]; q=g/ng; qprev=torch.zeros_like(q) for j in range(k): w=hvp(model, loss, q) if j: w = w - bet[-1]*qprev a=torch.dot(q,w); w=w-a*q # Full reorthogonalization keeps the small basis numerically valid. for old in qs: w=w-torch.dot(old,w)*old b=torch.linalg.vector_norm(w) qs.append(q); al.append(a) if float(b)<1e-9 or j==k-1: break bet.append(b); qprev=q; q=w/b Q=torch.stack(qs, dim=1); kk=Q.shape[1] T=torch.diag(torch.stack(al)) for i,b in enumerate(bet): T[i,i+1]=b; T[i+1,i]=b return Q,T,kk def cubic_solution(T, c, sigma): # Solve (T+lambda I)z=-c, lambda=(sigma/2)||z||. d=T.shape[0] if d==0: return torch.zeros(0), 0.0 I=torch.eye(d, dtype=T.dtype) def solve(lam): return torch.linalg.solve(T+(lam+1e-10)*I, -c) def phi(lam): z=solve(lam); return lam-(sigma/2)*float(torch.linalg.vector_norm(z)) lo=0.0; hi=1.0 # Positive lambda also handles negative curvature and singular T. while phi(hi)<0 and hi<1e8: hi*=2 for _ in range(70): mid=(lo+hi)/2 if phi(mid)>0: hi=mid else: lo=mid lam=(lo+hi)/2; return solve(lam), lam def model_value(g, Hs, s, sigma): return torch.dot(g,s)+0.5*torch.dot(s,Hs)+sigma/6*torch.linalg.vector_norm(s)**3 def cubic_step(model, x, y, sigma, k=6): loss,g=loss_grad(model,x,y,graph=True) Q,T,kk=lanczos(model,loss,g,k) if Q.shape[1]==0: return float(loss), False, 0, sigma, 0.0 c=Q.T@g; z,lam=cubic_solution(T,c,sigma); s=Q@z old=pvec(model); setp(model,old+s) with torch.no_grad(): newloss=float(((model(x)-y)**2).mean()) mv=float(torch.dot(c,z)+0.5*torch.dot(z,T@z)+sigma/6*torch.linalg.vector_norm(z)**3) actual=float(loss)-newloss rho=actual/(-mv) if mv < -1e-14 else -1.0 accept=actual>0.0 and rho>0.0 if not accept: setp(model,old) if not accept or rho<0.1: sigma=min(1e4, sigma*2) elif rho>0.75: sigma=max(1e-6, sigma/2) return (newloss if accept else float(loss)), accept, kk, sigma, rho def math_check(): # Exact quadratic: Lanczos with k=n recovers H, and stationarity residual is small. torch.manual_seed(SEED) n=9; A=torch.randn(n,n); H=A.T@A-1.5*torch.eye(n); g=torch.randn(n); sigma=2.3 Q,T,_=lanczos_matrix(H,g,n) z,lam=cubic_solution(T,Q.T@g,sigma); s=Q@z resid=torch.linalg.vector_norm(g+H@s+(sigma/2)*torch.linalg.vector_norm(s)*s) # Cubic upper-bound identity for a synthetic Lipschitz-Hessian scalar remainder. L=4.0; r=torch.randn(n); remainder=L/6*torch.linalg.vector_norm(r)**3 return float(resid), float(torch.linalg.vector_norm(Q.T@Q-torch.eye(n))), float(remainder) def lanczos_matrix(H,g,k): qs=[]; al=[]; bet=[]; q=g/torch.linalg.vector_norm(g); prev=torch.zeros_like(q) for j in range(k): w=H@q if j: w-=bet[-1]*prev a=q@w; w-=a*q for old in qs: w-=old@(w)*old b=torch.linalg.vector_norm(w); qs.append(q); al.append(a) if float(b)<1e-10 or j==k-1: break bet.append(b); prev=q; q=w/b Q=torch.stack(qs,1); T=torch.diag(torch.stack(al)) for i,b in enumerate(bet): T[i,i+1]=b; T[i+1,i]=b return Q,T,len(qs) def make_data(): torch.manual_seed(SEED+1) x=torch.linspace(-2,2,160).reshape(-1,1) X=torch.cat([x, torch.sin(2*x)],1) y=0.7*torch.sin(3*x)+0.25*x*x+0.08*torch.randn_like(x) return X,y def run(): r1,r2,_=math_check() X,y=make_data() def train(kind): torch.manual_seed(SEED+2); m=MLP(); t0=time.perf_counter(); accepts=0; hvps=0; sig=1.0 losses=[] if kind=='adam': opt=torch.optim.Adam(m.parameters(),lr=0.025) for _ in range(80): opt.zero_grad(); loss=((m(X)-y)**2).mean(); loss.backward(); opt.step(); losses.append(float(loss)) else: for _ in range(35): loss,ok,used,sig,rho=cubic_step(m,X,y,sig,6); accepts+=int(ok); hvps+=used; losses.append(loss) return losses[-1],min(losses),accepts,hvps,time.perf_counter()-t0,losses b=train('adam'); c=train('cubic') print('math_check_stationarity_residual %.3e orthogonality_error %.3e' % (r1,r2)) print('baseline_adam final %.6f best %.6f accepted n/a hvps n/a time %.3fs' % (b[0],b[1],b[4])) print('idea_cubic final %.6f best %.6f accepted %d/35 hvps %d time %.3fs' % (c[0],c[1],c[2],c[3],c[4])) print('loss_curves_adam',','.join('%.5f'%v for v in b[5][::10])) print('loss_curves_cubic',','.join('%.5f'%v for v in c[5][::5])) with open('results.txt','w') as f: f.write('math_check_stationarity_residual %.3e orthogonality_error %.3e\n' % (r1,r2)) f.write('adam final %.9f best %.9f time %.6f\n' % (b[0],b[1],b[4])) f.write('cubic final %.9f best %.9f accepted %d/35 hvps %d time %.6f\n' % (c[0],c[1],c[2],c[3],c[4])) f.write('adam_curve '+','.join('%.9f'%v for v in b[5])+'\n') f.write('cubic_curve '+','.join('%.9f'%v for v in c[5])+'\n') if __name__=='__main__': run()