Better graphing script

This commit is contained in:
City
2023-07-31 21:28:55 +02:00
parent a347f39606
commit 994bb1a04b
+16 -2
View File
@@ -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')