Buckets:
| """Claim 4 -- Sec. 3.1 / Fig. 2: continuous entropy-regularised OT with eps = 0.01. | |
| Baseline : kernel SGD on the dual (12) of Genevay et al. (2016), iterates (24)-(25). | |
| Proposed : kernel SGD on the approximate semi-dual (14)-(15), iterates of App. B.1. | |
| Metric : optimality gap W_hat - E_{X~mu_hat} h_eps(X, v) (App. B.3), where W_hat is | |
| obtained by SGD on the semi-discrete problem on a fixed test set. | |
| Independent numpy re-implementation from the equations in the paper (the authors' | |
| PyTorch code was read only for the distribution/hyper-parameter settings). | |
| Setup (Sec. 3.1 + App. B.2): mu = N(1, 1/sqrt(8 pi)), nu = 0.5 N(0, sqrt(0.02)) + | |
| 0.5 N(2, sqrt(1/(2 pi))), c(x,y) = (x-y)^2, Gaussian kernel exp(-|x-x'|^2 / sigma^2). | |
| """ | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import time | |
| from multiprocessing import Pool | |
| import numpy as np | |
| from scipy.special import expit, logsumexp | |
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| OUT = os.path.join(ROOT, "outputs") | |
| TRAJ = os.path.join(OUT, "eot_traj") | |
| os.makedirs(TRAJ, exist_ok=True) | |
| EPS = 0.01 | |
| N_TEST = 10000 | |
| FLOAT64_EXP_MAX = 709.782712893384 # log(np.finfo(np.float64).max) | |
| def sample_mu(rng, n): | |
| return rng.normal(1.0, 1.0 / math.sqrt(8 * math.pi), size=n) | |
| def sample_nu(rng, n): | |
| comp = rng.integers(0, 2, size=n) | |
| means = np.array([0.0, 2.0]) | |
| stds = np.sqrt(np.array([0.02, 1.0 / (2 * math.pi)])) | |
| return rng.normal(means[comp], stds[comp]) | |
| def test_sets(): | |
| rng = np.random.default_rng(42) | |
| return sample_mu(rng, N_TEST), sample_nu(rng, N_TEST) | |
| def eval_semidual(v_test, X_ts, Y_ts, eps=EPS, chunk=500): | |
| """E_{X ~ mu_hat} [ mean(v) - eps log( mean_j exp((v_j - c(X,y_j))/eps) ) - eps ].""" | |
| n = len(Y_ts) | |
| tot = 0.0 | |
| for s in range(0, len(X_ts), chunk): | |
| x = X_ts[s : s + chunk][:, None] | |
| z = (v_test[None, :] - (x - Y_ts[None, :]) ** 2) / eps | |
| tot += float(np.sum(logsumexp(z, axis=1) - math.log(n))) | |
| return float(np.mean(v_test) - eps * tot / len(X_ts) - eps) | |
| # ------------------------------------------------------------------ reference value | |
| def run_reference(seed, n_iter=200000, lr=5.0): | |
| X_ts, Y_ts = test_sets() | |
| rng = np.random.default_rng(1000 + seed) | |
| v = rng.normal(size=N_TEST) | |
| v_avg = v.copy() | |
| for k in range(1, n_iter + 1): | |
| x = sample_mu(rng, 1)[0] | |
| z = (v - (x - Y_ts) ** 2) / EPS | |
| p = np.exp(z - logsumexp(z)) | |
| v += lr * (1.0 / N_TEST - p) | |
| v_avg += (v - v_avg) / k | |
| val = eval_semidual(v_avg, X_ts, Y_ts) | |
| with open(os.path.join(TRAJ, f"ref_seed{seed}.json"), "w") as fh: | |
| json.dump({"seed": seed, "n_iter": n_iter, "lr": lr, "value": val}, fh) | |
| return val | |
| # ------------------------------------------------------------------ kernel SGD | |
| def kernel_sgd( | |
| mode, | |
| sigma_sq, | |
| C, | |
| seed, | |
| rho=None, | |
| n_iter=100001, | |
| eps=EPS, | |
| eval_every=10000, | |
| alpha_init="paper_code", | |
| ): | |
| """mode = 'dual' (baseline, Eq. 12) or 'semidual' (proposed, Eq. 15).""" | |
| X_ts, Y_ts = test_sets() | |
| rng = np.random.default_rng(seed) | |
| X_tr, Y_tr = sample_mu(rng, n_iter), sample_nu(rng, n_iter) | |
| c_tr = (X_tr - Y_tr) ** 2 | |
| coeffs = np.zeros(n_iter) | |
| v_test = np.zeros(N_TEST) | |
| k_prev = 0 | |
| iters, gaps, vals = [], [], [] | |
| overflow_iter, nonfinite_iter = None, None | |
| max_exp_arg = -np.inf | |
| log_rho = math.log(rho) if rho is not None else None | |
| alpha = ( | |
| -c_tr[0] / eps if (mode == "semidual" and alpha_init == "paper_code") else 0.0 | |
| ) | |
| t0 = time.time() | |
| for k in range(n_iter): | |
| if k > 0: | |
| ky = np.exp(-((Y_tr[k] - Y_tr[:k]) ** 2) / sigma_sq) | |
| v_at_yk = float(ky @ coeffs[:k]) | |
| else: | |
| v_at_yk = 0.0 | |
| if mode == "dual": | |
| if k > 0: | |
| kx = np.exp(-((X_tr[k] - X_tr[:k]) ** 2) / sigma_sq) | |
| u_at_xk = float(kx @ coeffs[:k]) | |
| else: | |
| u_at_xk = 0.0 | |
| arg = (u_at_xk + v_at_yk - c_tr[k]) / eps | |
| max_exp_arg = max(max_exp_arg, arg) | |
| if arg > FLOAT64_EXP_MAX and overflow_iter is None: | |
| overflow_iter = k | |
| with np.errstate(over="ignore"): | |
| g = 1.0 - math.exp(arg) if arg < 745 else -np.inf | |
| else: | |
| z = (v_at_yk - c_tr[k]) / eps - alpha | |
| max_exp_arg = max(max_exp_arg, z) | |
| g = 1.0 - expit(z + log_rho) / rho | |
| alpha += C * (-eps * g) / math.sqrt(k + 1) | |
| coeffs[k] = C * g / math.sqrt(k + 1) | |
| if not np.isfinite(coeffs[k]) and nonfinite_iter is None: | |
| nonfinite_iter = k | |
| coeffs[k] = 0.0 | |
| break | |
| if k > 0 and (k % eval_every == 0 or (k in (1, 10, 100, 1000))): | |
| for s0 in range(k_prev, k, 2000): | |
| s1 = min(s0 + 2000, k) | |
| blk = np.exp(-((Y_ts[:, None] - Y_tr[None, s0:s1]) ** 2) / sigma_sq) | |
| v_test = v_test + blk @ coeffs[s0:s1] | |
| k_prev = k | |
| val = eval_semidual(v_test, X_ts, Y_ts) | |
| iters.append(k) | |
| vals.append(val) | |
| rec = { | |
| "mode": mode, | |
| "sigma_sq": sigma_sq, | |
| "C": C, | |
| "seed": seed, | |
| "rho": rho, | |
| "n_iter": n_iter, | |
| "iters": iters, | |
| "values": vals, | |
| "overflow_iter": overflow_iter, | |
| "nonfinite_iter": nonfinite_iter, | |
| "max_exp_arg": float(max_exp_arg), | |
| "alpha_init": alpha_init, | |
| "final_alpha": float(alpha) if mode == "semidual" else None, | |
| "seconds": time.time() - t0, | |
| } | |
| tag = f"{mode}_sig{sigma_sq}_C{C}_rho{rho}_seed{seed}_{alpha_init}" | |
| with open(os.path.join(TRAJ, f"{tag}.json"), "w") as fh: | |
| json.dump(rec, fh) | |
| return rec | |
| def _ref_worker(seed): | |
| return run_reference(seed) | |
| def _sgd_worker(args): | |
| return kernel_sgd(*args) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--procs", type=int, default=40) | |
| ap.add_argument("--seeds", type=int, default=20) | |
| ap.add_argument("--ref-seeds", type=int, default=10) | |
| ap.add_argument("--n-iter", type=int, default=100001) | |
| args = ap.parse_args() | |
| # --- Remark 3.1 overflow thresholds (direct numerical check) -------------- | |
| rem = {} | |
| for dt, name in [(np.float64, "float64"), (np.float32, "float32")]: | |
| lo, hi = 0.0, 20.0 | |
| for _ in range(200): | |
| mid = 0.5 * (lo + hi) | |
| with np.errstate(over="ignore"): | |
| y = np.exp(np.array([mid / EPS], dtype=dt)) | |
| if np.isfinite(y[0]): | |
| lo = mid | |
| else: | |
| hi = mid | |
| rem[name] = { | |
| "z_overflow_threshold": lo, | |
| "paper_states": 7.1 if dt is np.float64 else 0.89, | |
| } | |
| with open(os.path.join(OUT, "claim4_remark31.json"), "w") as fh: | |
| json.dump(rem, fh, indent=2) | |
| print("Remark 3.1 thresholds:", rem, flush=True) | |
| with Pool(args.procs) as pool: | |
| print("reference value ...", flush=True) | |
| ref_vals = pool.map(_ref_worker, list(range(args.ref_seeds))) | |
| W_hat = float(np.max(ref_vals)) | |
| with open(os.path.join(OUT, "claim4_reference.json"), "w") as fh: | |
| json.dump({"values": ref_vals, "W_hat": W_hat}, fh, indent=2) | |
| print("W_hat =", W_hat, ref_vals, flush=True) | |
| tasks = [] | |
| # baseline, best reported stepsize C = 1e-3 | |
| for sig in [0.1, 1.0]: | |
| for s in range(args.seeds): | |
| tasks.append(("dual", sig, 1e-3, s, None, args.n_iter)) | |
| # baseline with larger stepsizes -> instability / overflow probe | |
| for sig in [0.1, 1.0, 10.0]: | |
| for C in [1e-2, 1e-1, 1.0]: | |
| for s in range(4): | |
| tasks.append(("dual", sig, C, s, None, args.n_iter)) | |
| # proposed | |
| for rho, C in [(0.03, 1.0), (0.1, 1.0), (0.3, 10.0)]: | |
| for s in range(args.seeds): | |
| tasks.append(("semidual", 10.0, C, s, rho, args.n_iter)) | |
| # proposed with a neutral alpha init (sensitivity) | |
| for s in range(4): | |
| tasks.append( | |
| ("semidual", 10.0, 1.0, s, 0.1, args.n_iter, EPS, 10000, "zero") | |
| ) | |
| print(f"{len(tasks)} kernel-SGD runs", flush=True) | |
| recs = pool.map(_sgd_worker, tasks) | |
| with open(os.path.join(OUT, "claim4_runs.json"), "w") as fh: | |
| json.dump({"W_hat": W_hat, "runs": recs}, fh) | |
| print("done", flush=True) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 8.5 kB
- Xet hash:
- 501e942bbb7b9c7cf48596217044347d52115d6fdf244791b8e62b0166463845
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.