Exponentially Growing Learning Rate with Update-Norm Restarts / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5
  6def scalar_restart(tau=0.1, r=0.01, beta=1.0, steps=500):
  7    x = 1.0
  8    prev = None
  9    clock = 0
 10    rows = []
 11    restarts = []
 12    for n in range(steps):
 13        eta = tau * math.exp(r * clock)
 14        u = -eta * x
 15        if prev is not None and abs(u) >= beta * math.exp(r) * abs(prev):
 16            restarts.append(n)
 17            clock = 0
 18            eta = tau
 19            u = -eta * x
 20        x += u
 21        rows.append((n, x, u, eta, clock))
 22        prev = u
 23        clock += 1
 24    return np.asarray(rows), restarts
 25
 26
 27def quadratic_restart(eigs, x0, tau=0.05, r=0.01, beta=1.2, steps=1200):
 28    x = np.asarray(x0, dtype=float).copy()
 29    prev_norm = None
 30    clock = 0
 31    restarts = []
 32    history = []
 33    for n in range(steps):
 34        eta = tau * math.exp(r * clock)
 35        u = -eta * np.asarray(eigs) * x
 36        norm = np.linalg.norm(u)
 37        restarted = False
 38        if prev_norm is not None and norm >= beta * math.exp(r) * prev_norm:
 39            clock = 0
 40            eta = tau
 41            u = -eta * np.asarray(eigs) * x
 42            norm = np.linalg.norm(u)
 43            restarts.append(n)
 44            restarted = True
 45        x += u
 46        history.append((n, 0.5*np.sum(np.asarray(eigs)*x*x), norm, eta, clock, restarted))
 47        prev_norm = norm
 48        clock += 1
 49    return np.asarray(history, dtype=object), restarts
 50
 51
 52def run_torch(seed=7, epochs=12, batch_size=128):
 53    try:
 54        import torch
 55        from torch import nn
 56        from torch.utils.data import DataLoader, TensorDataset
 57        from sklearn.datasets import load_digits
 58        from sklearn.model_selection import train_test_split
 59        device = 'cuda' if torch.cuda.is_available() else 'cpu'
 60        torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 61        d = load_digits()
 62        X = (d.data.astype('float32') / 16.0)
 63        y = d.target.astype('int64')
 64        xa, xv, ya, yv = train_test_split(X, y, test_size=0.25, random_state=seed, stratify=y)
 65        train = DataLoader(TensorDataset(torch.tensor(xa), torch.tensor(ya)), batch_size=batch_size, shuffle=True)
 66        valid = DataLoader(TensorDataset(torch.tensor(xv), torch.tensor(yv)), batch_size=batch_size)
 67        def make():
 68            return nn.Sequential(nn.Linear(64, 64), nn.Tanh(), nn.Linear(64, 10)).to(device)
 69        def evaluate(model):
 70            model.eval(); loss_sum=correct=total=0
 71            with torch.no_grad():
 72                for xb,yb in valid:
 73                    z=model(xb.to(device)); loss_sum += nn.functional.cross_entropy(z,yb.to(device), reduction='sum').item()
 74                    correct += (z.argmax(1)==yb.to(device)).sum().item(); total += len(yb)
 75            return loss_sum/total, correct/total
 76        def train_one(kind):
 77            torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 78            model=make(); base=0.01; r=0.012; beta=1.3; clock=0; prev=None; restarts=0; peak=[]; losses=[]
 79            if kind == 'fixed':
 80                opt=torch.optim.SGD(model.parameters(), lr=0.08)
 81            for ep in range(epochs):
 82                model.train()
 83                for xb,yb in train:
 84                    xb,yb=xb.to(device),yb.to(device); opt=None if kind!='fixed' else opt
 85                    model.zero_grad(set_to_none=True); loss=nn.functional.cross_entropy(model(xb),yb); loss.backward()
 86                    if kind == 'fixed': opt.step(); continue
 87                    grads=[p.grad for p in model.parameters() if p.grad is not None]
 88                    gn=torch.sqrt(sum((g*g).sum() for g in grads)).item()
 89                    eta=base*math.exp(r*clock); un=eta*gn
 90                    if prev is not None and un >= beta*math.exp(r)*prev:
 91                        restarts += 1; clock=0; eta=base; un=eta*gn
 92                    with torch.no_grad():
 93                        for p in model.parameters():
 94                            if p.grad is not None: p.add_(p.grad, alpha=-eta)
 95                    prev=un; peak.append(eta); clock += 1
 96                vl,va=evaluate(model); losses.append((ep+1,vl,va))
 97            return {'final_loss':losses[-1][1], 'final_acc':losses[-1][2], 'restarts':restarts, 'max_lr':max(peak) if peak else 0, 'losses':losses}
 98        return {'device':device, 'fixed':train_one('fixed'), 'restart':train_one('restart')}
 99    except Exception as e:
100        return {'error': repr(e)}
101
102
103def main():
104    # Directly verify the stated scalar product formula x_{n+1}=prod_{k=0}^n(1-tau exp(rk)).
105    rows,_=scalar_restart(tau=0.1, r=0.01, beta=1e9, steps=30)
106    prod=1.0; max_err=0.0
107    for n in range(30):
108        prod *= (1-0.1*math.exp(0.01*n))
109        max_err=max(max_err, abs(rows[n,1]-prod))
110    # Compare no-restart exponential growth against restart on an ill-conditioned quadratic.
111    no=quadratic_restart([1.,100.], [1.,1.], tau=.005, r=.01, beta=1e9, steps=500)
112    yes=quadratic_restart([1.,100.], [1.,1.], tau=.005, r=.01, beta=1.2, steps=500)
113    out={'formula_max_abs_error':max_err,
114         'scalar_first_restart':scalar_restart()[1][:5],
115         'quadratic_no_restart_final_loss':float(no[0][-1,1]),
116         'quadratic_restart_final_loss':float(yes[0][-1,1]),
117         'quadratic_restart_count':len(yes[1]),
118         'digits':run_torch()}
119    Path('results.json').write_text(json.dumps(out, indent=2))
120    print(json.dumps(out, indent=2))
121
122if __name__=='__main__': main()