File size: 3,497 Bytes
bdce880
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""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()