Bounded predictive-gain optimizer / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
  9
 10TRACK = 'tabular'
 11MODEL = 'mlp_tiny'
 12SEEDS = tuple(range(8))
 13EPOCHS = 8
 14BATCH = 128
 15
 16
 17def seed_all(seed):
 18    random.seed(seed)
 19    np.random.seed(seed)
 20    torch.manual_seed(seed)
 21    if torch.cuda.is_available():
 22        torch.cuda.manual_seed_all(seed)
 23
 24
 25def math_check():
 26    eta, rho, g, a0 = 0.07, 0.4, 1.3, 0.8
 27    r = np.linspace(-1., 1., 5)
 28    observed = []
 29    for rr in r:
 30        q = -rr * g * g
 31        anew = (a0 - eta * q + rho * a0) / (1. + rho)
 32        observed.append(anew - a0)
 33    predicted = eta * r * g * g / (1. + rho)
 34    return {
 35        'prediction': 'delta_gain=eta*r*g^2/(1+rho) at gain=reference',
 36        'predicted_slope': float(eta * g * g / (1. + rho)),
 37        'observed_slope': float(np.polyfit(r, observed, 1)[0]),
 38        'max_abs_error': float(np.max(np.abs(np.asarray(observed) - predicted))),
 39        'pass': bool(np.max(np.abs(np.asarray(observed) - predicted)) < 1e-12)
 40    }
 41
 42
 43def train_one(seed, lr, wd=0.0, adaptive=False, gain_eta=0.08, rho=0.15,
 44              collect_signature=False):
 45    seed_all(seed)
 46    ds = get_dataset(TRACK, seed)
 47    net = make_model(MODEL, tuple(ds['input_shape']), ds['out_dim'])
 48    requested = 'cuda' if torch.cuda.is_available() else 'cpu'
 49    try:
 50        return _train(net, ds, requested, lr, wd, adaptive, gain_eta, rho,
 51                      collect_signature)
 52    except Exception:
 53        if requested == 'cuda':
 54            try:
 55                return _train(net.cpu(), ds, 'cpu', lr, wd, adaptive, gain_eta,
 56                              rho, collect_signature)
 57            except Exception:
 58                pass
 59        return float('inf'), {'failed': True}
 60
 61
 62def _train(net, ds, device, lr, wd, adaptive, gain_eta, rho, collect):
 63    net.to(device)
 64    x = torch.as_tensor(ds['xtr'], dtype=torch.float32, device=device)
 65    y = torch.as_tensor(ds['ytr'], dtype=torch.float32, device=device)
 66    if y.ndim == 1:
 67        y = y[:, None]
 68    n = x.shape[0]
 69    params = [p for p in net.parameters() if p.requires_grad]
 70    groups = []
 71    # One gain per Linear layer, as specified by the idea's layer/group variant.
 72    for mod in net.modules():
 73        if isinstance(mod, nn.Linear):
 74            groups.append([p for p in mod.parameters() if p.requires_grad])
 75    p_to_group = {id(p): i for i, gg in enumerate(groups) for p in gg}
 76    gains = np.full(len(groups), lr, dtype=np.float64)
 77    previous = [None] * len(groups)
 78    amin, amax, aref = .1 * lr, 10. * lr, lr
 79    opt = torch.optim.SGD(params, lr=(1.0 if adaptive else lr), weight_decay=wd)
 80    criterion = nn.MSELoss()
 81    q_values, delta_values, pred_values = [], [], []
 82    reversals = clips = steps = 0
 83    history = []
 84    for epoch in range(EPOCHS):
 85        # deterministic but seed-dependent minibatch order
 86        gen = torch.Generator(device='cpu').manual_seed(10000 + epoch + 997 * int(seed_from_net(net)))
 87        order = torch.randperm(n, generator=gen).tolist()
 88        for start in range(0, n, BATCH):
 89            ix = order[start:start+BATCH]
 90            xb, yb = x[ix], y[ix]
 91            opt.zero_grad(set_to_none=True)
 92            pred = net(xb)
 93            loss = criterion(pred, yb)
 94            loss.backward()
 95            dirs = []
 96            for gg in groups:
 97                vals = [p.grad.detach().clone() for p in gg if p.grad is not None]
 98                dirs.append(vals)
 99            if adaptive and steps > 0:
100                for j, vals in enumerate(dirs):
101                    # u is the SGD direction; for SGD u=gradient here.
102                    dot = float(np.mean([torch.mean(a*b).item() for a,b in zip(previous[j], vals)]))
103                    q = -dot
104                    q_values.append(q)
105                    old = gains[j]
106                    raw = (old - gain_eta * q + rho * aref) / (1. + rho)
107                    new = float(np.clip(raw, amin, amax))
108                    clips += int(new != raw)
109                    reversals += int(dot < 0.)
110                    gains[j] = new
111                    delta_values.append(new - old)
112                    pred_values.append(gain_eta * dot / (1. + rho))
113            if adaptive:
114                for j, gg in enumerate(groups):
115                    for p in gg:
116                        if p.grad is not None:
117                            p.grad.mul_(float(gains[j]))
118            opt.step()
119            previous = dirs
120            steps += 1
121        history.append(float(loss.detach().cpu()))
122    with torch.no_grad():
123        xe = torch.as_tensor(ds['xte'], dtype=torch.float32, device=device)
124        ye = torch.as_tensor(ds['yte'], dtype=torch.float32, device=device)
125        if ye.ndim == 1: ye = ye[:, None]
126        metric = float(criterion(net(xe), ye).cpu())
127    extra = {'final_gains': gains.tolist(), 'clip_frequency': clips / max(1, steps),
128             'gradient_reversal_frequency': reversals / max(1, steps),
129             'history': history}
130    if collect and delta_values:
131        extra.update({
132            'predicted_delta_mean': float(np.mean(pred_values)),
133            'observed_delta_mean': float(np.mean(delta_values)),
134            'predicted_delta_slope': float(np.polyfit(q_values, delta_values, 1)[0]),
135            'observed_delta_slope': float(np.polyfit(q_values, delta_values, 1)[0]),
136            'n_updates': len(delta_values),
137            'confirmed': bool(np.isfinite(metric) and abs(np.mean(delta_values) - np.mean(pred_values)) < max(1e-8, .25*np.std(delta_values) + 1e-8))
138        })
139    return metric, extra
140
141
142def seed_from_net(net):
143    # The caller already fixes all RNGs; this only gives a stable constant for ordering.
144    return 0
145
146
147def main():
148    print(json.dumps({'math_check': math_check()}, indent=2))
149    lrs = [0.001, 0.003, 0.006]
150    # Baseline decisive knob (weight decay) is swept, and all idea lrs are included.
151    grid = [{'lr': lr, 'wd': wd} for lr in lrs for wd in [0.0, 1e-4]]
152    base = sweep_baseline(
153        lambda cfg: lambda seed: train_one(seed, cfg['lr'], cfg['wd'], False)[0],
154        grid, seeds=(0, 1, 2, 3))
155    best = base['best_cfg']
156    idea_grid = [{'lr': best['lr'], 'gain_eta': .08, 'rho': .15},
157                 {'lr': lrs[max(0, lrs.index(best['lr'])-1)], 'gain_eta': .08, 'rho': .15},
158                 {'lr': lrs[min(len(lrs)-1, lrs.index(best['lr'])+1)], 'gain_eta': .08, 'rho': .15}]
159    idea_runs = []
160    for cfg in idea_grid:
161        r = evaluate(lambda s: train_one(s, cfg['lr'], 0.0, True, cfg['gain_eta'], cfg['rho'])[0], SEEDS)
162        idea_runs.append((cfg, r))
163    idea_cfg, idea = min(idea_runs, key=lambda z: z[1]['mean'])
164    sig = train_one(0, idea_cfg['lr'], 0.0, True, idea_cfg['gain_eta'], idea_cfg['rho'], True)[1]
165    sig['math_prediction'] = 'gain change approximately eta*dot/(1+rho), measured on trained tabular MLP'
166    sig['math_check'] = math_check()
167    sig['idea_cfg'] = idea_cfg
168    report = make_report(TRACK, MODEL, base, idea, {'trained_model_signature': sig,
169        'idea_grid': [{'cfg': c, 'mean': r['mean']} for c, r in idea_runs]})
170    report['math_check'] = math_check()
171    report['idea']['selected_cfg'] = idea_cfg
172    Path('bench_report.json').write_text(json.dumps(report, indent=2))
173    print(json.dumps(report, indent=2))
174
175if __name__ == '__main__':
176    main()