steady-rans-surrogates / code /ezflow_v3 /gnn /plot_loss_live.py
BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
2.74 kB
"""Live training-loss plot. Reads each run's progress.jsonl and redraws a PNG
every ~20 s (atomic write, so VSCode's image preview never catches a half-file).
Open the PNG in VSCode and it auto-refreshes.
python ezflow_v3/gnn/plot_loss_live.py # loop, updates every 20 s
python ezflow_v3/gnn/plot_loss_live.py --once # single draw (test)
"""
import json, os, sys, time
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
RUNS = r"C:\dev\EZFlow\ezflow_v3\gnn\_runs"
OUT = r"C:\dev\ezflow_eval\live_loss.png"
# name, run-tag, color, linewidth, alpha (the live v3 run drawn boldest)
SERIES = [("v3 GeoReNet s0 (live)", "rans_v5_s0", "#d62728", 2.6, 1.0)]
LIVE_TAG = "rans_v5_s0"
def load(tag):
p = os.path.join(RUNS, tag, "progress.jsonl")
eps, ls = [], []
if not os.path.exists(p):
return eps, ls
with open(p) as f:
for line in f:
line = line.strip()
if not line:
continue
try:
r = json.loads(line)
eps.append(r["epoch"]); ls.append(r["loss"])
except Exception:
continue # tolerate a partially-written last line
return eps, ls
def draw():
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(13, 5))
cur = None
for name, tag, color, lw, alpha in SERIES:
eps, ls = load(tag)
if not eps:
continue
ax1.plot(eps, ls, color=color, lw=lw, alpha=alpha, label=name)
ax2.plot(eps, ls, color=color, lw=lw, alpha=alpha)
if tag == LIVE_TAG:
cur = (eps[-1], ls[-1])
for ax in (ax1, ax2):
ax.scatter([eps[-1]], [ls[-1]], color=color, zorder=5, s=40)
ax1.set_yscale("log")
ax1.set_title("Train loss — full history (log scale)")
ax2.set_title("Train loss — zoom (linear, < 0.6)")
ax2.set_ylim(0, 0.6)
for ax in (ax1, ax2):
ax.set_xlabel("epoch"); ax.set_ylabel("weighted MSE (train)")
ax.grid(alpha=0.3)
ax1.legend(loc="upper right")
ts = time.strftime("%Y-%m-%d %H:%M:%S")
title = f"EZFlow training loss | updated {ts}"
if cur:
title += f" | {LIVE_TAG} ep{cur[0]} loss={cur[1]:.4f}"
fig.suptitle(title, fontsize=12)
fig.tight_layout(rect=[0, 0, 1, 0.96])
tmp = OUT + ".tmp.png"
fig.savefig(tmp, dpi=110)
plt.close(fig)
os.replace(tmp, OUT) # atomic
if __name__ == "__main__":
once = "--once" in sys.argv
while True:
try:
draw()
print(f"drew {OUT} @ {time.strftime('%H:%M:%S')}", flush=True)
except Exception as e:
print("draw error:", e, flush=True)
if once:
break
time.sleep(20)