Krylov Block-Cubic Optimizer / experiment.py

Mechanism failed

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