kpshinnik's picture
download
raw
9.31 kB
#!/usr/bin/env python3
"""
Faithful implementation of OpPMD (Optimistic-Predict Mirror Descent) and the
Jin & Sidford (2020) SMD baseline, from the ICML 2026 paper "Efficiently Solving
Discounted MDPs with Predictions on Transition Matrices" (OpenReview 0nrxgFZEEq).
All equations follow ALGO_EXTRACTION.md (Eq 4,7,8,9,10,11,12,13).
MDP minimax reformulation (Eq 4):
min_{v in V} max_{mu in U} f(v,mu) = (1-gamma) q^T v + mu^T((gamma P - Ihat) v + r)
V = {||v||_inf <= 1/(1-gamma)}, U = simplex Delta^N.
"""
import numpy as np
class MDP:
def __init__(self, P, r, q, gamma, pair_state):
self.P = np.asarray(P, float) # (N,S) transition rows
self.r = np.asarray(r, float) # (N,)
self.q = np.asarray(q, float) # (S,)
self.gamma = float(gamma)
self.pair_state = np.asarray(pair_state, int) # (N,) state of each pair
self.N, self.S = self.P.shape
# Ihat (N,S): Ihat[l, pair_state[l]] = 1 (lifts v to pair space)
self.Ihat = np.zeros((self.N, self.S))
self.Ihat[np.arange(self.N), self.pair_state] = 1.0
# actions per state
self.A = [np.where(self.pair_state == i)[0] for i in range(self.S)]
self.Vmax = 1.0 / (1.0 - self.gamma)
self._vstar = None
def f(self, v, mu):
return (1 - self.gamma) * self.q @ v + mu @ ((self.gamma * self.P - self.Ihat) @ v + self.r)
def v_star(self):
"""Exact optimal value via value iteration (cached)."""
if getattr(self, "_vstar", None) is not None:
return self._vstar
v = np.zeros(self.S)
for _ in range(100000):
Q = self.r + self.gamma * self.P @ v
vn = np.array([Q[self.A[i]].max() for i in range(self.S)])
if np.max(np.abs(vn - v)) < 1e-14:
v = vn; break
v = vn
self._vstar = v
return v
def policy_value(self, pi):
"""Exact value of a (possibly stochastic) policy pi (N,) with sum over A_i = 1."""
Ppi = np.zeros((self.S, self.S)); rpi = np.zeros(self.S)
for i in range(self.S):
for l in self.A[i]:
Ppi[i] += pi[l] * self.P[l]
rpi[i] += pi[l] * self.r[l]
return np.linalg.solve(np.eye(self.S) - self.gamma * Ppi, rpi)
def duality_gap(self, vbar, mubar):
"""GAP(vbar,mubar) = max_mu f(vbar,mu) - min_v f(v,mubar), evaluated with TRUE P."""
g = self.gamma
# max over mu in simplex: linear -> pick best coordinate
coeff_mu = (g * self.P - self.Ihat) @ vbar + self.r # (N,)
max_mu = (1 - g) * self.q @ vbar + coeff_mu.max()
# min over v in box [-Vmax,Vmax]^S
grad_v = (1 - g) * self.q + (g * self.P - self.Ihat).T @ mubar # (S,)
min_v = mubar @ self.r - self.Vmax * np.abs(grad_v).sum()
return max_mu - min_v
def _softmax_log(logw):
m = logw.max()
w = np.exp(logw - m)
return w / w.sum()
def oppmd(mdp, P_hat, T, seed=0, track_every=None):
"""Algorithm 1 (OpPMD). Returns dict with gap, policy_value_gap, etc.
track_every: if set, record (t, gap, pv) at those checkpoints."""
rng = np.random.default_rng(seed)
N, S, g = mdp.N, mdp.S, mdp.gamma
P, r, q, Ihat, Vmax = mdp.P, mdp.r, mdp.q, mdp.Ihat, mdp.Vmax
ImgP_hat = Ihat - g * np.asarray(P_hat, float) # (Ihat - gamma Phat) (N,S)
v = np.zeros(S)
mu = np.full(N, 1.0 / N)
gbar_cur = ImgP_hat @ v - r # gbar_1^mu (Eq 13 init)
v_sum = np.zeros(S); mu_sum = np.zeros(N)
Sv = 0.0; Smu = 0.0
# storage for variance-reduced mu-estimator (Eq 12)
z_pair = np.empty(T, int); z_i = np.empty(T, int); z_j = np.empty(T, int)
sqrt2_2 = np.sqrt(2) / 2
checkpoints = {}
track_set = set(track_every) if track_every else set()
for t in range(1, T + 1):
v_sum += v; mu_sum += mu # accumulate iterate averages (v_t, mu_t)
# ---- v-side stochastic gradient (Eq 11) ----
l = rng.choice(N, p=mu) # (i,a) ~ mu_t
i = mdp.pair_state[l]
j = rng.choice(S, p=P[l]) # j ~ p(.|i,a)
ip = rng.choice(S, p=q) # i' ~ q
gv = np.zeros(S)
gv[ip] += (1 - g); gv[j] += g; gv[i] -= 1.0 # (1-g)e_i' + g e_j - e_i
# ---- v learning rate (Eq 7) & update (Eq 8) ----
Sv += gv @ gv
eta_v = sqrt2_2 * (np.sqrt(S) * Vmax) / np.sqrt(Sv + 1e-18)
v_next = np.clip(v - eta_v * gv, -Vmax, Vmax)
# ---- mu-side variance-reduced estimator (Eq 12), evaluated at v_t ----
l2 = rng.integers(N) # (i,a) ~ Uniform(N)
i2 = mdp.pair_state[l2]
j2 = rng.choice(S, p=P[l2]) # j ~ p(.|i,a)
z_pair[t-1] = l2; z_i[t-1] = i2; z_j[t-1] = j2
# residual for all stored samples at current v: N*(v_i - gamma v_j - r_pair)/t
resid = N * (v[z_i[:t]] - g * v[z_j[:t]] - r[z_pair[:t]]) / t
gmu = np.bincount(z_pair[:t], weights=resid, minlength=N) # (N,)
# ---- predicted gradient (Eq 13) & surprise ----
gbar_next = ImgP_hat @ v_next - r
surprise = gmu - gbar_cur # g~_t^mu - gbar_t^mu
Smu += (np.abs(surprise).max()) ** 2
eta_mu = sqrt2_2 * np.sqrt(np.log(N)) / np.sqrt(Smu + 1e-18)
# ---- optimistic entropic mirror-descent mu-update (Eq 10) ----
d = gmu - gbar_cur + gbar_next # optimistic direction
mu = _softmax_log(np.log(mu + 1e-300) - eta_mu * d)
v = v_next
gbar_cur = gbar_next
if t in track_set or t == T:
vbar = v_sum / t; mubar = mu_sum / t
gap = mdp.duality_gap(vbar, mubar)
pi = _extract_policy(mdp, mubar)
pv_gap = float(mdp.q @ (mdp.v_star() - mdp.policy_value(pi)))
checkpoints[t] = {"gap": float(gap), "pv_gap": pv_gap}
vbar = v_sum / T; mubar = mu_sum / T
return {"vbar": vbar, "mubar": mubar, "gap": mdp.duality_gap(vbar, mubar),
"checkpoints": checkpoints}
def jinsidford(mdp, T, eps, seed=0, track_every=None):
"""Baseline: Jin & Sidford (2020) SMD-DMDP — no prediction, fixed eps-dependent rates.
v-side: same estimator, fixed eta_v = eps/8.
mu-side: single-sample estimator, entropic MD, fixed eta_mu = eps/(36((1-g)^-2+1)N)."""
rng = np.random.default_rng(seed)
N, S, g = mdp.N, mdp.S, mdp.gamma
P, r, q, Vmax = mdp.P, mdp.r, mdp.q, mdp.Vmax
eta_v = eps / 8.0
eta_mu = eps / (36.0 * ((1 - g) ** -2 + 1) * N)
v = np.zeros(S); mu = np.full(N, 1.0 / N)
v_sum = np.zeros(S); mu_sum = np.zeros(N)
checkpoints = {}
track_set = set(track_every) if track_every else set()
for t in range(1, T + 1):
v_sum += v; mu_sum += mu
# v-side (Eq 11 estimator, fixed rate)
l = rng.choice(N, p=mu); i = mdp.pair_state[l]
j = rng.choice(S, p=P[l]); ip = rng.choice(S, p=q)
gv = np.zeros(S); gv[ip] += (1 - g); gv[j] += g; gv[i] -= 1.0
v_next = np.clip(v - eta_v * gv, -Vmax, Vmax)
# mu-side single-sample estimator (no reuse, no prediction)
l2 = rng.integers(N); i2 = mdp.pair_state[l2]; j2 = rng.choice(S, p=P[l2])
gmu = np.zeros(N); gmu[l2] = N * (v[i2] - g * v[j2] - r[l2])
mu = _softmax_log(np.log(mu + 1e-300) - eta_mu * gmu)
v = v_next
if t in track_set or t == T:
vbar = v_sum / t; mubar = mu_sum / t
gap = mdp.duality_gap(vbar, mubar)
pi = _extract_policy(mdp, mubar)
pv_gap = float(mdp.q @ (mdp.v_star() - mdp.policy_value(pi)))
checkpoints[t] = {"gap": float(gap), "pv_gap": pv_gap}
vbar = v_sum / T; mubar = mu_sum / T
return {"vbar": vbar, "mubar": mubar, "gap": mdp.duality_gap(vbar, mubar),
"checkpoints": checkpoints}
def _extract_policy(mdp, mubar):
pi = np.zeros(mdp.N)
for i in range(mdp.S):
idx = mdp.A[i]; m = mubar[idx].sum()
pi[idx] = mubar[idx] / m if m > 1e-12 else 1.0 / len(idx)
return pi
def dist(P, P_hat):
"""Dist(P,Phat) = max_{i,a} sum_j |phat - p| (Claim 5)."""
return float(np.abs(np.asarray(P) - np.asarray(P_hat)).sum(axis=1).max())
# ---- the paper's MDP instance (Appendix D) ----
def paper_mdp(gamma=0.5):
P = np.array([[1, 0, 0], [0.4, 0, 0.6], [0, 1, 0], [0, 0.4, 0.6],
[0.4, 0.4, 0.2], [0.2, 0.2, 0.6]], float)
r = np.array([0.001, 0.5, 0.001, 0.5, 1, 1])
q = np.array([0.4, 0.4, 0.2])
pair_state = [0, 0, 1, 1, 2, 2]
return MDP(P, r, q, gamma, pair_state)
P_HAT_NAC = np.array([[0, 1, 0], [0, 1, 0], [1, 0, 0], [1, 0, 0], [1, 0, 0], [0, 1, 0]], float)
if __name__ == "__main__":
mdp = paper_mdp()
print("v* =", mdp.v_star(), " q.v* =", mdp.q @ mdp.v_star())
print("Dist(P, Phat_NAC) =", dist(mdp.P, P_HAT_NAC))
print("Dist(P, P) =", dist(mdp.P, mdp.P))
for name, Phat in [("AC (Phat=P)", mdp.P), ("NAC (bad Phat)", P_HAT_NAC)]:
res = oppmd(mdp, Phat, T=4000, seed=0)
print(f"OpPMD-{name}: gap={res['gap']:.4f}")
resb = jinsidford(mdp, T=4000, eps=0.05, seed=0)
print(f"JinSidford baseline: gap={resb['gap']:.4f}")

Xet Storage Details

Size:
9.31 kB
·
Xet hash:
236bef32b603d89004996a77eacb446fa7bdac8ac4ea694d0af473bc33bcb1e3

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.