DooABLe / scripts /plot_results.py
pranamanam's picture
Upload 309 files
81ae663 verified
Raw
History Blame Contribute Delete
4.37 kB
"""Create vector plots and a LaTeX table from recorded benchmark measurements."""
import argparse
from pathlib import Path
import json
import matplotlib
matplotlib.use("Agg")
from matplotlib import pyplot as plt, font_manager
import pandas as pd
ROOT = Path(__file__).resolve().parents[1]
COLORS = ["#4E95C0", "#245778", "#A7D3EE", "#718594", "#7EB4D5", "#364C5B", "#D7EAF6"]
LABELS = {
"dooable": "DooABLe",
"tb_uniform": "TB, uniform backward",
"tb_exact": "TB, exact backward",
"uniform": "Uniform executable",
"reference_tilt": "Reference tilt",
"exact": "Exact joint law",
"zero_cost": "Zero route cost",
"duplicate_endpoints": "Duplicate endpoints",
"unnormalized": "Unnormalized backward",
}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--run", required=True)
parser.add_argument("--output")
args = parser.parse_args()
run = Path(args.run)
out = Path(args.output) if args.output else run / "plots"
out.mkdir(parents=True, exist_ok=True)
font = ROOT / "assets/fonts/Ubuntu-Regular.ttf"
if not font.exists():
font = ROOT / "paper/figures/fonts/Ubuntu-Regular.ttf"
font_manager.fontManager.addfont(str(font))
plt.rcParams.update(
{
"font.family": "Ubuntu",
"font.size": 9,
"mathtext.fontset": "cm",
"pdf.fonttype": 42,
"axes.spines.top": False,
"axes.spines.right": False,
}
)
df = pd.read_csv(run / "metrics.csv")
methods = list(dict.fromkeys(df.method))
fig, axes = plt.subplots(1, 3, figsize=(10, 3.2))
specifications = [
("endpoint_tv", "Endpoint TV"),
("conditional_free_energy_gap", "Conditional gap (cost units)"),
("mean_cost", "Mean execution cost"),
]
for ax, (metric, label) in zip(axes, specifications):
summary = df.groupby("method")[metric].agg(["mean", "sem"]).reindex(methods)
ax.bar(
range(len(methods)),
summary["mean"],
yerr=summary["sem"].fillna(0),
color=COLORS[: len(methods)],
capsize=2,
)
ax.set_ylabel(label)
ax.set_xticks(
range(len(methods)), [LABELS[x] for x in methods], rotation=40, ha="right"
)
fig.tight_layout()
fig.savefig(out / "benchmark.pdf")
fig.savefig(out / "benchmark.png", dpi=220)
plt.close(fig)
histories = []
for path in sorted(run.glob("*_seed*/training.json")):
frame = pd.read_json(path)
frame["method"] = path.parent.name.rsplit("_seed", 1)[0]
frame["seed"] = int(path.parent.name.rsplit("_seed", 1)[1])
histories.append(frame)
if histories:
history = pd.concat(histories)
fig, ax = plt.subplots(figsize=(5.3, 3.1))
for color, (method, group) in zip(COLORS, history.groupby("method")):
stats = group.groupby("step").endpoint_tv.agg(["mean", "sem"])
x = stats.index.to_numpy()
mean = stats["mean"].to_numpy()
sem = stats["sem"].fillna(0).to_numpy()
ax.plot(x, mean, label=LABELS[method], color=color)
ax.fill_between(x, mean - sem, mean + sem, color=color, alpha=0.15)
ax.set_xlabel("Training updates")
ax.set_ylabel("Endpoint TV")
ax.legend(frameon=False)
fig.tight_layout()
fig.savefig(out / "training.pdf")
fig.savefig(out / "training.png", dpi=220)
plt.close(fig)
summary = df.groupby("method")[[x[0] for x in specifications]].agg(["mean", "sem"])
lines = [
"\\begin{tabular}{lccc}",
"\\toprule",
"Method & TV & Conditional gap & Mean cost \\\\",
"\\midrule",
]
for method in methods:
cells = []
for metric, _ in specifications:
value = summary.loc[method, (metric, "mean")]
sem = summary.loc[method, (metric, "sem")]
cells.append(
f"${value:.4g}$" if pd.isna(sem) else f"${value:.4g} \\pm {sem:.2g}$"
)
lines.append(LABELS[method] + " & " + " & ".join(cells) + " \\\\")
lines.extend(["\\bottomrule", "\\end{tabular}"])
(out / "measurements.tex").write_text("\n".join(lines) + "\n")
print(json.dumps({"rows": len(df), "plots": str(out)}))
if __name__ == "__main__":
main()