Buckets:
| #!/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.