kpshinnik's picture
download
raw
3.76 kB
#!/usr/bin/env python3
"""Reproduce the paper's Figure 1 (Duality Gap vs T) and Figure 2 (Value function vs T)
on the exact Appendix-D MDP, comparing OpPMD-AC, OpPMD-NAC, and SMD-DMDP-JINSID.
Averages over seeds; saves CSV + plots."""
import numpy as np, json, csv, os
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from oppmd import paper_mdp, oppmd, jinsidford, dist, P_HAT_NAC
os.makedirs("outputs", exist_ok=True)
mdp = paper_mdp(gamma=0.5)
EPS = 0.05
SEEDS = list(range(20))
Ts = [100, 200, 400, 800, 1600, 3200, 6400, 12800, 16000]
Tmax = max(Ts)
vstar_q = float(mdp.q @ mdp.v_star())
print(f"v*={mdp.v_star()}, q.v*={vstar_q:.4f}, Dist(P,Phat_NAC)={dist(mdp.P,P_HAT_NAC)}")
algos = {
"OpPMD-AC": lambda seed: oppmd(mdp, mdp.P, Tmax, seed=seed, track_every=Ts),
"OpPMD-NAC": lambda seed: oppmd(mdp, P_HAT_NAC, Tmax, seed=seed, track_every=Ts),
"SMD-DMDP-JINSID": lambda seed: jinsidford(mdp, Tmax, EPS, seed=seed, track_every=Ts),
}
# gap[algo][T] = list over seeds; pv[algo][T] = list of policy value (q.v^pi)
gap = {a: {T: [] for T in Ts} for a in algos}
pv = {a: {T: [] for T in Ts} for a in algos}
for a, fn in algos.items():
for seed in SEEDS:
res = fn(seed)
for T in Ts:
c = res["checkpoints"][T]
gap[a][T].append(c["gap"])
pv[a][T].append(vstar_q - c["pv_gap"]) # achieved value q.v^pi
print(f"{a}: gap@Tmax={np.mean(gap[a][Tmax]):.4f} value@Tmax={np.mean(pv[a][Tmax]):.4f}")
# ---- save CSV ----
with open("outputs/fig1_data.csv", "w", newline="") as f:
w = csv.writer(f); w.writerow(["algo", "T", "gap_mean", "gap_std", "value_mean", "value_std"])
for a in algos:
for T in Ts:
w.writerow([a, T, np.mean(gap[a][T]), np.std(gap[a][T]),
np.mean(pv[a][T]), np.std(pv[a][T])])
summary = {"vstar_q": vstar_q, "eps": EPS, "seeds": len(SEEDS), "gamma": 0.5,
"dist_nac": dist(mdp.P, P_HAT_NAC),
"gap_at_Tmax": {a: float(np.mean(gap[a][Tmax])) for a in algos},
"value_at_Tmax": {a: float(np.mean(pv[a][Tmax])) for a in algos}}
json.dump(summary, open("outputs/fig1_summary.json", "w"), indent=1)
# ---- Figure 1: duality gap ----
colors = {"OpPMD-AC": "#1f77b4", "OpPMD-NAC": "#ff7f0e", "SMD-DMDP-JINSID": "#2ca02c"}
plt.figure(figsize=(6, 4.2))
for a in algos:
m = np.array([np.mean(gap[a][T]) for T in Ts]); s = np.array([np.std(gap[a][T]) for T in Ts])
plt.plot(Ts, m, "o-", color=colors[a], label=a)
plt.fill_between(Ts, np.maximum(m - s, 1e-4), m + s, color=colors[a], alpha=0.15)
plt.xscale("log"); plt.yscale("log"); plt.xlabel("T (iterations, ~2T samples)")
plt.ylabel("Duality Gap GAP(v̄, μ̄)")
plt.title("Fig. 1 reproduction — Duality Gap vs T\n(paper MDP, γ=0.5, mean±std over %d seeds)" % len(SEEDS))
plt.legend(); plt.grid(True, which="both", alpha=0.3); plt.tight_layout()
plt.savefig("outputs/fig1_duality_gap.png", dpi=130); plt.close()
# ---- Figure 2: value function ----
plt.figure(figsize=(6, 4.2))
for a in algos:
m = np.array([np.mean(pv[a][T]) for T in Ts]); s = np.array([np.std(pv[a][T]) for T in Ts])
plt.plot(Ts, m, "o-", color=colors[a], label=a)
plt.fill_between(Ts, m - s, m + s, color=colors[a], alpha=0.15)
plt.axhline(vstar_q, ls="--", color="k", label="q·v* (optimal)")
plt.xscale("log"); plt.xlabel("T (iterations)"); plt.ylabel("Value q·v^π̄ (extracted policy)")
plt.title("Fig. 2 reproduction — Value function vs T\n(paper MDP, γ=0.5, mean±std over %d seeds)" % len(SEEDS))
plt.legend(); plt.grid(True, alpha=0.3); plt.tight_layout()
plt.savefig("outputs/fig2_value.png", dpi=130); plt.close()
print("saved outputs/fig1_duality_gap.png, outputs/fig2_value.png, fig1_data.csv")

Xet Storage Details

Size:
3.76 kB
·
Xet hash:
c81e8639286fe009f6d7daf35e83657a7c37e91eeb748d3df6e399e33dbe84f6

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