Finite-horizon Lyapunov regularization for neural updates / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 17
  6np.random.seed(SEED); random.seed(SEED)
  7
  8
  9def finite_horizon_sweep():
 10    # z_{k+1}=a z_k, V=.5 z^2. Theory: V_{k+M}/V_k=a^(2M),
 11    # and alpha-contraction boundary |a|=(1-alpha)^(1/(2M)).
 12    alpha=0.10
 13    rows=[]
 14    for M in (2,4,8):
 15        pred=(1-alpha)**(1/(2*M))
 16        aa=np.linspace(0.80,1.06,261)
 17        ratios=aa**(2*M)
 18        ok=ratios <= 1-alpha
 19        obs=aa[np.where(ok)[0][-1]]
 20        rows.append({'M':M,'predicted_boundary':pred,'observed_grid_boundary':float(obs),
 21                     'abs_error':float(abs(obs-pred))})
 22    # Scaling prediction: log(V_M/V_0)=2M log|a|.
 23    scale=[]
 24    for a in (0.82,0.90,0.98):
 25        vals=[]
 26        for M in (1,2,4,8):
 27            vals.append((M, float(a**(2*M)), float(2*M*math.log(a))))
 28        scale.append({'a':a,'values':vals})
 29    return rows, scale
 30
 31
 32def stochastic_mismatch_sweep():
 33    # For a=0, z_{k+1}=noise, residual D=V_{k+M}-V_k+alpha V_k.
 34    # E|D| scales quadratically with noise amplitude, as V is quadratic.
 35    alpha=.1; M=4; n=12000; burn=200
 36    out=[]
 37    for sigma in (.02,.05,.10,.20):
 38        rng=np.random.default_rng(100+int(sigma*1000))
 39        z=rng.normal(0,sigma,size=n+M+1)
 40        V=.5*z*z
 41        D=V[M:]-V[:-M]+alpha*V[:-M]
 42        d=D[burn:]
 43        ema=0.; beta=.9; eps=[]
 44        for x in np.abs(d):
 45            ema=beta*ema+(1-beta)*x; eps.append(ema)
 46        out.append({'sigma':sigma,'mean_abs_residual':float(np.mean(np.abs(d))),
 47                    'mean_ema_last_half':float(np.mean(eps[len(eps)//2:])),
 48                    'normalized_by_sigma2':float(np.mean(np.abs(d))/sigma**2)})
 49    return out
 50
 51
 52def unstable_ema_sweep():
 53    # Prediction: the EMA allowance is nearly harmless for stationary residuals,
 54    # but unstable exponential growth outruns it and creates violations.
 55    alpha=.1; M=4; n=120; beta=.9; out=[]
 56    for a in (.90,.98,1.00,1.01,1.03,1.06):
 57        z=0.02; V=[.5*z*z]
 58        for k in range(n+M): z=a*z; V.append(.5*z*z)
 59        D=np.array(V[M:])-np.array(V[:-M])+alpha*np.array(V[:-M])
 60        ema=0.; viol=[]; eps=[]
 61        for d in D:
 62            ema=beta*ema+(1-beta)*abs(d)
 63            viol.append(d-ema>0); eps.append(ema)
 64        out.append({'a':a,'theory_ratio':a**(2*M),
 65                    'violation_rate_last_half':float(np.mean(viol[len(viol)//2:])),
 66                    'last_residual_over_ema':float(D[-1]/max(eps[-1],1e-30))})
 67    return out
 68
 69
 70def mlp_experiment():
 71    # Tiny digits MLP. The idea is implemented as an online delayed gradient
 72    # Lyapunov penalty; current gradient is differentiable, old gradient and EMA
 73    # allowance are detached. This is a practical proxy for the optimizer test.
 74    import torch
 75    from sklearn.datasets import load_digits
 76    from sklearn.model_selection import train_test_split
 77    torch.manual_seed(SEED); np.random.seed(SEED)
 78    device='cuda' if torch.cuda.is_available() else 'cpu'
 79    try:
 80        X,y=load_digits(return_X_y=True)
 81        X=X.astype('float32')/16.; y=y.astype('int64')
 82        xt,xv,yt,yv=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
 83        def run(use_reg):
 84            torch.manual_seed(SEED)
 85            model=torch.nn.Sequential(torch.nn.Linear(64,64),torch.nn.Tanh(),torch.nn.Linear(64,10)).to(device)
 86            opt=torch.optim.SGD(model.parameters(),lr=.35,momentum=.0)
 87            lossfn=torch.nn.CrossEntropyLoss(); M=4; alpha=.1; lam=.03; beta=.9
 88            qs=[]; eps=0.; losses=[]; gnorms=[]; violations=[]
 89            order=np.arange(len(xt)); rng=np.random.default_rng(SEED)
 90            for epoch in range(12):
 91                rng.shuffle(order)
 92                for start in range(0,len(order),64):
 93                    ids=order[start:start+64]
 94                    xb=torch.tensor(xt[ids],device=device); yb=torch.tensor(yt[ids],device=device)
 95                    opt.zero_grad(set_to_none=True)
 96                    loss=lossfn(model(xb),yb)
 97                    grads=torch.autograd.grad(loss,tuple(model.parameters()),create_graph=use_reg,retain_graph=True)
 98                    flat=torch.cat([g.reshape(-1) for g in grads])
 99                    V=.5*(flat*flat).sum()
100                    q_det=flat.detach()
101                    penalty=torch.zeros((),device=device)
102                    if use_reg and len(qs)>=M:
103                        old=qs[-M]
104                        Vold=.5*(old*old).sum()
105                        D=V-Vold+alpha*Vold
106                        eps=beta*eps+(1-beta)*float(abs(D.detach()).cpu())
107                        r=D-eps
108                        penalty=lam*torch.relu(r).clamp(max=10.)**2/(1.+Vold)
109                        violations.append(float((r.detach()>0).cpu()))
110                    total=loss+penalty
111                    total.backward(); opt.step()
112                    qs.append(q_det)
113                    losses.append(float(loss.detach().cpu())); gnorms.append(float(torch.linalg.vector_norm(flat.detach()).cpu()))
114            with torch.no_grad():
115                pred=model(torch.tensor(xv,device=device)).argmax(1).cpu().numpy()
116            acc=float(np.mean(pred==yv)); tail=np.array(losses[-100:]); gn=np.array(gnorms[-100:])
117            return {'accuracy':acc,'final_loss':float(np.mean(tail)),
118                    'loss_std_tail':float(np.std(tail)),'grad_std_tail':float(np.std(gn)),
119                    'violation_rate':float(np.mean(violations)) if violations else 0.0}
120        try:
121            b=run(False); r=run(True)
122        except Exception:
123            if device=='cuda':
124                torch.cuda.empty_cache(); device='cpu'; b=run(False); r=run(True)
125            else: raise
126        return {'device':device,'baseline':b,'idea':r}
127    except Exception as e:
128        return {'error':repr(e)}
129
130if __name__=='__main__':
131    result={'finite_horizon_boundary':finite_horizon_sweep()[0],
132            'ratio_scaling':finite_horizon_sweep()[1],
133            'stochastic_mismatch':stochastic_mismatch_sweep(),
134            'unstable_ema':unstable_ema_sweep(),
135            'mlp':mlp_experiment()}
136    Path('results.json').write_text(json.dumps(result,indent=2))
137    print(json.dumps(result,indent=2))