Spectrally safeguarded DFP / experiment.py
Beats tuned baseline
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))