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