Exponentially Growing Learning Rate with Update-Norm Restarts / experiment.py
Mechanism confirmed, baseline not beaten
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()