steady-rans-surrogates / code /ezflow_v3 /gnn /ensemble_eval.py
BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
7.75 kB
"""Ensemble evaluation of the 3 trained models (GeoReNet, Transolver, HybridFlow).
All three predict on the SAME sampled nodes and share the SAME norms, so we combine
per node / per channel in raw physical units. We test the training-free schemes:
single models, uniform mean(3), all pairwise means, per-channel val-weighted mean,
per-channel hard selection (best model per channel by val R^2).
R^2 is accumulated incrementally (SS_res / SS_tot) so we never hold all predictions
in memory. Same deterministic split / norms as paper2/eval_table.py.
python ezflow_v3/gnn/ensemble_eval.py --device cpu
"""
import os
os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE")
os.environ.setdefault("OMP_NUM_THREADS", "4")
import sys, csv, argparse
import numpy as np, torch
sys.path.insert(0, r"C:\dev\EZFlow")
from torch_geometric.loader import DataLoader
from ezflow_v3.gnn.etl import CaseDatasetV2
from ezflow_v3.gnn import features as F
from ezflow_v3.gnn.train_v5 import split
from ezflow_v3.gnn.model_v5 import MeshGraphNetV5
RUNS = r"C:\dev\EZFlow\ezflow_v3\gnn\_runs"
CACHE = r"C:\dev\ezflow_eval\cache_v3"
CH = ["u", "v", "w", "p", "logk", "logom", "lognut"]
NC = len(CH)
MODELS = ["G", "T", "H"] # GeoReNet, Transolver, HybridFlow
TAG = {"G": "rans_v5", "T": "transolver", "H": "hybrid"}
def pick_ckpt(run):
for n in ("best.pt", "model.pt", "ckpt.pt"):
p = os.path.join(run, n)
if os.path.exists(p):
return p
return None
def build(tag, dev):
ck = torch.load(pick_ckpt(os.path.join(RUNS, tag)), map_location="cpu", weights_only=False)
ar = ck["args"]
if ar.get("model") == "transolver":
from ezflow_v3.baselines.transolver_wrap import TransolverWrapper
m = 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)
elif ar.get("model") == "hybrid":
from ezflow_v3.gnn.model_hybrid import HybridFlow
m = HybridFlow(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM, hidden=int(ar["hidden"]),
K=8, out_dim=F.TARGET_DIM, slice_num=64)
else:
m = MeshGraphNetV5(F.NODE_FEATURE_DIM, 4, F.GLOBAL_DIM, hidden=int(ar["hidden"]),
K=int(ar["K"]), 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"))
m.load_state_dict(ck["model"]); m.eval().to(dev)
return m, int(ck.get("epoch", -1))
class Acc:
"""Incremental per-channel R^2 over a set, for several prediction combos."""
def __init__(self):
self.n = 0; self.sy = np.zeros(NC); self.sy2 = np.zeros(NC)
self.sres = {}
def add(self, y, preds): # y:[N,NC] (once); preds: dict name->[N,NC]
self.n += len(y); self.sy += y.sum(0); self.sy2 += (y * y).sum(0)
for k, p in preds.items():
d = y - p
self.sres.setdefault(k, np.zeros(NC))
self.sres[k] += (d * d).sum(0)
def r2(self, key): # per-channel R^2 array
sstot = self.sy2 - self.sy ** 2 / max(self.n, 1)
return 1.0 - self.sres[key] / sstot
@torch.no_grad()
def model_preds(models, graphs, ymt, yst, gmt, gst, dev):
"""Yield (y_raw[N,NC], {G,T,H: pred_raw[N,NC]}) per graph, node-aligned."""
for b in DataLoader(graphs, batch_size=1):
b = b.to(dev)
b.global_feat = (b.global_feat - gmt) / gst # shared norm, all models
y = b.y.cpu().numpy()
preds = {}
for name in MODELS:
out = models[name](b) * yst + ymt
preds[name] = out.cpu().numpy()
yield y, preds
def combos(preds, w=None, sel=None):
"""Build all ensemble predictions from single-model preds dict (numpy)."""
G, T, H = preds["G"], preds["T"], preds["H"]
out = {"G": G, "T": T, "H": H,
"mean(GTH)": (G + T + H) / 3.0,
"mean(G,T)": (G + T) / 2.0,
"mean(G,H)": (G + H) / 2.0,
"mean(T,H)": (T + H) / 2.0}
if w is not None: # per-channel weighted mean (weights [3,NC])
out["wmean(val)"] = w[0] * G + w[1] * T + w[2] * H
if sel is not None: # per-channel hard pick (sel:[NC] in 0/1/2)
stack = np.stack([G, T, H], 0) # [3,N,NC]
out["select(val)"] = np.take_along_axis(
stack, sel[None, None, :].repeat(G.shape[0], 1), axis=0)[0]
return out
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--device", default="cpu")
ap.add_argument("--out", default=r"C:\dev\ezflow_eval\ensemble_eval.csv")
a = ap.parse_args()
dev = torch.device(a.device if (a.device == "cpu" or torch.cuda.is_available()) else "cpu")
print(f"device {dev}", flush=True)
ds = CaseDatasetV2(CACHE); _, val, ood = split(ds)
sets = [("val", val), ("ood_Re", ood["ood_Re_band"]),
("ood_modelnet", ood["ood_modelnet"]), ("ood_family", ood["ood_family"])]
print("split:", {k: len(v) for k, v in sets}, flush=True)
models = {}; eps = {}
for name in MODELS:
models[name], eps[name] = build(TAG[name], dev)
print(f"loaded {name} ({TAG[name]}) ep{eps[name]}", flush=True)
nz = np.load(os.path.join(RUNS, TAG["G"], "norms.npz"))
ymt = torch.tensor(nz["y_mean"], device=dev); yst = torch.tensor(nz["y_std"], device=dev)
gmt = torch.tensor(nz["g_mean"], device=dev); gst = torch.tensor(nz["g_std"], device=dev)
# ---- phase 1: per-channel val R^2 of each single model -> weights + selection ----
av = Acc()
for y, preds in model_preds(models, val, ymt, yst, gmt, gst, dev):
av.add(y, preds)
r2val = np.stack([av.r2("G"), av.r2("T"), av.r2("H")], 0) # [3,NC]
w = np.clip(r2val, 0, None) ** 2
w = w / w.sum(0, keepdims=True) # [3,NC] weights
sel = r2val.argmax(0) # [NC] best model per channel
print("per-channel val R^2 (G/T/H):", flush=True)
for c in range(NC):
print(f" {CH[c]:7s} G {r2val[0,c]:.3f} T {r2val[1,c]:.3f} H {r2val[2,c]:.3f}"
f" -> best {MODELS[sel[c]]}", flush=True)
# ---- phase 2: all combos on all 4 sets ----
accs = {name: Acc() for name, _ in sets}
for name, graphs in sets:
for y, preds in model_preds(models, graphs, ymt, yst, gmt, gst, dev):
accs[name].add(y, combos(preds, w=w, sel=sel))
keys = ["G", "T", "H", "mean(GTH)", "mean(G,T)", "mean(G,H)", "mean(T,H)",
"wmean(val)", "select(val)"]
label = {"G": "GeoReNet", "T": "Transolver", "H": "HybridFlow"}
hdr = f"{'combo':14s} " + " ".join(f"{s:>12s}" for s, _ in sets) + " mean4"
print("\n" + hdr); print("-" * len(hdr), flush=True)
rows = []
for k in keys:
means = [float(np.mean(accs[name].r2(k))) for name, _ in sets]
m4 = float(np.mean(means))
nm = label.get(k, k)
print(f"{nm:14s} " + " ".join(f"{v:12.4f}" for v in means) + f" {m4:8.4f}", flush=True)
rows.append(dict(combo=nm, val=round(means[0], 4), ood_Re=round(means[1], 4),
ood_modelnet=round(means[2], 4), ood_family=round(means[3], 4),
mean4=round(m4, 4)))
with open(a.out, "w", newline="") as f:
wtr = csv.DictWriter(f, fieldnames=["combo", "val", "ood_Re", "ood_modelnet",
"ood_family", "mean4"])
wtr.writeheader(); wtr.writerows(rows)
print(f"\nwrote {a.out}\nENSEMBLE_DONE", flush=True)
if __name__ == "__main__":
main()