steady-rans-surrogates / easy_eval.py
BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
3.5 kB
"""One command to reproduce the fixed split table from the released weights.
python easy_eval.py --cache ./data/cache_v3 --weights ./weights [--device cuda] [--runs r1,r2]
Rebuilds each architecture from its stored args, evaluates the deterministic
val / OOD object / OOD Reynolds splits, and prints mean field R2 per run.
Writes easy_eval_results.csv next to this script.
"""
import os
os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"; os.environ.setdefault("OMP_NUM_THREADS", "4")
import sys, csv, argparse
import numpy as np
import torch
HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(HERE, "code"))
from torch_geometric.loader import DataLoader
from ezflow_v3.gnn.etl import CaseDatasetV2
from ezflow_v3.gnn.train_v5 import split, r2
from ezflow_v3.gnn.eval_run import build
DEFAULT_RUNS = ("geore_fieldonly_s0,hybrid_s0,tpp_s0,pfaff_s0,"
"geore_fieldonly_s1,geore_fieldonly_s2,hybrid_s1,hybrid_s2,"
"tpp_s1,tpp_s2,pfaff_s1,pfaff_s2,"
"hybrid_nolocal_s0,hybrid_nolocal_h300_s0,geore_noglobal_s0")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--cache", required=True, help="path to the cache_v3 folder of the dataset repo")
ap.add_argument("--weights", default=os.path.join(HERE, "weights"))
ap.add_argument("--runs", default=DEFAULT_RUNS, help="comma separated run folder names")
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
a = ap.parse_args()
dev = torch.device(a.device if (torch.cuda.is_available() or a.device == "cpu") else "cpu")
print(f"loading dataset from {a.cache} ...", flush=True)
ds = CaseDatasetV2(a.cache)
_, val, ood = split(ds)
splits = [("val", val), ("ood_object", ood["ood_modelnet"]), ("ood_Re", ood["ood_Re_band"])]
for nm, g in splits:
print(f" split {nm}: {len(g)} cases", flush=True)
rows = []
for run in [r.strip() for r in a.runs.split(",") if r.strip()]:
rd = os.path.join(a.weights, run)
nz = np.load(os.path.join(rd, "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)
ck = torch.load(os.path.join(rd, "model.pt"), map_location=dev, weights_only=False)
model = build(ck["args"]).to(dev); model.load_state_dict(ck["model"]); model.eval()
@torch.no_grad()
def ev(graphs):
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())
rc = r2(np.concatenate(P), np.concatenate(T))
return float(np.mean([float(x) for x in rc]))
row = {"run": run, "model": ck["args"].get("model")}
line = f"{run:26s}"
for nm, g in splits:
row[nm] = round(ev(g), 4)
line += f" {nm} {row[nm]:.4f}"
rows.append(row)
print(line, flush=True)
out_csv = os.path.join(HERE, "easy_eval_results.csv")
with open(out_csv, "w", newline="") as f:
w = csv.DictWriter(f, fieldnames=list(rows[0].keys())); w.writeheader(); w.writerows(rows)
print(f"wrote {out_csv}", flush=True)
if __name__ == "__main__":
main()