steady-rans-surrogates / code /ezflow_v3 /gnn /eval_force_watch.py
BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
9.55 kB
"""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()