From 994bb1a04b659218509ab925b08101ad2a4f9d56 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 31 Jul 2023 21:28:55 +0200 Subject: [PATCH] Better graphing script --- log_loss.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/log_loss.py b/log_loss.py index a80499d..3111cc0 100644 --- a/log_loss.py +++ b/log_loss.py @@ -1,4 +1,5 @@ import os +import math import matplotlib.pyplot as plt files = [f"models/{x}" for x in os.listdir("models") if x.endswith(".csv")] @@ -10,11 +11,24 @@ for fp in files: if not step: step = [int(x.split(",")[0]) for x in lines] name = fp.split("/")[1].split("_")[0] - data[name] = [float(x.split(",")[1]) for x in lines] + data[name] = ( + [int(x.split(",")[0]) for x in lines], + [math.log(float(x.split(",")[1])) for x in lines], + ) + +# https://stackoverflow.com/a/49357445 +def smooth(scalars, weight): + last = scalars[0] + smoothed = list() + for point in scalars: + smoothed_val = last * weight + (1 - weight) * point + smoothed.append(smoothed_val) + last = smoothed_val + return smoothed fig, ax = plt.subplots() ax.grid() for name, val in data.items(): - ax.plot(step, val, label=name) + ax.plot(val[0], smooth(val[1], 0.7), label=name) plt.legend(loc="upper right") plt.savefig('loss.png', dpi=300, bbox_inches='tight')