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