Krylov Block-Cubic Optimizer / experiment.py
Mechanism failed
1import random, time
2import numpy as np
3import torch
4import torch.nn as nn
5
6SEED = 47
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8torch.set_num_threads(12)
9torch.set_default_dtype(torch.float64)
10
11class MLP(nn.Module):
12 def __init__(self):
13 super().__init__()
14 self.net = nn.Sequential(nn.Linear(2, 16), nn.Tanh(), nn.Linear(16, 1))
15 def forward(self, x): return self.net(x)
16
17def pvec(model):
18 return torch.cat([p.detach().reshape(-1) for p in model.parameters()])
19
20def setp(model, v):
21 off = 0
22 with torch.no_grad():
23 for p in model.parameters():
24 n = p.numel(); p.copy_(v[off:off+n].reshape_as(p)); off += n
25
26def pieces(model, loss, create_graph=True):
27 return torch.autograd.grad(loss, tuple(model.parameters()), create_graph=create_graph,
28 retain_graph=create_graph)
29
30def flat(xs): return torch.cat([x.reshape(-1) for x in xs])
31
32def split(model, v):
33 out=[]; off=0
34 for p in model.parameters():
35 n=p.numel(); out.append(v[off:off+n].reshape_as(p)); off += n
36 return out
37
38def loss_grad(model, x, y, graph=False):
39 loss = ((model(x)-y)**2).mean()
40 g = flat(pieces(model, loss, create_graph=graph))
41 return loss, g
42
43def hvp(model, loss, v):
44 gs = pieces(model, loss, create_graph=True)
45 dot = sum((a*b).sum() for a,b in zip(gs, split(model, v)))
46 hs = torch.autograd.grad(dot, tuple(model.parameters()), retain_graph=True)
47 return flat(hs)
48
49def lanczos(model, loss, g, k=6):
50 ng = torch.linalg.vector_norm(g)
51 if float(ng) < 1e-12: return torch.zeros((g.numel(),0)), torch.zeros((0,0)), 0
52 qs=[]; al=[]; bet=[]; q=g/ng; qprev=torch.zeros_like(q)
53 for j in range(k):
54 w=hvp(model, loss, q)
55 if j: w = w - bet[-1]*qprev
56 a=torch.dot(q,w); w=w-a*q
57 # Full reorthogonalization keeps the small basis numerically valid.
58 for old in qs: w=w-torch.dot(old,w)*old
59 b=torch.linalg.vector_norm(w)
60 qs.append(q); al.append(a)
61 if float(b)<1e-9 or j==k-1: break
62 bet.append(b); qprev=q; q=w/b
63 Q=torch.stack(qs, dim=1); kk=Q.shape[1]
64 T=torch.diag(torch.stack(al))
65 for i,b in enumerate(bet): T[i,i+1]=b; T[i+1,i]=b
66 return Q,T,kk
67
68def cubic_solution(T, c, sigma):
69 # Solve (T+lambda I)z=-c, lambda=(sigma/2)||z||.
70 d=T.shape[0]
71 if d==0: return torch.zeros(0), 0.0
72 I=torch.eye(d, dtype=T.dtype)
73 def solve(lam):
74 return torch.linalg.solve(T+(lam+1e-10)*I, -c)
75 def phi(lam):
76 z=solve(lam); return lam-(sigma/2)*float(torch.linalg.vector_norm(z))
77 lo=0.0; hi=1.0
78 # Positive lambda also handles negative curvature and singular T.
79 while phi(hi)<0 and hi<1e8: hi*=2
80 for _ in range(70):
81 mid=(lo+hi)/2
82 if phi(mid)>0: hi=mid
83 else: lo=mid
84 lam=(lo+hi)/2; return solve(lam), lam
85
86def model_value(g, Hs, s, sigma):
87 return torch.dot(g,s)+0.5*torch.dot(s,Hs)+sigma/6*torch.linalg.vector_norm(s)**3
88
89def cubic_step(model, x, y, sigma, k=6):
90 loss,g=loss_grad(model,x,y,graph=True)
91 Q,T,kk=lanczos(model,loss,g,k)
92 if Q.shape[1]==0: return float(loss), False, 0, sigma, 0.0
93 c=Q.T@g; z,lam=cubic_solution(T,c,sigma); s=Q@z
94 old=pvec(model); setp(model,old+s)
95 with torch.no_grad(): newloss=float(((model(x)-y)**2).mean())
96 mv=float(torch.dot(c,z)+0.5*torch.dot(z,T@z)+sigma/6*torch.linalg.vector_norm(z)**3)
97 actual=float(loss)-newloss
98 rho=actual/(-mv) if mv < -1e-14 else -1.0
99 accept=actual>0.0 and rho>0.0
100 if not accept: setp(model,old)
101 if not accept or rho<0.1: sigma=min(1e4, sigma*2)
102 elif rho>0.75: sigma=max(1e-6, sigma/2)
103 return (newloss if accept else float(loss)), accept, kk, sigma, rho
104
105def math_check():
106 # Exact quadratic: Lanczos with k=n recovers H, and stationarity residual is small.
107 torch.manual_seed(SEED)
108 n=9; A=torch.randn(n,n); H=A.T@A-1.5*torch.eye(n); g=torch.randn(n); sigma=2.3
109 Q,T,_=lanczos_matrix(H,g,n)
110 z,lam=cubic_solution(T,Q.T@g,sigma); s=Q@z
111 resid=torch.linalg.vector_norm(g+H@s+(sigma/2)*torch.linalg.vector_norm(s)*s)
112 # Cubic upper-bound identity for a synthetic Lipschitz-Hessian scalar remainder.
113 L=4.0; r=torch.randn(n); remainder=L/6*torch.linalg.vector_norm(r)**3
114 return float(resid), float(torch.linalg.vector_norm(Q.T@Q-torch.eye(n))), float(remainder)
115
116def lanczos_matrix(H,g,k):
117 qs=[]; al=[]; bet=[]; q=g/torch.linalg.vector_norm(g); prev=torch.zeros_like(q)
118 for j in range(k):
119 w=H@q
120 if j: w-=bet[-1]*prev
121 a=q@w; w-=a*q
122 for old in qs: w-=old@(w)*old
123 b=torch.linalg.vector_norm(w); qs.append(q); al.append(a)
124 if float(b)<1e-10 or j==k-1: break
125 bet.append(b); prev=q; q=w/b
126 Q=torch.stack(qs,1); T=torch.diag(torch.stack(al))
127 for i,b in enumerate(bet): T[i,i+1]=b; T[i+1,i]=b
128 return Q,T,len(qs)
129
130def make_data():
131 torch.manual_seed(SEED+1)
132 x=torch.linspace(-2,2,160).reshape(-1,1)
133 X=torch.cat([x, torch.sin(2*x)],1)
134 y=0.7*torch.sin(3*x)+0.25*x*x+0.08*torch.randn_like(x)
135 return X,y
136
137def run():
138 r1,r2,_=math_check()
139 X,y=make_data()
140 def train(kind):
141 torch.manual_seed(SEED+2); m=MLP(); t0=time.perf_counter(); accepts=0; hvps=0; sig=1.0
142 losses=[]
143 if kind=='adam':
144 opt=torch.optim.Adam(m.parameters(),lr=0.025)
145 for _ in range(80):
146 opt.zero_grad(); loss=((m(X)-y)**2).mean(); loss.backward(); opt.step(); losses.append(float(loss))
147 else:
148 for _ in range(35):
149 loss,ok,used,sig,rho=cubic_step(m,X,y,sig,6); accepts+=int(ok); hvps+=used; losses.append(loss)
150 return losses[-1],min(losses),accepts,hvps,time.perf_counter()-t0,losses
151 b=train('adam'); c=train('cubic')
152 print('math_check_stationarity_residual %.3e orthogonality_error %.3e' % (r1,r2))
153 print('baseline_adam final %.6f best %.6f accepted n/a hvps n/a time %.3fs' % (b[0],b[1],b[4]))
154 print('idea_cubic final %.6f best %.6f accepted %d/35 hvps %d time %.3fs' % (c[0],c[1],c[2],c[3],c[4]))
155 print('loss_curves_adam',','.join('%.5f'%v for v in b[5][::10]))
156 print('loss_curves_cubic',','.join('%.5f'%v for v in c[5][::5]))
157 with open('results.txt','w') as f:
158 f.write('math_check_stationarity_residual %.3e orthogonality_error %.3e\n' % (r1,r2))
159 f.write('adam final %.9f best %.9f time %.6f\n' % (b[0],b[1],b[4]))
160 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]))
161 f.write('adam_curve '+','.join('%.9f'%v for v in b[5])+'\n')
162 f.write('cubic_curve '+','.join('%.9f'%v for v in c[5])+'\n')
163
164if __name__=='__main__': run()