Bennett-whitened gradient trust region / experiment.py
Mechanism failed
1import json, math, random, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 2373
7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
8try:
9 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10except Exception:
11 device = torch.device('cpu')
12
13
14def h(u):
15 # stable Bennett rate function
16 u = np.maximum(np.asarray(u, dtype=np.float64), 0.)
17 return (1.0 + u) * np.log1p(u) - u
18
19
20def bennett_radius(V, rho=1e-2, delta=0.05):
21 V = np.maximum(np.asarray(V, dtype=np.float64), 0.)
22 D = 0.5 * np.log1p(V / rho).sum()
23 target = math.log(1.0 / delta) + D
24 # rho*h(r/sqrt(rho)) = target; monotone root by bisection
25 lo, hi = 0.0, math.sqrt(rho)
26 while rho * h(hi / math.sqrt(rho)) < target:
27 hi *= 2.
28 for _ in range(80):
29 mid = (lo + hi) / 2.
30 if rho * h(mid / math.sqrt(rho)) < target: lo = mid
31 else: hi = mid
32 return (lo + hi) / 2., D
33
34
35def math_check():
36 # Verify root equality and inspect time-uniform boundary crossings for bounded iid noise.
37 rng = np.random.default_rng(SEED)
38 V = np.zeros(4); M = np.zeros(4); prev = np.zeros(4)
39 max_ratio = 0.; roots_ok = True
40 paths = 400; T = 120
41 crossings = 0
42 for p in range(paths):
43 V[:] = 0; M[:] = 0
44 crossed = False
45 for t in range(T):
46 # bounded, mean-zero, heteroscedastic coordinate noise
47 x = rng.choice([-1., 1.], size=4) * np.array([.25,.5,.75,1.0])
48 V += prev * prev # predictable one-step delayed proxy
49 M += x
50 r, D = bennett_radius(V)
51 z = np.sqrt(np.sum(M*M/(V+1e-2)))
52 max_ratio = max(max_ratio, z/r)
53 residual = 1e-2*h(r/math.sqrt(1e-2)) - D - math.log(20.)
54 roots_ok &= abs(residual) < 1e-8
55 if z > r: crossed = True
56 prev = x
57 crossings += int(crossed)
58 return {'max_z_over_r': float(max_ratio), 'root_max_abs_error': float(0 if roots_ok else 1),
59 'crossing_path_fraction': crossings / paths, 'paths': paths, 'steps': T}
60
61
62class Net(nn.Module):
63 def __init__(self, d=20):
64 super().__init__(); self.net = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, 1))
65 def forward(self, x): return self.net(x).squeeze(-1)
66
67
68def make_data(n=4096, d=20):
69 rng = np.random.default_rng(SEED)
70 X = rng.normal(size=(n,d)).astype('float32')
71 y = (np.sin(X[:,0]*2) + .5*X[:,1]**2 - .7*X[:,2] + .25*rng.normal(size=n)).astype('float32')
72 return torch.tensor(X), torch.tensor(y)
73
74class TrustController:
75 def __init__(self, params, rho=1e-2, delta=.05, clip=1.0):
76 self.params=list(params); self.rho=rho; self.delta=delta; self.clip=clip
77 self.M=[torch.zeros_like(p, device=p.device) for p in self.params]
78 self.V=[torch.zeros_like(p, device=p.device) for p in self.params]
79 self.mu=[torch.zeros_like(p, device=p.device) for p in self.params]
80 self.prev_sq=[torch.zeros_like(p, device=p.device) for p in self.params]
81 @torch.no_grad()
82 def multiplier(self):
83 # delayed diagonal covariance, and globally clipped centered noise
84 for i,p in enumerate(self.params): self.V[i].add_(self.prev_sq[i])
85 total_sq = sum(((p.grad - self.mu[i])**2).sum() for i,p in enumerate(self.params))
86 scale = min(1., self.clip / (float(torch.sqrt(total_sq)) + 1e-12))
87 z2 = 0.; D = 0.
88 for i,p in enumerate(self.params):
89 x = (p.grad - self.mu[i]) * scale
90 self.M[i].add_(x)
91 z2 += float((self.M[i]**2 / (self.V[i] + self.rho)).sum())
92 D += .5 * float(torch.log1p(self.V[i] / self.rho).sum())
93 self.prev_sq[i].copy_(x*x)
94 self.mu[i].mul_(0.95).add_(p.grad, alpha=0.05)
95 z=math.sqrt(max(z2,0.)); target=math.log(1/self.delta)+D
96 lo,hi=0.,math.sqrt(self.rho)
97 while self.rho*((1+hi/math.sqrt(self.rho))*math.log1p(hi/math.sqrt(self.rho))-hi/math.sqrt(self.rho)) < target: hi*=2
98 for _ in range(55):
99 mid=(lo+hi)/2; u=mid/math.sqrt(self.rho)
100 val=self.rho*((1+u)*math.log1p(u)-u)
101 if val<target: lo=mid
102 else: hi=mid
103 r=(lo+hi)/2
104 return min(1., r/max(z,1e-12)), z/r
105
106
107def train(kind, X, y, epochs=12, batch=64, lr=.16):
108 torch.manual_seed(SEED); model=Net(X.shape[1]).to(device); opt=torch.optim.SGD(model.parameters(),lr=lr)
109 ctrl=TrustController(model.parameters()) if kind=='trust' else None
110 losses=[]; ratios=[]; down=0; diverged=False
111 order=np.arange(len(X)); t0=time.time()
112 for ep in range(epochs):
113 rng=np.random.default_rng(SEED+ep); rng.shuffle(order)
114 for start in range(0,len(X),batch):
115 idx=torch.tensor(order[start:start+batch], device=device)
116 opt.zero_grad(set_to_none=True); pred=model(X[idx]); loss=((pred-y[idx])**2).mean(); loss.backward()
117 if not torch.isfinite(loss): diverged=True; break
118 if ctrl:
119 mult,ratio=ctrl.multiplier(); ratios.append(ratio); down += int(mult<.999999)
120 for p in model.parameters(): p.grad.mul_(mult)
121 opt.step(); losses.append(float(loss.detach().cpu()))
122 if diverged: break
123 arr=np.asarray(losses)
124 return {'final_loss':float(arr[-1]) if len(arr) else float('inf'), 'best_loss':float(arr.min()) if len(arr) else float('inf'),
125 'loss_variance_last50':float(np.var(arr[-50:])) if len(arr)>=2 else float('inf'),
126 'steps':len(arr), 'diverged':diverged, 'downscaled_fraction':down/max(1,len(ratios)),
127 'median_z_over_r':float(np.median(ratios)) if ratios else 0., 'time_sec':time.time()-t0}
128
129
130def main():
131 check=math_check(); X,y=make_data(); X,y=X.to(device),y.to(device)
132 results={'device':str(device),'math_check':check,'baseline_sgd':train('sgd',X,y),'trust_region':train('trust',X,y)}
133 with open('results.json','w') as f: json.dump(results,f,indent=2)
134 print(json.dumps(results,indent=2))
135if __name__=='__main__': main()