"""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)