import json, math, random from pathlib import Path import numpy as np def scalar_restart(tau=0.1, r=0.01, beta=1.0, steps=500): x = 1.0 prev = None clock = 0 rows = [] restarts = [] for n in range(steps): eta = tau * math.exp(r * clock) u = -eta * x if prev is not None and abs(u) >= beta * math.exp(r) * abs(prev): restarts.append(n) clock = 0 eta = tau u = -eta * x x += u rows.append((n, x, u, eta, clock)) prev = u clock += 1 return np.asarray(rows), restarts def quadratic_restart(eigs, x0, tau=0.05, r=0.01, beta=1.2, steps=1200): x = np.asarray(x0, dtype=float).copy() prev_norm = None clock = 0 restarts = [] history = [] for n in range(steps): eta = tau * math.exp(r * clock) u = -eta * np.asarray(eigs) * x norm = np.linalg.norm(u) restarted = False if prev_norm is not None and norm >= beta * math.exp(r) * prev_norm: clock = 0 eta = tau u = -eta * np.asarray(eigs) * x norm = np.linalg.norm(u) restarts.append(n) restarted = True x += u history.append((n, 0.5*np.sum(np.asarray(eigs)*x*x), norm, eta, clock, restarted)) prev_norm = norm clock += 1 return np.asarray(history, dtype=object), restarts def run_torch(seed=7, epochs=12, batch_size=128): try: import torch from torch import nn from torch.utils.data import DataLoader, TensorDataset from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split device = 'cuda' if torch.cuda.is_available() else 'cpu' torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) d = load_digits() X = (d.data.astype('float32') / 16.0) y = d.target.astype('int64') xa, xv, ya, yv = train_test_split(X, y, test_size=0.25, random_state=seed, stratify=y) train = DataLoader(TensorDataset(torch.tensor(xa), torch.tensor(ya)), batch_size=batch_size, shuffle=True) valid = DataLoader(TensorDataset(torch.tensor(xv), torch.tensor(yv)), batch_size=batch_size) def make(): return nn.Sequential(nn.Linear(64, 64), nn.Tanh(), nn.Linear(64, 10)).to(device) def evaluate(model): model.eval(); loss_sum=correct=total=0 with torch.no_grad(): for xb,yb in valid: z=model(xb.to(device)); loss_sum += nn.functional.cross_entropy(z,yb.to(device), reduction='sum').item() correct += (z.argmax(1)==yb.to(device)).sum().item(); total += len(yb) return loss_sum/total, correct/total def train_one(kind): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) model=make(); base=0.01; r=0.012; beta=1.3; clock=0; prev=None; restarts=0; peak=[]; losses=[] if kind == 'fixed': opt=torch.optim.SGD(model.parameters(), lr=0.08) for ep in range(epochs): model.train() for xb,yb in train: xb,yb=xb.to(device),yb.to(device); opt=None if kind!='fixed' else opt model.zero_grad(set_to_none=True); loss=nn.functional.cross_entropy(model(xb),yb); loss.backward() if kind == 'fixed': opt.step(); continue grads=[p.grad for p in model.parameters() if p.grad is not None] gn=torch.sqrt(sum((g*g).sum() for g in grads)).item() eta=base*math.exp(r*clock); un=eta*gn if prev is not None and un >= beta*math.exp(r)*prev: restarts += 1; clock=0; eta=base; un=eta*gn with torch.no_grad(): for p in model.parameters(): if p.grad is not None: p.add_(p.grad, alpha=-eta) prev=un; peak.append(eta); clock += 1 vl,va=evaluate(model); losses.append((ep+1,vl,va)) return {'final_loss':losses[-1][1], 'final_acc':losses[-1][2], 'restarts':restarts, 'max_lr':max(peak) if peak else 0, 'losses':losses} return {'device':device, 'fixed':train_one('fixed'), 'restart':train_one('restart')} except Exception as e: return {'error': repr(e)} def main(): # Directly verify the stated scalar product formula x_{n+1}=prod_{k=0}^n(1-tau exp(rk)). rows,_=scalar_restart(tau=0.1, r=0.01, beta=1e9, steps=30) prod=1.0; max_err=0.0 for n in range(30): prod *= (1-0.1*math.exp(0.01*n)) max_err=max(max_err, abs(rows[n,1]-prod)) # Compare no-restart exponential growth against restart on an ill-conditioned quadratic. no=quadratic_restart([1.,100.], [1.,1.], tau=.005, r=.01, beta=1e9, steps=500) yes=quadratic_restart([1.,100.], [1.,1.], tau=.005, r=.01, beta=1.2, steps=500) out={'formula_max_abs_error':max_err, 'scalar_first_restart':scalar_restart()[1][:5], 'quadratic_no_restart_final_loss':float(no[0][-1,1]), 'quadratic_restart_final_loss':float(yes[0][-1,1]), 'quadratic_restart_count':len(yes[1]), 'digits':run_torch()} Path('results.json').write_text(json.dumps(out, indent=2)) print(json.dumps(out, indent=2)) if __name__=='__main__': main()