| """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) |
| mr, cdm, clm = eval_all(model, ood_mn, n) |
| rr, _, _ = eval_all(model, ood_re, n) |
| 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): |
| 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() |
|
|