Wasserstein Speed-Limit Controller / wasserstein_controller.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 1729
  6
  7
  8def sliced_w2_sq(a, b, n_proj=128, rng=None):
  9    rng = np.random.default_rng() if rng is None else rng
 10    d = a.shape[1]
 11    q = rng.normal(size=(n_proj, d))
 12    q /= np.linalg.norm(q, axis=1, keepdims=True)
 13    pa = np.sort(a @ q.T, axis=0)
 14    pb = np.sort(b @ q.T, axis=0)
 15    return float(d * np.mean((pa - pb) ** 2))
 16
 17
 18def exact_translation_sweep():
 19    """Verify W2^2 <= D dt Sigma for a distribution translated at velocity u.
 20
 21    For p_t(x)=p_0(x-u t), v=u, so W2^2=||u||^2 dt^2 and
 22    D dt Sigma = D dt * (||u||^2/D)dt = ||u||^2 dt^2 exactly.
 23    This avoids confusing paired particle displacement with distribution W2.
 24    """
 25    rng = np.random.default_rng(SEED)
 26    rows = []
 27    for d in (1, 2, 8):
 28        for D in (0.05, 0.2, 0.8):
 29            for dt in (0.02, 0.08):
 30                for speed in (0.4, 1.2):
 31                    n = 1024
 32                    x0 = rng.normal(size=(n, d))
 33                    u = np.full(d, speed / np.sqrt(d))
 34                    x1 = x0 + dt * u
 35                    w2 = sliced_w2_sq(x0, x1, 128, rng)
 36                    sigma = float(np.dot(u, u) / D)
 37                    rhs = D * dt * sigma * dt
 38                    predicted = float(np.dot(u, u) * dt * dt)
 39                    rows.append({"d": d, "D": D, "dt": dt, "speed": speed,
 40                                 "w2": w2, "rhs": rhs, "predicted": predicted,
 41                                 "ratio": w2 / rhs,
 42                                 "w2_over_speed2_dt2": w2 / (speed * speed * dt * dt)})
 43    ratios = np.array([r["ratio"] for r in rows])
 44    scaling = np.array([r["w2_over_speed2_dt2"] for r in rows])
 45    return rows, {"median_ratio": float(np.median(ratios)),
 46                  "max_abs_ratio_error": float(np.max(np.abs(ratios - 1))),
 47                  "median_speed_dt_scaling": float(np.median(scaling)),
 48                  "speed_dt_scaling_range": [float(np.min(scaling)), float(np.max(scaling))]}
 49
 50
 51def controller_update(eta, ratio, eta_min, eta_max, alpha=0.35, target=0.72):
 52    # Safety-oriented interpretation of the prose: r>target decreases eta;
 53    # r<target permits a cautious increase, with multiplicative clipping.
 54    factor = np.clip((target / max(ratio, 1e-8)) ** alpha, 0.5, 1.25)
 55    return float(np.clip(eta * factor, eta_min, eta_max))
 56
 57
 58def quadratic_run(eta0, controlled, seed, steps=240, K=32, d=4):
 59    rng = np.random.default_rng(seed)
 60    H = np.diag(np.linspace(0.5, 4.0, d))
 61    x = rng.normal(0, 2, size=(K, d))
 62    eta = eta0
 63    T = 0.002
 64    D = T
 65    eta_min, eta_max = eta0 * 0.01, eta0 * 2.0
 66    losses, etas, rs, disp = [], [], [], []
 67    old = x.copy()
 68    for t in range(steps):
 69        grad = x @ H.T
 70        x = x - eta * grad + np.sqrt(2 * eta * T) * rng.normal(size=x.shape)
 71        losses.append(float(np.mean(0.5 * np.sum((x @ H) * x, axis=1))))
 72        etas.append(eta)
 73        if (t + 1) % 8 == 0:
 74            interval = 8 * eta
 75            delta = x - old
 76            sigma = float(np.mean((delta / interval) ** 2) / D)
 77            rhs = D * interval * sigma * interval
 78            w2 = sliced_w2_sq(old, x, 64, rng)
 79            ratio = w2 / max(rhs, 1e-12)
 80            rs.append(ratio)
 81            disp.append(w2)
 82            if controlled:
 83                eta = controller_update(eta, ratio, eta_min, eta_max)
 84            old = x.copy()
 85    return {"final_loss": losses[-1], "min_loss": min(losses),
 86            "max_loss": max(losses), "final_eta": eta,
 87            "mean_ratio": float(np.mean(rs)), "max_ratio": float(np.max(rs)),
 88            "losses": losses, "etas": etas, "ratios": rs, "displacements": disp}
 89
 90
 91def main():
 92    rows, toy_summary = exact_translation_sweep()
 93    sweep = []
 94    for eta in (0.10, 0.30, 0.45, 0.55, 0.70):
 95        sweep.append({"eta0": eta,
 96                      "baseline": quadratic_run(eta, False, SEED),
 97                      "controller": quadratic_run(eta, True, SEED)})
 98    out = {
 99        "seed": SEED,
100        "toy": {"rows": rows, "summary": toy_summary},
101        "quadratic": sweep,
102        "predictions": {
103            "P1_bound": "For translation, W2^2/(D dt Sigma)=1 independent of D.",
104            "P2_scaling": "W2^2 scales as speed^2*dt^2; the dimension-corrected sliced estimator is invariant to D and dimension for fixed total speed.",
105            "P3_optimizer_boundary": "For H eigenvalue 4, noiseless GD boundary is eta=2/4=0.5; high-eta baseline should diverge and controller should reduce eta."
106        }
107    }
108    Path("results.json").write_text(json.dumps(out, indent=2))
109    print(json.dumps({"toy_summary": toy_summary,
110      "quadratic_summary": [{"eta0": r["eta0"],
111        "base_final": r["baseline"]["final_loss"],
112        "ctrl_final": r["controller"]["final_loss"],
113        "base_max_ratio": r["baseline"]["max_ratio"],
114        "ctrl_max_ratio": r["controller"]["max_ratio"],
115        "ctrl_final_eta": r["controller"]["final_eta"]} for r in sweep]}, indent=2))
116
117if __name__ == "__main__":
118    main()