| """Model-agnostic evaluation of a trained run on the FIXED val / OOD splits. |
| Works for georenet, transolver, transolverpp, and hybrid (reads the architecture |
| from the checkpoint's stored args). Reports mean field R2 (and per channel) on |
| val, OOD-ModelNet, and OOD-Re, in the same target space the training uses. |
| |
| python ezflow_v3/gnn/eval_run.py --run tpp_s0 --cache /content/cache_v3 |
| """ |
| import os |
| os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"; os.environ.setdefault("OMP_NUM_THREADS", "4") |
| import sys, argparse |
| import numpy as np, torch |
| HERE = os.path.dirname(os.path.abspath(__file__)) |
| REPO = os.path.dirname(os.path.dirname(HERE)) |
| sys.path.insert(0, REPO) |
| from torch_geometric.loader import DataLoader |
| from ezflow_v3.gnn import features as F |
| from ezflow_v3.gnn.etl import CaseDatasetV2 |
| from ezflow_v3.gnn.train_v5 import split, r2 |
| from ezflow_v3.gnn.model_v5 import MeshGraphNetV5 |
|
|
|
|
| def build(ar): |
| m = ar.get("model", "georenet") |
| if m == "transolver": |
| from ezflow_v3.baselines.transolver_wrap import TransolverWrapper |
| return TransolverWrapper(node_in=F.NODE_FEATURE_DIM, space_dim=3, global_dim=F.GLOBAL_DIM, |
| out_dim=F.TARGET_DIM, n_hidden=256, n_layers=8, slice_num=64) |
| if m == "transolverpp": |
| from ezflow_v3.baselines.transolverpp_wrap import TransolverPlusWrapper |
| return TransolverPlusWrapper(node_in=F.NODE_FEATURE_DIM, space_dim=3, global_dim=F.GLOBAL_DIM, |
| out_dim=F.TARGET_DIM, n_hidden=256, n_layers=8, slice_num=64) |
| if m == "hybrid": |
| from ezflow_v3.gnn.model_hybrid import HybridFlow |
| return HybridFlow(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM, hidden=int(ar.get("hidden", 160)), |
| K=8, out_dim=F.TARGET_DIM, slice_num=64, use_local=not ar.get("no_local", False)) |
| return MeshGraphNetV5(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM, hidden=int(ar.get("hidden", 128)), |
| K=int(ar.get("K", 12)), 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 main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--run", required=True, help="run tag (under gnn/_runs) or a full run dir path") |
| ap.add_argument("--cache", default=r"/content/cache_v3") |
| ap.add_argument("--device", default="cuda") |
| a = ap.parse_args() |
| run = a.run if os.path.isdir(a.run) else os.path.join(HERE, "_runs", a.run) |
| dev = torch.device(a.device if (torch.cuda.is_available() or a.device == "cpu") else "cpu") |
|
|
| nz = np.load(os.path.join(run, "norms.npz")) |
| ym = torch.tensor(nz["y_mean"], device=dev); ys = torch.tensor(nz["y_std"], device=dev) |
| gm = torch.tensor(nz["g_mean"], device=dev); gs = torch.tensor(nz["g_std"], device=dev) |
| ckp = os.path.join(run, "ckpt.pt") |
| if not os.path.exists(ckp): |
| ckp = os.path.join(run, "model.pt") |
| ck = torch.load(ckp, map_location=dev, weights_only=False) |
| ar = ck["args"] |
| model = build(ar).to(dev); model.load_state_dict(ck["model"]); model.eval() |
|
|
| ds = CaseDatasetV2(a.cache); _, val, ood = split(ds) |
|
|
| @torch.no_grad() |
| def ev(graphs): |
| if not graphs: |
| return float("nan"), [] |
| P, T = [], [] |
| for b in DataLoader(graphs, batch_size=1): |
| b = b.to(dev); b.global_feat = (b.global_feat - gm) / gs |
| out = model(b); node = out[0] if isinstance(out, tuple) else out |
| P.append((node * ys + ym).cpu().numpy()); T.append(b.y.cpu().numpy()) |
| P = np.concatenate(P); T = np.concatenate(T) |
| rc = r2(P, T) |
| return float(np.mean([float(x) for x in rc])), [round(float(x), 3) for x in rc] |
|
|
| print(f"run={os.path.basename(run)} model={ar.get('model')} epoch={ck.get('epoch')} device={dev}", flush=True) |
| print("channels: [u v w p logk logomega lognutRe]", flush=True) |
| for name, g in [("val", val), ("ood_modelnet", ood["ood_modelnet"]), ("ood_Re", ood["ood_Re_band"])]: |
| mr, rc = ev(g) |
| print(f" {name:13s} meanR2 {mr:.4f} perchan {rc}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|