Removable-Pole Negative-Shifted Optimizer / stage2_negative_shift_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
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, sweep_baseline, evaluate, make_report
  8
  9SEEDS = tuple(range(8))
 10LRS = [0.001, 0.003, 0.01]
 11EPOCHS, BATCH = 18, 128
 12NU_MULTS = [0.0, 0.05, 0.15]
 13
 14
 15def seed_all(seed):
 16    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 17    if torch.cuda.is_available():
 18        try: torch.cuda.manual_seed_all(seed)
 19        except Exception: pass
 20
 21
 22def get_data(seed):
 23    return get_dataset('tabular', seed, n_train=400, n_test=400)
 24
 25
 26def hessian_scale(model, ds, dev):
 27    """Power iteration on the minibatch Hessian, giving an actual curvature scale."""
 28    model.zero_grad(set_to_none=True)
 29    x, y = ds['xtr'][:128].to(dev), ds['ytr'][:128].to(dev)
 30    loss = nn.MSELoss()(model(x), y)
 31    gs = torch.autograd.grad(loss, tuple(model.parameters()), create_graph=True)
 32    vs = [torch.randn_like(p) for p in model.parameters()]
 33    norm = torch.sqrt(sum((v*v).sum() for v in vs))
 34    vs = [v / norm for v in vs]
 35    val = 1e-4
 36    for _ in range(4):
 37        dot = sum((g*v).sum() for g, v in zip(gs, vs))
 38        hv = torch.autograd.grad(dot, tuple(model.parameters()), retain_graph=True)
 39        norm = torch.sqrt(sum((h*h).sum() for h in hv))
 40        val = float(norm.detach().cpu())
 41        vs = [h / (norm + 1e-12) for h in hv]
 42    model.zero_grad(set_to_none=True)
 43    return max(val, 1e-5)
 44
 45
 46def baseline_one(cfg, seed):
 47    seed_all(seed); ds = get_data(seed)
 48    model = make_model('mlp_tiny', ds['input_shape'], ds['out_dim'])
 49    _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH,
 50                               weight_decay=cfg['weight_decay'], log=lambda *a: None)
 51    return float('inf') if metric is None else metric
 52
 53
 54def shifted_one(cfg, seed, collect=False):
 55    seed_all(seed); ds = get_data(seed); model = make_model('mlp_tiny', ds['input_shape'], ds['out_dim'])
 56    try:
 57        dev = 'cuda' if torch.cuda.is_available() else 'cpu'
 58        model = model.to(dev)
 59        initial = [p.detach().clone() for p in model.parameters()]
 60        scale = hessian_scale(model, ds, dev); nu = cfg['nu_mult'] * scale
 61        opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 62        x, y = ds['xtr'].to(dev), ds['ytr'].to(dev); lossf = nn.MSELoss(); hist = []
 63        for _ in range(EPOCHS):
 64            for ix in torch.randperm(len(x), device=dev).split(BATCH):
 65                opt.zero_grad(set_to_none=True); loss = lossf(model(x[ix]), y[ix]); loss.backward()
 66                # Apply the negative quadratic term to displacement delta=p-p0.
 67                with torch.no_grad():
 68                    for p, p0 in zip(model.parameters(), initial):
 69                        if p.grad is not None: p.grad.add_(-nu * (p - p0))
 70                torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0); opt.step()
 71            hist.append(float(loss.detach().cpu()))
 72        with torch.no_grad():
 73            test = float(lossf(model(ds['xte'].to(dev)), ds['yte'].to(dev)).cpu())
 74            disp = float(torch.sqrt(sum(((p-p0)**2).sum() for p,p0 in zip(model.parameters(), initial))).cpu())
 75        if collect: return test, {'model': model, 'ds': ds, 'scale': scale, 'nu': nu, 'disp': disp, 'history': hist}
 76        return test
 77    except Exception:
 78        try: torch.cuda.empty_cache()
 79        except Exception: pass
 80        return float('inf')
 81
 82
 83def main():
 84    base_grid = [{'lr': lr, 'weight_decay': wd} for lr in LRS for wd in [0.0, 1e-4]]
 85    baseline = sweep_baseline(lambda c: lambda s: baseline_one(c, s), base_grid, seeds=(0,1,2,3))
 86    # Same lr union and baseline's selected central regularization; only nu differs.
 87    wd = baseline['best_cfg']['weight_decay']
 88    idea_grid = [{'lr': lr, 'weight_decay': wd, 'nu_mult': nu} for lr in LRS for nu in NU_MULTS]
 89    tried = []
 90    for cfg in idea_grid:
 91        r = evaluate(lambda s, c=cfg: shifted_one(c, s), seeds=(0,1,2,3))
 92        tried.append({'cfg': cfg, 'mean': r['mean']})
 93    best_cfg = min(idea_grid, key=lambda c: next(q['mean'] for q in tried if q['cfg'] == c))
 94    idea = evaluate(lambda s: shifted_one(best_cfg, s), seeds=SEEDS)
 95    # Signature is measured on trained models: positive-shift run versus zero-shift run.
 96    sig_cfg = next(c for c in idea_grid if c['lr'] == best_cfg['lr'] and c['nu_mult'] == 0.15)
 97    a = shifted_one(sig_cfg, 0, collect=True); z = shifted_one({**sig_cfg, 'nu_mult': 0.0}, 0, collect=True)
 98    predicted = 1.0 + sig_cfg['lr'] * a[1]['nu']
 99    observed = (a[1]['disp'] / max(z[1]['disp'], 1e-12))
100    signature = {'predicted': {'one_step_displacement_factor': predicted, 'nu': a[1]['nu'], 'curvature_scale': a[1]['scale']},
101                 'observed': {'trained_positive_shift_disp': a[1]['disp'], 'trained_zero_shift_disp': z[1]['disp'], 'whole_run_ratio': observed},
102                 'confirmed': bool(np.isfinite(observed) and abs(observed-predicted) / max(abs(predicted),1e-9) < 0.25)}
103    rep = make_report('tabular', 'mlp_tiny', baseline, idea, signature)
104    rep['idea_sweep'] = {'grid': tried, 'best_cfg': best_cfg, 'shared_lr_union': LRS}
105    Path('bench_report.json').write_text(json.dumps(rep, indent=2)); print(json.dumps(rep, indent=2))
106
107if __name__ == '__main__': main()