Semiglobal-PL Phase Scheduler / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 2015
  6
  7def lqr(k):
  8    return (1.0 + k * k) / (2.0 * k)
  9
 10def lqr_grad(k):
 11    return 0.5 * (1.0 - 1.0 / (k * k))
 12
 13def lqr_q(k):
 14    return abs(lqr_grad(k)) / math.sqrt(max(lqr(k) - 1.0, 1e-30))
 15
 16def lqr_check():
 17    # Directly verify the claimed q formula and gradient-flow behavior.
 18    ks = np.geomspace(1.001, 1000.0, 400)
 19    formula_err = []
 20    for k in ks:
 21        closed = (k + 1.0) / (math.sqrt(2.0) * k ** 1.5)
 22        formula_err.append(abs(lqr_q(k) - closed) / closed)
 23    trajectories = []
 24    for k0 in (1.05, 1.2, 2.0, 10.0):
 25        eta, k = 1e-4, k0
 26        gaps, qs, times = [], [], []
 27        for step in range(30000):
 28            if step % 100 == 0:
 29                gaps.append(lqr(k) - 1.0); qs.append(lqr_q(k)); times.append(step * eta)
 30            k = max(1.0000001, k - eta * lqr_grad(k))
 31        # local late-trajectory log-gap slope versus the contemporaneous q^2
 32        y = np.log(np.maximum(gaps, 1e-300)); x = np.asarray(times)
 33        slope = float(np.polyfit(x[-100:], y[-100:], 1)[0])
 34        q2 = float(np.mean(np.asarray(qs[-100:]) ** 2))
 35        trajectories.append({"k0": k0, "log_gap_rate": slope,
 36                             "minus_mean_q2": -q2,
 37                             "relative_error": abs(slope + q2) / q2})
 38    # One-step Euler check of the discrete prediction log gap ~= -eta*q^2.
 39    discrete = []
 40    k = 1.01
 41    for eta in (1e-4, 2e-4, 5e-4, 1e-3, 2e-3):
 42        old = lqr(k) - 1.0
 43        knew = k - eta * lqr_grad(k)
 44        observed = math.log((lqr(knew) - 1.0) / old)
 45        predicted = -eta * lqr_q(k) ** 2
 46        discrete.append({"eta": eta, "observed": observed, "predicted": predicted,
 47                         "relative_error": abs(observed - predicted) / abs(predicted)})
 48    far_rate = lqr_grad(500.0) ** 2
 49    return {"max_q_formula_relative_error": float(max(formula_err)),
 50            "flow_rate_checks": trajectories,
 51            "discrete_rate_checks": discrete,
 52            "far_field_rate_at_k500": float(far_rate),
 53            "checks_pass": bool(max(formula_err) < 1e-8 and
 54                                  max(r["relative_error"] for r in trajectories) < .01 and
 55                                  max(r["relative_error"] for r in discrete) < .02 and
 56                                  abs(far_rate - .25) < .01)}
 57
 58def make_data(n=1200):
 59    rng = np.random.RandomState(SEED)
 60    x = rng.randn(n, 2).astype(np.float32)
 61    y = (x[:, 0] * x[:, 1] + .25 * x[:, 0] - .15 * x[:, 1] > 0).astype(np.float32)
 62    return x, y
 63
 64def mlp_run(scheduled, seed, steps=500, force_cpu=False):
 65    import torch
 66    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 67    requested = "cpu" if force_cpu else ("cuda" if torch.cuda.is_available() else "cpu")
 68    try:
 69        dev = torch.device(requested)
 70        x, y = make_data()
 71        xt, yt = torch.tensor(x[:900], device=dev), torch.tensor(y[:900], device=dev)
 72        xv, yv = torch.tensor(x[900:], device=dev), torch.tensor(y[900:], device=dev)
 73        model = torch.nn.Sequential(torch.nn.Linear(2, 32), torch.nn.Tanh(),
 74                                    torch.nn.Linear(32, 1)).to(dev)
 75        opt = torch.optim.SGD(model.parameters(), lr=.08)
 76        lossfn = torch.nn.BCEWithLogitsLoss()
 77        fhat = None; qhist = []; gaphist = []; consecutive = 0; switched = False
 78        train, val, lrs, qs = [], [], [], []
 79        clip_count = 0
 80        for step in range(steps):
 81            opt.zero_grad(set_to_none=True)
 82            loss = lossfn(model(xt).squeeze(-1), yt); loss.backward()
 83            g = math.sqrt(sum(float((p.grad.detach() ** 2).sum().cpu())
 84                              for p in model.parameters() if p.grad is not None))
 85            lv = float(loss.detach().cpu())
 86            if fhat is None: fhat = lv
 87            else: fhat = min(fhat, .98 * fhat + .02 * lv)
 88            gap = max(lv - fhat, 1e-8); q = g / math.sqrt(gap)
 89            qhist.append(q); gaphist.append(gap); qs.append(q)
 90            if scheduled and len(qhist) >= 20:
 91                gaps = np.asarray(gaphist[-100:]); qvals = np.asarray(qhist[-100:])
 92                candidate = gap <= np.percentile(gaps, 35) and q >= max(np.percentile(qvals, 10) / 2, 1e-5)
 93                consecutive = consecutive + 1 if candidate else 0
 94                if consecutive >= 5 and not switched:
 95                    for group in opt.param_groups: group["lr"] *= 1.5
 96                    switched = True; switch_step = step
 97            # Identical clipping in both arms makes the comparison isolate scheduling.
 98            before = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
 99            clip_count += int(float(before) > 1.0)
100            opt.step()
101            with torch.no_grad():
102                train.append(lv); val.append(float(lossfn(model(xv).squeeze(-1), yv).cpu()))
103            lrs.append(opt.param_groups[0]["lr"])
104        tail = np.asarray(val[-100:])
105        return {"final_train": train[-1], "final_val": val[-1], "best_val": min(val),
106                "val_last100_std": float(tail.std()), "switched": switched,
107                "switch_step": int(switch_step) if switched else None,
108                "lr_final": opt.param_groups[0]["lr"], "clip_steps": clip_count,
109                "device": str(dev), "mean_q_last100": float(np.mean(qs[-100:]))}
110    except Exception:
111        if requested == "cuda":
112            try: torch.cuda.empty_cache()
113            except Exception: pass
114            return mlp_run(scheduled, seed, steps, force_cpu=True)
115        raise
116
117def main():
118    out = {"math_check": lqr_check(), "mlp": {}}
119    for name, scheduled in (("fixed_sgd", False), ("semiglobal_scheduler", True)):
120        runs = [mlp_run(scheduled, seed) for seed in (SEED, SEED + 1, SEED + 2)]
121        for r in runs: r.pop("device", None)
122        for key in ("final_val", "best_val", "val_last100_std", "lr_final", "mean_q_last100"):
123            vals = [r[key] for r in runs]
124            out["mlp"].setdefault(name, {})[key + "_mean"] = float(np.mean(vals))
125            out["mlp"][name][key + "_std"] = float(np.std(vals))
126        out["mlp"][name]["runs"] = runs
127    Path("results.json").write_text(json.dumps(out, indent=2))
128    print(json.dumps(out, indent=2))
129
130if __name__ == "__main__": main()