"""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()