"""Plot mean +/- standard deviation across seeds from synthetic_accuracy.csv.

The figure is drawn at 6.5 x 3.6 in and saved cropped to its content (about
5.6 x 3.7 in). Check your venue's column width and rescale before submitting.
All data is synthetic and generated by make_dataset.py.
"""
import csv
from collections import defaultdict

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np

matplotlib.rcParams["svg.fonttype"] = "none"  # keep text as text
matplotlib.rcParams["svg.hashsalt"] = "skillgild-example"  # reproducible ids
plt.rcParams.update(
    {
        "font.family": "sans-serif",
        "font.size": 10,
        "axes.edgecolor": "#8f8bb8",
        "axes.labelcolor": "#2b2a3d",
        "xtick.color": "#5d5b78",
        "ytick.color": "#5d5b78",
        "text.color": "#2b2a3d",
    }
)

runs = defaultdict(lambda: defaultdict(list))
with open("synthetic_accuracy.csv") as f:
    for row in csv.DictReader(f):
        runs[row["method"]][int(row["step"])].append(float(row["val_accuracy_pct"]))

fig, ax = plt.subplots(figsize=(6.5, 3.6), layout="constrained")
colors = {"Baseline": "#7a7799", "Example method": "#2a7de1"}
for method, by_step in runs.items():
    steps = sorted(by_step)
    xs = np.array(steps) / 1000
    mean = np.array([np.mean(by_step[s]) for s in steps])
    std = np.array([np.std(by_step[s], ddof=1) for s in steps])
    c = colors[method]
    ax.fill_between(xs, mean - std, mean + std, color=c, alpha=0.16, lw=0)
    ax.plot(xs, mean, color=c, lw=2.2, marker="o", ms=4.5, mfc="white", mew=1.6)
    ax.annotate(
        f"{method}\n{mean[-1]:.1f}%",
        (xs[-1], mean[-1]),
        xytext=(8, 7 if method == "Example method" else -9),
        textcoords="offset points",
        va="center",
        color=c,
        fontweight="bold",
        fontsize=9.5,
        linespacing=1.3,
        annotation_clip=False,
    )

ax.set_xlim(0, 16)
ax.set_ylim(0, 90)
ax.set_xticks([0, 2, 4, 8, 12, 16])
ax.set_xlabel("Training steps (thousands)")
ax.set_ylabel("Validation accuracy (%)")
ax.grid(axis="y", color="#e4e1f2", lw=0.8)
ax.set_axisbelow(True)
ax.spines[["top", "right"]].set_visible(False)
ax.margins(x=0)
fig.suptitle(
    "Synthetic example data: mean ± s.d. over 3 seeds",
    x=0.01,
    ha="left",
    fontsize=9,
    color="#5d5b78",
)
fig.get_layout_engine().set(rect=(0, 0, 0.84, 1))  # room for end labels

fig.savefig("accuracy_curve.svg", bbox_inches="tight", pad_inches=0.1)
fig.savefig("accuracy_curve.png", dpi=200, bbox_inches="tight", pad_inches=0.1)
