"""CPU eval watcher for the v3 force-head run. Loads the dataset/split ONCE, then every N epochs snapshots the live checkpoint and reports, on val / OOD-ModelNet / OOD-Re: mean field R^2 plus the Cd and Cl R^2 from the force head. Appends a CSV and redraws a trend PNG. Pure CPU, so it does not slow the GPU training. python ezflow_v3/gnn/eval_force_watch.py --run rans_v5_s0 --epochs 150 --every 5 """ import os os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"; os.environ["OMP_NUM_THREADS"] = "4" import sys, csv, json, time, shutil, tempfile, argparse import numpy as np, torch torch.set_num_threads(4) import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt sys.path.insert(0, r"C:\dev\EZFlow") from torch_geometric.loader import DataLoader from ezflow_v3.gnn.etl import CaseDatasetV2 from ezflow_v3.gnn import features as F from ezflow_v3.gnn.train_v5 import split, r2 as r2c from ezflow_v3.gnn.model_v5 import MeshGraphNetV5 CACHE = r"C:\dev\ezflow_eval\cache_v3" RUNS = r"C:\dev\EZFlow\ezflow_v3\gnn\_runs" def build(ar): return MeshGraphNetV5(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM, hidden=int(ar["hidden"]), K=int(ar["K"]), out_dim=F.TARGET_DIM, use_film=not ar.get("no_film", False), use_global=not ar.get("no_global", False), agg=ar.get("agg", "meansum"), force_head=ar.get("force_head", False)) def eval_all(model, graphs, n): """ONE forward pass per graph -> (mean field R2, Cd R2, Cl R2). The model returns both node fields and the Cd/Cl head in a single call, so we never re-run the forward pass (halves the previous field-then-force eval cost).""" if not graphs: return float("nan"), float("nan"), float("nan") has_f = n.get("fm") is not None FP, FT, GP, GT = [], [], [], [] with torch.no_grad(): for b in DataLoader(graphs, batch_size=1): b.global_feat = (b.global_feat - n["gm"]) / n["gs"] out = model(b) node = out[0] if isinstance(out, tuple) else out FP.append((node * n["ys"] + n["ym"]).numpy()); FT.append(b.y.numpy()) if has_f and isinstance(out, tuple): GP.append((out[1] * n["fs"] + n["fm"]).numpy()) GT.append(np.array([[float(np.asarray(b.cd).reshape(-1)[0]), float(np.asarray(b.cl).reshape(-1)[0])]], dtype=np.float32)) field = float(np.mean(r2c(np.concatenate(FP), np.concatenate(FT)))) if GP: rr = r2c(np.concatenate(GP), np.concatenate(GT)) return field, float(rr[0]), float(rr[1]) return field, float("nan"), float("nan") def train_loss_at(run, ep): p = os.path.join(run, "progress.jsonl"); best = None if os.path.exists(p): for line in open(p): line = line.strip() if not line: continue try: r = json.loads(line) except Exception: continue if r.get("epoch") == ep: best = r.get("loss") return best def main(): ap = argparse.ArgumentParser() ap.add_argument("--run", default="rans_v5_s0") ap.add_argument("--epochs", type=int, default=150) ap.add_argument("--every", type=int, default=5) ap.add_argument("--label", default=None) a = ap.parse_args() run = os.path.join(RUNS, a.run) label = a.label or a.run csv_path = rf"C:\dev\ezflow_eval\{a.run}_force_trend.csv" png_path = rf"C:\dev\ezflow_eval\{a.run}_force_trend.png" hb_path = os.path.join(run, "heartbeat.json") cols = ["epoch", "train_loss", "val", "ood_modelnet", "ood_Re", "cd_val", "cl_val", "cd_mn", "cl_mn", "time"] print(f"watching '{a.run}', epochs {a.epochs}, eval every {a.every}", flush=True) print("loading dataset/split once ...", flush=True) ds = CaseDatasetV2(CACHE); _, val, ood = split(ds) ood_mn, ood_re = ood["ood_modelnet"], ood["ood_Re_band"] print(f"split: val {len(val)} | ood_modelnet {len(ood_mn)} | ood_Re {len(ood_re)}", flush=True) def hb(): try: return json.load(open(hb_path)) except Exception: return {} def read_csv(): if not os.path.exists(csv_path): return [] return list(csv.DictReader(open(csv_path, encoding="utf-8-sig"))) def draw(): rows = read_csv() if not rows: return eps = [int(r["epoch"]) for r in rows] fig, ax = plt.subplots(2, 3, figsize=(15, 8)) panels = [("val", "field R2 val"), ("ood_modelnet", "field R2 OOD-ModelNet"), ("ood_Re", "field R2 OOD-Re")] for (k, t), x in zip(panels, ax[0]): x.plot(eps, [float(r[k]) for r in rows], "-o", color="#1f77b4", lw=2) x.set_title(t); x.set_xlabel("epoch"); x.set_ylabel("R2"); x.grid(alpha=0.3) ax[1][0].plot(eps, [float(r["cd_val"]) for r in rows], "-o", color="#d62728", label="val") ax[1][0].plot(eps, [float(r["cd_mn"]) for r in rows], "-o", color="#ff9896", label="OOD-MN") ax[1][0].set_title("Cd R2"); ax[1][0].legend(fontsize=8) ax[1][1].plot(eps, [float(r["cl_val"]) for r in rows], "-o", color="#2ca02c", label="val") ax[1][1].plot(eps, [float(r["cl_mn"]) for r in rows], "-o", color="#98df8a", label="OOD-MN") ax[1][1].set_title("Cl R2"); ax[1][1].legend(fontsize=8) tl = [r["train_loss"] for r in rows if r["train_loss"] not in (None, "", "None")] if tl: ax[1][2].plot(eps[-len(tl):], [float(x) for x in tl], "-o", color="#9467bd") ax[1][2].set_title("train loss"); ax[1][2].set_yscale("log") for row in ax: for x in row: x.set_xlabel("epoch"); x.grid(alpha=0.3) fig.suptitle(f"{label} | v3 force-head run | {time.strftime('%Y-%m-%d %H:%M:%S')}", fontsize=12) fig.tight_layout(rect=[0, 0, 1, 0.97]) tmp = png_path + ".tmp.png"; fig.savefig(tmp, dpi=110); plt.close(fig); os.replace(tmp, png_path) done = set(int(r["epoch"]) for r in read_csv()) last = max(done) if done else None first = True while True: h = hb(); cur = int(h.get("epoch", -1)) trigger = first or last is None or cur >= (last + a.every) or cur >= a.epochs if trigger and cur > 0 and os.path.exists(os.path.join(run, "ckpt.pt")): first = False tmp = os.path.join(tempfile.gettempdir(), f"_fsnap_{a.run}.pt") shutil.copyfile(os.path.join(run, "ckpt.pt"), tmp) ck = torch.load(tmp, map_location="cpu", weights_only=False) try: os.remove(tmp) except OSError: pass ep = int(ck.get("epoch", -1)) + 1 if ep not in done: ar = ck["args"] nz = np.load(os.path.join(run, "norms.npz")) n = {"ym": torch.tensor(nz["y_mean"]), "ys": torch.tensor(nz["y_std"]), "gm": torch.tensor(nz["g_mean"]), "gs": torch.tensor(nz["g_std"]), "fm": torch.tensor(nz["f_mean"]) if "f_mean" in nz else None, "fs": torch.tensor(nz["f_std"]) if "f_std" in nz else None} model = build(ar); model.load_state_dict(ck["model"]); model.eval() t0 = time.time() vr, cdv, clv = eval_all(model, val, n) # one pass: field + Cd/Cl mr, cdm, clm = eval_all(model, ood_mn, n) # one pass: field + Cd/Cl rr, _, _ = eval_all(model, ood_re, n) # field only (force not tracked on Re axis) row = dict(epoch=ep, train_loss=train_loss_at(run, ep), val=round(vr, 4), ood_modelnet=round(mr, 4), ood_Re=round(rr, 4), cd_val=round(cdv, 4), cl_val=round(clv, 4), cd_mn=round(cdm, 4), cl_mn=round(clm, 4), time=time.strftime("%H:%M:%S")) wrote = False for _ in range(60): # tolerate a transient lock (e.g. CSV open in a viewer) try: new = not os.path.exists(csv_path) with open(csv_path, "a", newline="") as f: w = csv.DictWriter(f, fieldnames=cols) if new: w.writeheader() w.writerow(row) wrote = True; break except PermissionError: time.sleep(2) done.add(ep); last = ep if not wrote: print(f"WARN: ep{ep} CSV locked >2min, row skipped (continuing)", flush=True) try: draw() except Exception as ex: print(f"draw skipped (locked PNG?): {ex}", flush=True) print(f"ep{ep:>4} val {vr:.3f} mn {mr:.3f} Re {rr:.3f} | " f"Cd[val {cdv:.3f} mn {cdm:.3f}] Cl[val {clv:.3f} mn {clm:.3f}] ({time.time()-t0:.0f}s)", flush=True) if last is not None and last >= a.epochs - 1: print("reached final epoch -> stop watcher", flush=True); break try: age = (time.time() - os.path.getmtime(hb_path)) / 60.0 except OSError: age = 0.0 if age > 40 and last is not None: print("heartbeat stale >40min -> stop watcher", flush=True); break time.sleep(30) print("FORCE_TREND_DONE", flush=True) if __name__ == "__main__": main()