import json, math, random, time import numpy as np import torch SEED=3072 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type=='cuda': torch.cuda.get_device_properties(0) except Exception: device=torch.device('cpu') def contraction_check(): # Stable linear guiding policy: c+ = (I-alpha A)c, whose exact q is spectral norm. A=np.diag([1.,3.]); alpha=.25 M=np.eye(2)-alpha*A q=np.linalg.norm(M,2) c=np.array([.8,-.5]); x=c+np.array([.37,-.22]); initial=np.linalg.norm(x-c) ratios=[]; ds=[] for _ in range(6): x=M@x; c=M@c d=np.linalg.norm(x-c); ds.append(float(d)); ratios.append(float(d/initial)); initial=d # terminal feasibility: perturbation inside radius r remains inside after fallback r=.5; perturb=np.array([.4,.3]); c0=np.zeros(2); x0=c0+perturb terminal=[] for _ in range(5): x0=M@x0; c0=M@c0; terminal.append(float(np.linalg.norm(x0-c0))) return {'q_exact':float(q), 'ratios':ratios, 'distances':ds, 'geometric_bound_holds':bool(all(ds[i] <= q**(i+1)*np.linalg.norm(np.array([.37,-.22]))+1e-10 for i in range(6))), 'r':r,'fallback_distances':terminal,'stays_in_terminal_set':bool(max(terminal)<=r)} # A tiny flat MLP, allowing differentiable parameter unrolling. def unpack(w): a=w[:32].reshape(2,16); b=w[32:48].reshape(16); c=w[48:64].reshape(16,1); d=w[64:65] return a,b,c,d def net_loss(w,x,y): a,b,c,d=unpack(w); h=torch.tanh(x@a+b); logits=(h@c+d).squeeze(1) return torch.nn.functional.binary_cross_entropy_with_logits(logits,y) def adam_step(w,g,m,v,t,lr=.01): b1=.9; b2=.999; eps=1e-8 m=b1*m+(1-b1)*g; v=b2*v+(1-b2)*g*g mh=m/(1-b1**t); vh=v/(1-b2**t) return w-lr*mh/(torch.sqrt(vh)+eps),m,v def make_data(n=512): rng=np.random.RandomState(SEED+9); x=rng.randn(n,2).astype('float32') y=(x[:,0]*x[:,1]>0).astype('float32') x += .18*rng.randn(n,2).astype('float32') return torch.tensor(x,device=device),torch.tensor(y,device=device) def train(kind, X, Y, steps=180, H=2, rho=2.0, inner_lr=.08): torch.manual_seed(SEED+ (0 if kind=='adam' else 1)) w=(.15*torch.randn(65,device=device)).requires_grad_(); m=torch.zeros_like(w); v=torch.zeros_like(w) losses=[]; terminal=[]; grad_evals=0 n=len(X) for k in range(steps): ix=slice((k*7)%n, (k*7)%n+96) if (k*7)%n+96<=n else slice(0,96) x,y=X[ix],Y[ix] if kind=='adam': loss=net_loss(w,x,y); g=torch.autograd.grad(loss,w)[0]; grad_evals+=1 with torch.no_grad(): w,m,v=adam_step(w,g,m,v,k+1); w.requires_grad_() losses.append(float(loss.detach())); continue # Guide rollout: two Adam updates, using current minibatch gradients as a cheap policy. loss=net_loss(w,x,y); g=torch.autograd.grad(loss,w)[0]; grad_evals+=1 with torch.no_grad(): d0=-.01*g c1=w.detach()+d0 # same guide policy at c1; this is deliberately cheap and stateless c1r=c1.detach().requires_grad_(); l1=net_loss(c1r,x,y); g1=torch.autograd.grad(l1,c1r)[0]; grad_evals+=1 with torch.no_grad(): d1=-.01*g1.detach(); center=c1+d1 # Candidate controls initialized by guide, optimized against two losses + terminal penalty. v0=d0.detach().clone().requires_grad_(); v1=d1.detach().clone().requires_grad_() for _ in range(1): p1=w.detach()+v0; p2=p1+v1 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) q0,q1=torch.autograd.grad(J,(v0,v1)); grad_evals+=2 with torch.no_grad(): v0-=inner_lr*q0; v1-=inner_lr*q1 v0.requires_grad_(); v1.requires_grad_() with torch.no_grad(): deviation=torch.linalg.vector_norm((w.detach()+v0.detach()+v1.detach())-center.detach()) # Trust-region projection of the endpoint; preserve the first control direction. radius=.035 if deviation>radius: scale=radius/(deviation+1e-12) v0.mul_(scale); v1.mul_(scale) endpoint=(w.detach()+v0.detach()+v1.detach()) actual_deviation=torch.linalg.vector_norm(endpoint-center.detach()) w=w.detach()+v0.detach(); w.requires_grad_() losses.append(float(loss.detach())); terminal.append(float(actual_deviation.detach())) with torch.no_grad(): final=float(net_loss(w,X,Y)) return {'final_full_loss':final,'mean_last20':float(np.mean(losses[-20:])),'initial_loss':losses[0], 'max_terminal_deviation':float(max(terminal) if terminal else 0),'median_terminal_deviation':float(np.median(terminal) if terminal else 0), 'gradient_evaluations':grad_evals,'loss_curve':losses} def main(): check=contraction_check(); X,Y=make_data() 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 out={'device':str(device),'seconds':elapsed,'contraction_check':check,'baseline':base,'baseline_equal_budget':base_equal_budget,'trust_region':trust} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps({k:v for k,v in out.items() if k not in ('baseline','trust_region')} ,indent=2)) 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)) if __name__=='__main__': main()