import json import numpy as np import torch torch.set_num_threads(12) np.random.seed(7); torch.manual_seed(7) def dfp_update(H, s, y): sy = float(s @ y) Hy = H @ y return H - np.outer(Hy, Hy) / float(y @ Hy) + np.outer(s, s) / sy def eig_projector(H, r=2): w, V = np.linalg.eigh((H + H.T) * .5) r = min(r, len(w)) return V[:, :r] @ V[:, :r].T, w def safeguarded_update(H, s, y, old_P, eps=1e-3, tau=.25, kappa_max=1e4): raw = dfp_update(H, s, y) raw = (raw + raw.T) * .5 P, w = eig_projector(raw) q = 0.0 if old_P is None else float(np.linalg.norm(P - old_P, 2)) unsafe = w[0] < eps or w[-1] / max(w[0], 1e-12) > kappa_max or q > tau rho = 1.0 if unsafe: accepted = False for candidate in (0.5, 0.25, 0.1, 0.0): B = (1.0 - candidate) * H + candidate * raw ew = np.linalg.eigvalsh((B + B.T) * .5) if ew[0] >= eps and ew[-1] <= kappa_max * eps: raw, rho, accepted = B, candidate, True break if not accepted: raw, rho = H.copy(), 0.0 ew, V = np.linalg.eigh((raw + raw.T) * .5) ew = np.clip(ew, eps, kappa_max * eps) return (V * ew) @ V.T, P, q, bool(unsafe), rho def math_check(): rng = np.random.default_rng(3); n = 8 A = rng.normal(size=(n, n)); H = A.T @ A + .5*np.eye(n) s = rng.normal(size=n); y = rng.normal(size=n) y += (1.0 - s @ y) * s / (s @ s) H2 = dfp_update(H, s, y) secant_error = float(np.linalg.norm(H2 @ y - s)) # Keep the tiny eigendirection untouched so the safeguard must detect it. Hbad = np.diag([1e-8, 1., 2., 3.]) ss = np.array([0., 1., 0., 0.]); yy = ss.copy() Hsafe, _, _, triggered, _ = safeguarded_update(Hbad, ss, yy, None, eps=1e-3) safe_min = float(np.linalg.eigvalsh(Hsafe)[0]) return {'secant_error': secant_error, 'ordinary_min_eig': float(np.linalg.eigvalsh(H2)[0]), 'floor_min_eig': safe_min, 'safeguard_triggered': triggered, 'secant_pass': secant_error < 1e-9, 'floor_pass': triggered and safe_min >= 1e-3-1e-10} class TinyNet(torch.nn.Module): def __init__(self): super().__init__(); self.net = torch.nn.Sequential( torch.nn.Linear(2, 24), torch.nn.Tanh(), torch.nn.Linear(24, 2)) def forward(self, x): return self.net(x) def data(): rng = np.random.default_rng(11); n = 240 t = rng.uniform(0, 2*np.pi, n) r = np.where(np.arange(n) % 2 == 0, .8, 1.8) + .10*rng.normal(size=n) return torch.tensor(np.c_[r*np.cos(t), r*np.sin(t)].astype('float32')), torch.tensor(np.arange(n)%2, dtype=torch.long) def flat_params(m): return torch.cat([p.detach().reshape(-1) for p in m.parameters()]).numpy() def 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() def set_flat(m, v): k=0 with torch.no_grad(): for p in m.parameters(): z=p.numel(); p.copy_(torch.tensor(v[k:k+z], dtype=p.dtype).reshape_as(p)); k += z def run(kind, X, Y, steps): torch.manual_seed(19); m=TinyNet(); d=sum(p.numel() for p in m.parameters()) H=np.eye(d); oldP=None; g_evals=0; safeguards=0; min_seen=float('inf') losses=[] for _ in range(steps): m.zero_grad(); loss=torch.nn.functional.cross_entropy(m(X),Y); loss.backward() g=flat_grad(m).astype(float); x=flat_params(m).astype(float); g_evals += 1; losses.append(float(loss)) if kind == 'sgd': xnew=x-.12*g else: dvec=-H@g; alpha=1.; f0=float(loss); gd=float(g@dvec) for _ in range(12): set_flat(m,x+alpha*dvec) with torch.no_grad(): ft=float(torch.nn.functional.cross_entropy(m(X),Y)) if ft <= f0 + 1e-4*alpha*gd: break alpha *= .5 xnew=x+alpha*dvec; set_flat(m,xnew) m.zero_grad(); l2=torch.nn.functional.cross_entropy(m(X),Y); l2.backward() g2=flat_grad(m).astype(float); g_evals += 1; s=xnew-x; y=g2-g if s@y > 1e-10*np.linalg.norm(s)*np.linalg.norm(y): if kind == 'dfp': H=(dfp_update(H,s,y)+dfp_update(H,s,y).T)*.5 else: H,P,q,tr,rho=safeguarded_update(H,s,y,oldP); oldP=P; safeguards += int(tr) min_seen=min(min_seen, float(np.linalg.eigvalsh(H)[0])) x=xnew set_flat(m,x); m.zero_grad(); lf=torch.nn.functional.cross_entropy(m(X),Y); lf.backward() return {'final_loss':float(lf), 'final_grad_norm':float(np.linalg.norm(flat_grad(m))), 'best_loss':min(losses), 'gradient_evals':g_evals, 'min_H_eigenvalue':min_seen, 'safeguards':safeguards} if __name__ == '__main__': X,Y=data(); out={'math_check':math_check(), 'results':{ 'sgd':run('sgd',X,Y,100), 'dfp':run('dfp',X,Y,50), 'safeguarded_dfp':run('safeguarded_dfp',X,Y,50)}} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out,indent=2))