Policy-Guided Terminal Trust Region for Optimizers / experiment.py
Mechanism failed
1import json, math, random, time
2import numpy as np
3import torch
4
5SEED=3072
6random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
7try:
8 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
9 if device.type=='cuda': torch.cuda.get_device_properties(0)
10except Exception:
11 device=torch.device('cpu')
12
13def contraction_check():
14 # Stable linear guiding policy: c+ = (I-alpha A)c, whose exact q is spectral norm.
15 A=np.diag([1.,3.]); alpha=.25
16 M=np.eye(2)-alpha*A
17 q=np.linalg.norm(M,2)
18 c=np.array([.8,-.5]); x=c+np.array([.37,-.22]); initial=np.linalg.norm(x-c)
19 ratios=[]; ds=[]
20 for _ in range(6):
21 x=M@x; c=M@c
22 d=np.linalg.norm(x-c); ds.append(float(d)); ratios.append(float(d/initial)); initial=d
23 # terminal feasibility: perturbation inside radius r remains inside after fallback
24 r=.5; perturb=np.array([.4,.3]); c0=np.zeros(2); x0=c0+perturb
25 terminal=[]
26 for _ in range(5):
27 x0=M@x0; c0=M@c0; terminal.append(float(np.linalg.norm(x0-c0)))
28 return {'q_exact':float(q), 'ratios':ratios, 'distances':ds,
29 'geometric_bound_holds':bool(all(ds[i] <= q**(i+1)*np.linalg.norm(np.array([.37,-.22]))+1e-10 for i in range(6))),
30 'r':r,'fallback_distances':terminal,'stays_in_terminal_set':bool(max(terminal)<=r)}
31
32# A tiny flat MLP, allowing differentiable parameter unrolling.
33def unpack(w):
34 a=w[:32].reshape(2,16); b=w[32:48].reshape(16); c=w[48:64].reshape(16,1); d=w[64:65]
35 return a,b,c,d
36def net_loss(w,x,y):
37 a,b,c,d=unpack(w); h=torch.tanh(x@a+b); logits=(h@c+d).squeeze(1)
38 return torch.nn.functional.binary_cross_entropy_with_logits(logits,y)
39def adam_step(w,g,m,v,t,lr=.01):
40 b1=.9; b2=.999; eps=1e-8
41 m=b1*m+(1-b1)*g; v=b2*v+(1-b2)*g*g
42 mh=m/(1-b1**t); vh=v/(1-b2**t)
43 return w-lr*mh/(torch.sqrt(vh)+eps),m,v
44
45def make_data(n=512):
46 rng=np.random.RandomState(SEED+9); x=rng.randn(n,2).astype('float32')
47 y=(x[:,0]*x[:,1]>0).astype('float32')
48 x += .18*rng.randn(n,2).astype('float32')
49 return torch.tensor(x,device=device),torch.tensor(y,device=device)
50
51def train(kind, X, Y, steps=180, H=2, rho=2.0, inner_lr=.08):
52 torch.manual_seed(SEED+ (0 if kind=='adam' else 1))
53 w=(.15*torch.randn(65,device=device)).requires_grad_(); m=torch.zeros_like(w); v=torch.zeros_like(w)
54 losses=[]; terminal=[]; grad_evals=0
55 n=len(X)
56 for k in range(steps):
57 ix=slice((k*7)%n, (k*7)%n+96) if (k*7)%n+96<=n else slice(0,96)
58 x,y=X[ix],Y[ix]
59 if kind=='adam':
60 loss=net_loss(w,x,y); g=torch.autograd.grad(loss,w)[0]; grad_evals+=1
61 with torch.no_grad(): w,m,v=adam_step(w,g,m,v,k+1); w.requires_grad_()
62 losses.append(float(loss.detach())); continue
63 # Guide rollout: two Adam updates, using current minibatch gradients as a cheap policy.
64 loss=net_loss(w,x,y); g=torch.autograd.grad(loss,w)[0]; grad_evals+=1
65 with torch.no_grad():
66 d0=-.01*g
67 c1=w.detach()+d0
68 # same guide policy at c1; this is deliberately cheap and stateless
69 c1r=c1.detach().requires_grad_(); l1=net_loss(c1r,x,y); g1=torch.autograd.grad(l1,c1r)[0]; grad_evals+=1
70 with torch.no_grad():
71 d1=-.01*g1.detach(); center=c1+d1
72 # Candidate controls initialized by guide, optimized against two losses + terminal penalty.
73 v0=d0.detach().clone().requires_grad_(); v1=d1.detach().clone().requires_grad_()
74 for _ in range(1):
75 p1=w.detach()+v0; p2=p1+v1
76 J=net_loss(w.detach(),x,y)+net_loss(p1,x,y)+net_loss(p2,x,y)+rho*.5*torch.sum((p2-center.detach())**2)
77 q0,q1=torch.autograd.grad(J,(v0,v1)); grad_evals+=2
78 with torch.no_grad(): v0-=inner_lr*q0; v1-=inner_lr*q1
79 v0.requires_grad_(); v1.requires_grad_()
80 with torch.no_grad():
81 deviation=torch.linalg.vector_norm((w.detach()+v0.detach()+v1.detach())-center.detach())
82 # Trust-region projection of the endpoint; preserve the first control direction.
83 radius=.035
84 if deviation>radius:
85 scale=radius/(deviation+1e-12)
86 v0.mul_(scale); v1.mul_(scale)
87 endpoint=(w.detach()+v0.detach()+v1.detach())
88 actual_deviation=torch.linalg.vector_norm(endpoint-center.detach())
89 w=w.detach()+v0.detach(); w.requires_grad_()
90 losses.append(float(loss.detach())); terminal.append(float(actual_deviation.detach()))
91 with torch.no_grad(): final=float(net_loss(w,X,Y))
92 return {'final_full_loss':final,'mean_last20':float(np.mean(losses[-20:])),'initial_loss':losses[0],
93 'max_terminal_deviation':float(max(terminal) if terminal else 0),'median_terminal_deviation':float(np.median(terminal) if terminal else 0),
94 'gradient_evaluations':grad_evals,'loss_curve':losses}
95
96def main():
97 check=contraction_check(); X,Y=make_data()
98 t=time.time(); base=train('adam',X,Y,steps=180); base_equal_budget=train('adam',X,Y,steps=720); trust=train('trust',X,Y,steps=180); elapsed=time.time()-t
99 out={'device':str(device),'seconds':elapsed,'contraction_check':check,'baseline':base,'baseline_equal_budget':base_equal_budget,'trust_region':trust}
100 with open('results.json','w') as f: json.dump(out,f,indent=2)
101 print(json.dumps({k:v for k,v in out.items() if k not in ('baseline','trust_region')} ,indent=2))
102 print(json.dumps({'baseline':{k:v for k,v in base.items() if k!='loss_curve'},'baseline_equal_budget':{k:v for k,v in base_equal_budget.items() if k!='loss_curve'},'trust_region':{k:v for k,v in trust.items() if k!='loss_curve'}},indent=2))
103if __name__=='__main__': main()