| """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"] |
| 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): |
| 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): |
| 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 |
| 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: |
| out["wmean(val)"] = w[0] * G + w[1] * T + w[2] * H |
| if sel is not None: |
| stack = np.stack([G, T, H], 0) |
| 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) |
|
|
| |
| 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) |
| w = np.clip(r2val, 0, None) ** 2 |
| w = w / w.sum(0, keepdims=True) |
| sel = r2val.argmax(0) |
| 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) |
|
|
| |
| 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() |
|
|