Spectrally safeguarded DFP / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json
  2import numpy as np
  3import torch
  4
  5torch.set_num_threads(12)
  6np.random.seed(7); torch.manual_seed(7)
  7
  8
  9def dfp_update(H, s, y):
 10    sy = float(s @ y)
 11    Hy = H @ y
 12    return H - np.outer(Hy, Hy) / float(y @ Hy) + np.outer(s, s) / sy
 13
 14
 15def eig_projector(H, r=2):
 16    w, V = np.linalg.eigh((H + H.T) * .5)
 17    r = min(r, len(w))
 18    return V[:, :r] @ V[:, :r].T, w
 19
 20
 21def safeguarded_update(H, s, y, old_P, eps=1e-3, tau=.25, kappa_max=1e4):
 22    raw = dfp_update(H, s, y)
 23    raw = (raw + raw.T) * .5
 24    P, w = eig_projector(raw)
 25    q = 0.0 if old_P is None else float(np.linalg.norm(P - old_P, 2))
 26    unsafe = w[0] < eps or w[-1] / max(w[0], 1e-12) > kappa_max or q > tau
 27    rho = 1.0
 28    if unsafe:
 29        accepted = False
 30        for candidate in (0.5, 0.25, 0.1, 0.0):
 31            B = (1.0 - candidate) * H + candidate * raw
 32            ew = np.linalg.eigvalsh((B + B.T) * .5)
 33            if ew[0] >= eps and ew[-1] <= kappa_max * eps:
 34                raw, rho, accepted = B, candidate, True
 35                break
 36        if not accepted:
 37            raw, rho = H.copy(), 0.0
 38    ew, V = np.linalg.eigh((raw + raw.T) * .5)
 39    ew = np.clip(ew, eps, kappa_max * eps)
 40    return (V * ew) @ V.T, P, q, bool(unsafe), rho
 41
 42
 43def math_check():
 44    rng = np.random.default_rng(3); n = 8
 45    A = rng.normal(size=(n, n)); H = A.T @ A + .5*np.eye(n)
 46    s = rng.normal(size=n); y = rng.normal(size=n)
 47    y += (1.0 - s @ y) * s / (s @ s)
 48    H2 = dfp_update(H, s, y)
 49    secant_error = float(np.linalg.norm(H2 @ y - s))
 50    # Keep the tiny eigendirection untouched so the safeguard must detect it.
 51    Hbad = np.diag([1e-8, 1., 2., 3.])
 52    ss = np.array([0., 1., 0., 0.]); yy = ss.copy()
 53    Hsafe, _, _, triggered, _ = safeguarded_update(Hbad, ss, yy, None, eps=1e-3)
 54    safe_min = float(np.linalg.eigvalsh(Hsafe)[0])
 55    return {'secant_error': secant_error,
 56            'ordinary_min_eig': float(np.linalg.eigvalsh(H2)[0]),
 57            'floor_min_eig': safe_min, 'safeguard_triggered': triggered,
 58            'secant_pass': secant_error < 1e-9,
 59            'floor_pass': triggered and safe_min >= 1e-3-1e-10}
 60
 61
 62class TinyNet(torch.nn.Module):
 63    def __init__(self):
 64        super().__init__(); self.net = torch.nn.Sequential(
 65            torch.nn.Linear(2, 24), torch.nn.Tanh(), torch.nn.Linear(24, 2))
 66    def forward(self, x): return self.net(x)
 67
 68
 69def data():
 70    rng = np.random.default_rng(11); n = 240
 71    t = rng.uniform(0, 2*np.pi, n)
 72    r = np.where(np.arange(n) % 2 == 0, .8, 1.8) + .10*rng.normal(size=n)
 73    return torch.tensor(np.c_[r*np.cos(t), r*np.sin(t)].astype('float32')), torch.tensor(np.arange(n)%2, dtype=torch.long)
 74
 75
 76def flat_params(m): return torch.cat([p.detach().reshape(-1) for p in m.parameters()]).numpy()
 77def flat_grad(m): return torch.cat([(torch.zeros_like(p) if p.grad is None else p.grad).reshape(-1) for p in m.parameters()]).detach().numpy()
 78def set_flat(m, v):
 79    k=0
 80    with torch.no_grad():
 81        for p in m.parameters():
 82            z=p.numel(); p.copy_(torch.tensor(v[k:k+z], dtype=p.dtype).reshape_as(p)); k += z
 83
 84
 85def run(kind, X, Y, steps):
 86    torch.manual_seed(19); m=TinyNet(); d=sum(p.numel() for p in m.parameters())
 87    H=np.eye(d); oldP=None; g_evals=0; safeguards=0; min_seen=float('inf')
 88    losses=[]
 89    for _ in range(steps):
 90        m.zero_grad(); loss=torch.nn.functional.cross_entropy(m(X),Y); loss.backward()
 91        g=flat_grad(m).astype(float); x=flat_params(m).astype(float); g_evals += 1; losses.append(float(loss))
 92        if kind == 'sgd': xnew=x-.12*g
 93        else:
 94            dvec=-H@g; alpha=1.; f0=float(loss); gd=float(g@dvec)
 95            for _ in range(12):
 96                set_flat(m,x+alpha*dvec)
 97                with torch.no_grad(): ft=float(torch.nn.functional.cross_entropy(m(X),Y))
 98                if ft <= f0 + 1e-4*alpha*gd: break
 99                alpha *= .5
100            xnew=x+alpha*dvec; set_flat(m,xnew)
101            m.zero_grad(); l2=torch.nn.functional.cross_entropy(m(X),Y); l2.backward()
102            g2=flat_grad(m).astype(float); g_evals += 1; s=xnew-x; y=g2-g
103            if s@y > 1e-10*np.linalg.norm(s)*np.linalg.norm(y):
104                if kind == 'dfp': H=(dfp_update(H,s,y)+dfp_update(H,s,y).T)*.5
105                else:
106                    H,P,q,tr,rho=safeguarded_update(H,s,y,oldP); oldP=P; safeguards += int(tr)
107                min_seen=min(min_seen, float(np.linalg.eigvalsh(H)[0]))
108            x=xnew
109    set_flat(m,x); m.zero_grad(); lf=torch.nn.functional.cross_entropy(m(X),Y); lf.backward()
110    return {'final_loss':float(lf), 'final_grad_norm':float(np.linalg.norm(flat_grad(m))),
111            'best_loss':min(losses), 'gradient_evals':g_evals,
112            'min_H_eigenvalue':min_seen, 'safeguards':safeguards}
113
114
115if __name__ == '__main__':
116    X,Y=data(); out={'math_check':math_check(), 'results':{
117        'sgd':run('sgd',X,Y,100), 'dfp':run('dfp',X,Y,50),
118        'safeguarded_dfp':run('safeguarded_dfp',X,Y,50)}}
119    with open('results.json','w') as f: json.dump(out,f,indent=2)
120    print(json.dumps(out,indent=2))