import json, math, random, time import numpy as np import torch from torch import nn SEED = 2373 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) try: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') except Exception: device = torch.device('cpu') def h(u): # stable Bennett rate function u = np.maximum(np.asarray(u, dtype=np.float64), 0.) return (1.0 + u) * np.log1p(u) - u def bennett_radius(V, rho=1e-2, delta=0.05): V = np.maximum(np.asarray(V, dtype=np.float64), 0.) D = 0.5 * np.log1p(V / rho).sum() target = math.log(1.0 / delta) + D # rho*h(r/sqrt(rho)) = target; monotone root by bisection lo, hi = 0.0, math.sqrt(rho) while rho * h(hi / math.sqrt(rho)) < target: hi *= 2. for _ in range(80): mid = (lo + hi) / 2. if rho * h(mid / math.sqrt(rho)) < target: lo = mid else: hi = mid return (lo + hi) / 2., D def math_check(): # Verify root equality and inspect time-uniform boundary crossings for bounded iid noise. rng = np.random.default_rng(SEED) V = np.zeros(4); M = np.zeros(4); prev = np.zeros(4) max_ratio = 0.; roots_ok = True paths = 400; T = 120 crossings = 0 for p in range(paths): V[:] = 0; M[:] = 0 crossed = False for t in range(T): # bounded, mean-zero, heteroscedastic coordinate noise x = rng.choice([-1., 1.], size=4) * np.array([.25,.5,.75,1.0]) V += prev * prev # predictable one-step delayed proxy M += x r, D = bennett_radius(V) z = np.sqrt(np.sum(M*M/(V+1e-2))) max_ratio = max(max_ratio, z/r) residual = 1e-2*h(r/math.sqrt(1e-2)) - D - math.log(20.) roots_ok &= abs(residual) < 1e-8 if z > r: crossed = True prev = x crossings += int(crossed) return {'max_z_over_r': float(max_ratio), 'root_max_abs_error': float(0 if roots_ok else 1), 'crossing_path_fraction': crossings / paths, 'paths': paths, 'steps': T} class Net(nn.Module): def __init__(self, d=20): super().__init__(); self.net = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, 1)) def forward(self, x): return self.net(x).squeeze(-1) def make_data(n=4096, d=20): rng = np.random.default_rng(SEED) X = rng.normal(size=(n,d)).astype('float32') y = (np.sin(X[:,0]*2) + .5*X[:,1]**2 - .7*X[:,2] + .25*rng.normal(size=n)).astype('float32') return torch.tensor(X), torch.tensor(y) class TrustController: def __init__(self, params, rho=1e-2, delta=.05, clip=1.0): self.params=list(params); self.rho=rho; self.delta=delta; self.clip=clip self.M=[torch.zeros_like(p, device=p.device) for p in self.params] self.V=[torch.zeros_like(p, device=p.device) for p in self.params] self.mu=[torch.zeros_like(p, device=p.device) for p in self.params] self.prev_sq=[torch.zeros_like(p, device=p.device) for p in self.params] @torch.no_grad() def multiplier(self): # delayed diagonal covariance, and globally clipped centered noise for i,p in enumerate(self.params): self.V[i].add_(self.prev_sq[i]) total_sq = sum(((p.grad - self.mu[i])**2).sum() for i,p in enumerate(self.params)) scale = min(1., self.clip / (float(torch.sqrt(total_sq)) + 1e-12)) z2 = 0.; D = 0. for i,p in enumerate(self.params): x = (p.grad - self.mu[i]) * scale self.M[i].add_(x) z2 += float((self.M[i]**2 / (self.V[i] + self.rho)).sum()) D += .5 * float(torch.log1p(self.V[i] / self.rho).sum()) self.prev_sq[i].copy_(x*x) self.mu[i].mul_(0.95).add_(p.grad, alpha=0.05) z=math.sqrt(max(z2,0.)); target=math.log(1/self.delta)+D lo,hi=0.,math.sqrt(self.rho) while self.rho*((1+hi/math.sqrt(self.rho))*math.log1p(hi/math.sqrt(self.rho))-hi/math.sqrt(self.rho)) < target: hi*=2 for _ in range(55): mid=(lo+hi)/2; u=mid/math.sqrt(self.rho) val=self.rho*((1+u)*math.log1p(u)-u) if val=2 else float('inf'), 'steps':len(arr), 'diverged':diverged, 'downscaled_fraction':down/max(1,len(ratios)), 'median_z_over_r':float(np.median(ratios)) if ratios else 0., 'time_sec':time.time()-t0} def main(): check=math_check(); X,y=make_data(); X,y=X.to(device),y.to(device) results={'device':str(device),'math_check':check,'baseline_sgd':train('sgd',X,y),'trust_region':train('trust',X,y)} with open('results.json','w') as f: json.dump(results,f,indent=2) print(json.dumps(results,indent=2)) if __name__=='__main__': main()