Wasserstein Speed-Limit Controller / wasserstein_controller.py
Failed on benchmark
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()