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