SabaPivot's picture
download
raw
8.5 kB
"""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.