BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
5.99 kB
"""CPU-only actual-vs-predicted field plots for OOD cases (uses best.pt).
Renders, for a few held-out (OOD) cases, a z~0 mid-plane slice colored by the
true CFD field, the model prediction, and the absolute error -- for both velocity
magnitude |U| and pressure p. Pure CPU; safe to run alongside GPU training.
Run: python -m ezflow_v3.gnn.plot_ood --n 2
"""
from __future__ import annotations
import argparse, os, sys
import numpy as np
import torch
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import matplotlib.tri as mtri
os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE")
os.environ.setdefault("OMP_NUM_THREADS", "2")
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
from ezflow_v3.gnn.etl import CaseDatasetV2
from ezflow_v3.gnn.model_v5 import MeshGraphNetV5
from ezflow_v3.gnn import features as F
from ezflow_v3.gnn.train_v5 import split
def r2(t, p):
ss = ((t - p) ** 2).sum(); tot = ((t - t.mean()) ** 2).sum() + 1e-30
return float(1 - ss / tot)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--run", default=os.path.join(os.path.dirname(__file__), "_runs", "rans_v5"))
ap.add_argument("--cache", default=r"C:\dev\ezflow_eval\cache_v3")
ap.add_argument("--out", default=r"C:\dev\ezflow_eval\ood_plots")
ap.add_argument("--ckpt", default="best.pt")
ap.add_argument("--n", type=int, default=2)
a = ap.parse_args()
os.makedirs(a.out, exist_ok=True)
ck = torch.load(os.path.join(a.run, a.ckpt), map_location="cpu", weights_only=False)
model = MeshGraphNetV5(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM,
hidden=int(ck["args"]["hidden"]), K=int(ck["args"]["K"]),
out_dim=F.TARGET_DIM)
model.load_state_dict(ck["model"]); model.eval()
ep = ck.get("epoch", "?")
print(f"loaded {a.ckpt} (epoch {ep}, val_loss {ck.get('val_loss','?')})", flush=True)
nz = np.load(os.path.join(a.run, "norms.npz"))
ym = torch.tensor(nz["y_mean"]); ys = torch.tensor(nz["y_std"])
gm = torch.tensor(nz["g_mean"]); gs = torch.tensor(nz["g_std"])
ds = CaseDatasetV2(a.cache)
_, _, ood = split(ds)
for setname, graphs in ood.items():
sel = graphs[:a.n]
if not sel:
continue
fig, axes = plt.subplots(len(sel), 6, figsize=(23, 3.6 * len(sel) + 0.8), squeeze=False)
for r, g in enumerate(sel):
gg = g.clone(); gg.global_feat = (gg.global_feat - gm) / gs
with torch.no_grad():
pred = model(gg).numpy() * ys.numpy() + ym.numpy() # -> physical
gt = g.y.numpy()
pos = g.pos.numpy()
uT = np.linalg.norm(gt[:, :3], axis=1); uP = np.linalg.norm(pred[:, :3], axis=1)
pT, pP = gt[:, 3], pred[:, 3]
r2u, r2p = r2(uT, uP), r2(pT, pP) # R2 on the full graph
Re = float(getattr(g, "Re", 0.0))
# z~0 mid-plane slice (flow is in x-y), then crop to body + wake
z = pos[:, 2]; tol = 0.15
while (np.abs(z) < tol).sum() < 800 and tol < 1.0:
tol += 0.08
body = uT < 0.2
cx = np.median(pos[body, 0]) if body.any() else 0.0
cy = np.median(pos[body, 1]) if body.any() else 0.0
m = (np.abs(z) < tol) & (pos[:, 0] > cx - 2.5) & (pos[:, 0] < cx + 8.0) \
& (np.abs(pos[:, 1] - cy) < 3.0)
X, Y = pos[m, 0], pos[m, 1]
# triangulate the slice, mask triangles bridging the body/far gaps
tri = mtri.Triangulation(X, Y)
t = tri.triangles
e = np.concatenate([np.hypot(X[t[:, i]] - X[t[:, j]], Y[t[:, i]] - Y[t[:, j]])
for i, j in ((0, 1), (1, 2), (2, 0))])
maxe = np.maximum.reduce([np.hypot(X[t[:, 0]] - X[t[:, 1]], Y[t[:, 0]] - Y[t[:, 1]]),
np.hypot(X[t[:, 1]] - X[t[:, 2]], Y[t[:, 1]] - Y[t[:, 2]]),
np.hypot(X[t[:, 2]] - X[t[:, 0]], Y[t[:, 2]] - Y[t[:, 0]])])
tri.set_mask(maxe > 4.5 * np.median(e))
xlo, xhi = np.percentile(X, 1) - 0.3, np.percentile(X, 99) + 0.3
ylo, yhi = np.percentile(Y, 1) - 0.3, np.percentile(Y, 99) + 0.3
def panel(c, val, title, vmin, vmax, cmap):
lv = np.linspace(vmin, vmax, 25)
cs = axes[r, c].tricontourf(tri, val[m], levels=lv, cmap=cmap, extend="both")
axes[r, c].set_title(title, fontsize=9)
axes[r, c].set_aspect("equal"); axes[r, c].set_xticks([]); axes[r, c].set_yticks([])
axes[r, c].set_xlim(xlo, xhi); axes[r, c].set_ylim(ylo, yhi)
plt.colorbar(cs, ax=axes[r, c], fraction=0.046, pad=0.02)
vU = (float(min(uT.min(), uP.min())), float(max(uT.max(), uP.max())))
vP = (float(min(pT.min(), pP.min())), float(max(pT.max(), pP.max())))
eU = np.abs(uP - uT); eP = np.abs(pP - pT)
panel(0, uT, f"{g.cid}\nRe={Re:.0f} |U| TRUE", *vU, "viridis")
panel(1, uP, f"|U| PRED (R2={r2u:.3f})", *vU, "viridis")
panel(2, eU, "|U| |error|", 0.0, float(eU[m].max() + 1e-9), "magma")
panel(3, pT, "p TRUE", *vP, "coolwarm")
panel(4, pP, f"p PRED (R2={r2p:.3f})", *vP, "coolwarm")
panel(5, eP, "p |error|", 0.0, float(eP[m].max() + 1e-9), "magma")
print(f" {setname:12s} {g.cid:24s} Re={Re:7.0f} |U|R2={r2u:.3f} pR2={r2p:.3f} (slice n={int(m.sum())}, tol={tol:.2f})", flush=True)
fig.suptitle(f"OOD: {setname} - actual vs predicted ({a.ckpt} ep{ep}, z~0 slice)", fontsize=13)
fig.tight_layout(rect=[0, 0, 1, 0.96])
outp = os.path.join(a.out, f"ood_{setname}.png")
fig.savefig(outp, dpi=110); plt.close(fig)
print(f"saved {outp}", flush=True)
print("PLOT_OOD_DONE", flush=True)
if __name__ == "__main__":
main()