| """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() |
| 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) |
| Re = float(getattr(g, "Re", 0.0)) |
|
|
| |
| 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] |
|
|
| |
| 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() |
|
|