Pole-radius tuning for gradient tracking / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, math, copy, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10WORKERS = 4
 11EPOCHS = 24
 12CANDIDATES = [0.001, 0.003, 0.01]
 13
 14# Symmetric stochastic ring: each worker keeps half its value and averages neighbors.
 15def ring_matrix(n=4):
 16    W = np.zeros((n, n), dtype=float)
 17    for i in range(n):
 18        W[i, i] = 0.5
 19        W[i, (i-1) % n] += 0.25
 20        W[i, (i+1) % n] += 0.25
 21    return W
 22
 23W = ring_matrix()
 24LAMS = np.linalg.eigvalsh(W)[0:-1]  # exclude consensus eigenvalue 1
 25
 26
 27def roots(lam, q, alpha):
 28    return np.roots([1.0, q * alpha - 2.0 * lam,
 29                     lam * lam - q * alpha])
 30
 31
 32def pole_radius(alpha, lams=LAMS, qs=np.linspace(0.1, 10.0, 17)):
 33    return float(max(abs(r) for lam in lams for q in qs for r in roots(lam, q, alpha)))
 34
 35
 36def flat(tensors):
 37    return torch.cat([x.detach().reshape(-1) for x in tensors])
 38
 39
 40def grad_at(model, x, y):
 41    model.zero_grad(set_to_none=True)
 42    loss = ((model(x) - y) ** 2).mean()
 43    loss.backward()
 44    return [p.grad.detach().clone() for p in model.parameters()], float(loss.detach())
 45
 46
 47def curvature_interval(models, xs, ys):
 48    # Empirical gradient-difference Rayleigh quotients, as specified in the idea.
 49    vals = []
 50    for m, x, y in zip(models, xs, ys):
 51        base = [p.detach().clone() for p in m.parameters()]
 52        g0, _ = grad_at(m, x, y)
 53        torch.manual_seed(9137)
 54        delta = [0.01 * torch.randn_like(p) for p in m.parameters()]
 55        with torch.no_grad():
 56            for p, d in zip(m.parameters(), delta): p.add_(d)
 57        g1, _ = grad_at(m, x, y)
 58        with torch.no_grad():
 59            for p, b in zip(m.parameters(), base): p.copy_(b)
 60        dx = flat(delta); dg = flat([a-b for a,b in zip(g1,g0)])
 61        q = float(torch.dot(dg, dx) / (torch.dot(dx, dx) + 1e-12))
 62        if np.isfinite(q) and q > 1e-6: vals.append(q)
 63    if not vals: return 0.1, 10.0
 64    return max(0.05, float(np.quantile(vals, .2))), max(0.1, float(np.quantile(vals, .8)))
 65
 66
 67def choose_alpha(models, xs, ys):
 68    mu, hi = curvature_interval(models, xs, ys)
 69    qs = np.linspace(mu, hi, 17)
 70    scores = {a: pole_radius(a, LAMS, qs) for a in CANDIDATES}
 71    alpha = min(CANDIDATES, key=lambda a: scores[a])
 72    return alpha, mu, hi, scores
 73
 74
 75def train_diging(seed, alpha, tune=False, return_signature=False):
 76    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 77    d = get_dataset('tabular', seed, n_train=400, n_test=200)
 78    # Same architecture and same partition for both systems.
 79    xparts = list(torch.chunk(d['xtr'], WORKERS))
 80    yparts = list(torch.chunk(d['ytr'], WORKERS))
 81    models = [make_model('mlp_tiny', d['input_shape'], d['out_dim']) for _ in range(WORKERS)]
 82    xs = xparts; ys = yparts
 83    grads = []; losses = []
 84    for m,x,y in zip(models,xs,ys):
 85        g, loss = grad_at(m,x,y); grads.append(g); losses.append(loss)
 86    selected = alpha; mu=hi=None; scores=None
 87    if tune:
 88        selected, mu, hi, scores = choose_alpha(models, xs, ys)
 89    # y is the gradient tracker, represented as per-model parameter lists.
 90    tracker = [[g.clone() for g in gg] for gg in grads]
 91    obs_dis = []
 92    for _ in range(EPOCHS):
 93        old_params = [[p.detach().clone() for p in m.parameters()] for m in models]
 94        old_tracker = [[z.detach().clone() for z in tt] for tt in tracker]
 95        with torch.no_grad():
 96            for i,m in enumerate(models):
 97                for k,p in enumerate(m.parameters()):
 98                    comm = sum(W[i,j] * old_params[j][k] for j in range(WORKERS))
 99                    p.copy_(comm - selected * old_tracker[i][k])
100        new_grads=[]
101        for m,x,y in zip(models,xs,ys):
102            g, loss = grad_at(m,x,y); new_grads.append(g); losses.append(loss)
103        with torch.no_grad():
104            for i in range(WORKERS):
105                for k in range(len(tracker[i])):
106                    comm = sum(W[i,j] * old_tracker[j][k] for j in range(WORKERS))
107                    tracker[i][k].copy_(comm + new_grads[i][k] - grads[i][k])
108        grads = new_grads
109        tv = torch.stack([flat(t) for t in tracker])
110        obs_dis.append(float(torch.linalg.norm(tv - tv.mean(0,keepdim=True))))
111    with torch.no_grad():
112        pred = torch.stack([m(d['xte']) for m in models]).mean(0)
113        metric = float(((pred-d['yte'])**2).mean())
114    sig = None
115    if return_signature:
116        ratios = [obs_dis[i+1]/(obs_dis[i]+1e-12) for i in range(len(obs_dis)-1)]
117        observed = float(np.median(ratios[-8:])) if ratios else float('nan')
118        predicted = pole_radius(selected, LAMS, np.linspace(mu or .1, hi or 10., 17))
119        sig = {'alpha': selected, 'mu_est': mu, 'L_est': hi,
120               'predicted_pole_radius': predicted,
121               'observed_tracker_disagreement_ratio': observed,
122               'nondesignated_training_model': 'mlp_tiny_worker_ensemble',
123               'confirmed': bool(np.isfinite(observed) and abs(observed-predicted) <= max(.08, .25*predicted))}
124    return metric, sig
125
126
127def make_baseline(cfg):
128    return lambda seed: train_diging(seed, float(cfg['lr']), tune=False)[0]
129
130def make_idea(cfg):
131    # cfg lr is the shared candidate grid; tuning chooses the minimax candidate
132    return lambda seed: train_diging(seed, float(cfg['lr']), tune=True)[0]
133
134if __name__ == '__main__':
135    grid = [{'lr': a} for a in CANDIDATES]
136    base = sweep_baseline(make_baseline, grid, seeds=(0,1,2,3))
137    # Evaluate every idea candidate on all eight paired seeds; select best full mean.
138    idea_runs=[]
139    for cfg in grid:
140        r=evaluate(make_idea(cfg), seeds=SEEDS)
141        idea_runs.append((cfg,r))
142    idea_cfg, idea = min(idea_runs, key=lambda z:z[1]['mean'])
143    # Signature is measured on trained systems, not on the analytic toy alone.
144    _, signature = train_diging(0, idea_cfg['lr'], tune=True, return_signature=True)
145    report = make_report('tabular', 'mlp_tiny', base, idea,
146                         {'track_structure':'optimizer/decentralized regression', **signature,
147                          'pole_spectrum': LAMS.tolist(),
148                          'idea_grid': grid, 'selected_idea_cfg': idea_cfg,
149                          'candidate_full_means': {str(c['lr']): r['mean'] for c,r in idea_runs}})
150    Path('bench_report.json').write_text(json.dumps(report, indent=2))
151    print(json.dumps(report, indent=2))