"""Reproduce the synthetic-data figure from the recorded production skill session.

The coding agent creates this file locally; SkillGild validates the palette.
Run beside synthetic_accuracy.csv with the versions in requirements.txt.
"""

import csv
import json
from collections import defaultdict
from pathlib import Path

import matplotlib

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

ROOT = Path(__file__).resolve().parent
COLORS = {"Baseline": "#6D6A8B", "Example method": "#2467BD"}
matplotlib.rcParams.update(
    {
        "font.family": "DejaVu Serif",
        "font.size": 7.5,
        "text.color": "#2B2A3D",
        "axes.labelcolor": "#2B2A3D",
        "axes.edgecolor": "#5D5B78",
        "xtick.color": "#5D5B78",
        "ytick.color": "#5D5B78",
        "svg.fonttype": "none",
        "svg.hashsalt": "skillgild-production-plotting-example",
        "pdf.fonttype": 42,
    }
)

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

fig, ax = plt.subplots(figsize=(3.25, 2.5))
fig.subplots_adjust(left=0.17, right=0.70, bottom=0.18, top=0.84)
last_means = {}
for method, by_step in runs.items():
    steps = sorted(by_step)
    if any(len(by_step[step]) != 3 for step in steps):
        raise ValueError("Each method and checkpoint must contain exactly three seeds")
    xs = np.array(steps) / 1000
    mean = np.array([np.mean(by_step[step]) for step in steps])
    std = np.array([np.std(by_step[step], ddof=1) for step in steps])
    color = COLORS[method]
    ax.fill_between(xs, mean - std, mean + std, color=color, alpha=0.12, lw=0)
    ax.plot(
        xs,
        mean,
        color=color,
        linestyle="--" if method == "Baseline" else "-",
        marker="s" if method == "Baseline" else "o",
        markersize=2.8,
        linewidth=1.1,
        markerfacecolor="white",
        markeredgewidth=0.9,
    )
    ax.annotate(
        f"{'Example' if method == 'Example method' else method}\n{mean[-1]:.1f}%",
        (xs[-1], mean[-1]),
        xytext=(5, 6 if method == "Example method" else -7),
        textcoords="offset points",
        color=color,
        fontsize=7,
        va="center",
        annotation_clip=False,
    )
    last_means[method] = float(mean[-1])

ax.set_xlim(0, 16)
ax.set_ylim(0, 90)
ax.set_xticks([0, 4, 8, 12, 16])
ax.set_xlabel("Training steps (thousands)", fontsize=7.5)
ax.set_ylabel("Validation accuracy (%)", fontsize=7.5)
ax.grid(axis="y", color="#E4E1F2", linewidth=0.4)
ax.set_axisbelow(True)
ax.spines[["top", "right"]].set_visible(False)
fig.suptitle("Synthetic data: mean ± s.d., 3 seeds", x=0.04, ha="left", fontsize=7)

fig.savefig(ROOT / "accuracy.svg", metadata={"Date": None})
fig.savefig(ROOT / "accuracy.pdf", metadata={"CreationDate": None, "ModDate": None})
fig.savefig(ROOT / "accuracy.png", dpi=300)
plt.close(fig)

print(json.dumps({"final_mean_accuracy_pct": last_means, "size_inches": [3.25, 2.5]}))
