Bennett-whitened gradient trust region / experiment.py

Mechanism failed

Raw ⬇ ZIP
  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()