| """Claim 6: summarise the k x mask-ratio sweep on PTB-XL and compare with Table 3.""" |
|
|
| import json |
| import os |
| import re |
|
|
| import numpy as np |
| import plotly.graph_objects as go |
| from plotly.subplots import make_subplots |
|
|
| RES = "results/sweep" |
| KS, RS = [6, 8, 10], [0.5, 0.6, 0.7, 0.8] |
| |
| PAPER_F1 = {(6, 0.5): 54.81, (6, 0.6): 56.68, (6, 0.7): 55.41, (6, 0.8): 54.32, |
| (8, 0.5): 62.60, (8, 0.6): 63.51, (8, 0.7): 62.10, (8, 0.8): 59.19, |
| (10, 0.5): 51.76, (10, 0.6): 50.70, (10, 0.7): 56.33, (10, 0.8): 55.70} |
| PAPER_ACC = {(6, 0.5): 70.19, (6, 0.6): 70.71, (6, 0.7): 70.63, (6, 0.8): 69.77, |
| (8, 0.5): 74.07, (8, 0.6): 73.98, (8, 0.7): 73.86, (8, 0.8): 72.35, |
| (10, 0.5): 67.36, (10, 0.6): 66.41, (10, 0.7): 71.07, (10, 0.8): 70.42} |
|
|
|
|
| def main(): |
| rows = [] |
| for k in KS: |
| for r in RS: |
| p = os.path.join(RES, f"ptbxl_k{k}_r{r}.json") |
| if not (os.path.exists(p) and os.path.getsize(p) > 0): |
| continue |
| d = json.load(open(p)) |
| a = d["aggregate"] |
| rows.append({"k": k, "r": r, |
| "f1": a["f1"]["mean"], "f1_sd": a["f1"]["std"], |
| "acc": a["accuracy"]["mean"], "acc_sd": a["accuracy"]["std"], |
| "auroc": a["auroc"]["mean"], |
| "n_seeds": len(d["runs"]), |
| "paper_f1": PAPER_F1[(k, r)], "paper_acc": PAPER_ACC[(k, r)], |
| "sec": float(np.mean([x["seconds"] for x in d["runs"]]))}) |
| if not rows: |
| print("no sweep results yet") |
| return |
| print(f"{'k':>3}{'r':>6}{'F1 ours':>12}{'F1 paper':>11}{'Acc ours':>11}" |
| f"{'Acc paper':>11}{'seeds':>7}") |
| for x in rows: |
| print(f"{x['k']:>3}{x['r']:>6}{x['f1']:>9.2f}±{x['f1_sd']:<3.1f}" |
| f"{x['paper_f1']:>11.2f}{x['acc']:>8.2f}±{x['acc_sd']:<3.1f}" |
| f"{x['paper_acc']:>11.2f}{x['n_seeds']:>7}") |
|
|
| out = {"rows": rows} |
| if len(rows) == 12: |
| f1 = {(x["k"], x["r"]): x["f1"] for x in rows} |
| best = max(f1, key=f1.get) |
| out["best_cell_ours"] = {"k": best[0], "r": best[1], "f1": f1[best]} |
| out["best_cell_paper"] = {"k": 8, "r": 0.6, "f1": 63.51} |
| |
| out["best_k_per_r"] = {str(r): max(KS, key=lambda k: f1[(k, r)]) for r in RS} |
| out["best_r_per_k"] = {str(k): max(RS, key=lambda r: f1[(k, r)]) for k in KS} |
| out["mean_f1_by_k"] = {str(k): float(np.mean([f1[(k, r)] for r in RS])) |
| for k in KS} |
| out["mean_f1_by_r"] = {str(r): float(np.mean([f1[(k, r)] for k in KS])) |
| for r in RS} |
| out["paper_mean_f1_by_k"] = { |
| str(k): float(np.mean([PAPER_F1[(k, r)] for r in RS])) for k in KS} |
| out["paper_mean_f1_by_r"] = { |
| str(r): float(np.mean([PAPER_F1[(k, r)] for k in KS])) for r in RS} |
| |
| deg = {} |
| for k in KS: |
| lo = np.mean([f1[(k, r)] for r in (0.5, 0.6, 0.7)]) |
| deg[str(k)] = {"mean_r_le_0.7": float(lo), "r_0.8": f1[(k, 0.8)], |
| "delta": float(f1[(k, 0.8)] - lo)} |
| out["r08_degradation"] = deg |
| |
| ours_v = np.array([f1[(k, r)] for k in KS for r in RS]) |
| pap_v = np.array([PAPER_F1[(k, r)] for k in KS for r in RS]) |
| out["pearson_ours_vs_paper_grid"] = float(np.corrcoef(ours_v, pap_v)[0, 1]) |
| out["spearman_ours_vs_paper_grid"] = float(np.corrcoef( |
| np.argsort(np.argsort(ours_v)), np.argsort(np.argsort(pap_v)))[0, 1]) |
| print("\nbest cell (ours):", out["best_cell_ours"]) |
| print("mean F1 by k (ours):", {k: round(v, 2) |
| for k, v in out["mean_f1_by_k"].items()}) |
| print("mean F1 by k (paper):", {k: round(v, 2) |
| for k, v in out["paper_mean_f1_by_k"].items()}) |
| print("mean F1 by r (ours):", {k: round(v, 2) |
| for k, v in out["mean_f1_by_r"].items()}) |
| print("mean F1 by r (paper):", {k: round(v, 2) |
| for k, v in out["paper_mean_f1_by_r"].items()}) |
| print("r=0.8 vs r<=0.7:", {k: round(v["delta"], 2) |
| for k, v in deg.items()}) |
| print(f"grid correlation with Table 3: Pearson " |
| f"{out['pearson_ours_vs_paper_grid']:.3f}, " |
| f"Spearman {out['spearman_ours_vs_paper_grid']:.3f}") |
|
|
| json.dump(out, open("results/sweep_summary.json", "w"), indent=2) |
|
|
| |
| ours = np.full((3, 4), np.nan) |
| for x in rows: |
| ours[KS.index(x["k"]), RS.index(x["r"])] = x["f1"] |
| paper = np.array([[PAPER_F1[(k, r)] for r in RS] for k in KS]) |
| fig = make_subplots(rows=1, cols=2, horizontal_spacing=0.13, |
| subplot_titles=("this reproduction (PTB-XL, thinned)", |
| "paper (Table 3)")) |
| for j, m in enumerate((ours, paper)): |
| fig.add_heatmap(z=m, x=[f"r={r}" for r in RS], y=[f"k={k}" for k in KS], |
| colorscale=[[0, "#F2F7F8"], [1, "#17697B"]], showscale=False, |
| text=[[("" if np.isnan(v) else f"{v:.2f}") for v in row] |
| for row in m], |
| texttemplate="%{text}", row=1, col=j + 1) |
| fig.update_layout(template="plotly_white", height=420, |
| title="Claim 6 — macro F1 (%) over bottleneck size k and mask ratio r") |
| fig.write_html("results/claim6_sweep.html", include_plotlyjs="cdn") |
|
|
| with open("results/sweep_results.csv", "w") as f: |
| f.write("k,mask_ratio,n_seeds,f1,f1_sd,acc,acc_sd,auroc,paper_f1,paper_acc,sec\n") |
| for x in rows: |
| f.write(f"{x['k']},{x['r']},{x['n_seeds']},{x['f1']:.3f},{x['f1_sd']:.3f}," |
| f"{x['acc']:.3f},{x['acc_sd']:.3f},{x['auroc']:.3f}," |
| f"{x['paper_f1']},{x['paper_acc']},{x['sec']:.0f}\n") |
| print("\nwrote results/sweep_summary.json, sweep_results.csv, claim6_sweep.html") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|