BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
4.23 kB
"""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()