| """Turn the ADFTD ablation result JSONs into the Claim 3 / Claim 4 tables and figures.""" |
|
|
| import json |
| import os |
|
|
| import numpy as np |
| import plotly.graph_objects as go |
|
|
| RES = "results/adftd" |
| OUT = "results" |
|
|
| |
| PAPER = { |
| "tsfp_scratch": {"accuracy": 48.83, "f1": 45.80, "auroc": 65.29}, |
| "tsfp_rec": {"accuracy": 53.51, "f1": 50.48, "auroc": 70.21}, |
| "tsfp_rec_div": {"accuracy": 53.91, "f1": 51.79, "auroc": 72.11}, |
| "timae_scratch": {"accuracy": 50.27, "f1": 45.88}, |
| "timae_pretrain": {"accuracy": 50.48, "f1": 45.44}, |
| "simmtm_scratch": {"accuracy": 51.98, "f1": 43.66}, |
| "simmtm_pretrain": {"accuracy": 52.92, "f1": 45.59}, |
| } |
| LABEL = { |
| "tsfp_scratch": "TS-Fingerprint · Scratch", |
| "tsfp_rec": "TS-Fingerprint · Pre-trained (L_rec)", |
| "tsfp_rec_div": "TS-Fingerprint · Pre-trained (L_rec + L_div)", |
| "timae_scratch": "Ti-MAE · Scratch", |
| "timae_pretrain": "Ti-MAE · Pre-trained", |
| "simmtm_scratch": "SimMTM · Scratch", |
| "simmtm_pretrain": "SimMTM · Pre-trained", |
| } |
|
|
|
|
| def load(): |
| out = {} |
| for name in PAPER: |
| p = os.path.join(RES, f"{name}.json") |
| if os.path.exists(p) and os.path.getsize(p) > 0: |
| out[name] = json.load(open(p)) |
| return out |
|
|
|
|
| def fmt(agg, m): |
| return f"{agg[m]['mean']:.2f} ± {agg[m]['std']:.2f}" |
|
|
|
|
| def main(): |
| r = load() |
| print(f"available configs: {sorted(r)}\n") |
| rows = [] |
| for name in PAPER: |
| if name not in r: |
| continue |
| a = r[name]["aggregate"] |
| rows.append({ |
| "config": name, "label": LABEL[name], |
| "acc": a["accuracy"]["mean"], "acc_sd": a["accuracy"]["std"], |
| "f1": a["f1"]["mean"], "f1_sd": a["f1"]["std"], |
| "auroc": a["auroc"]["mean"], "auroc_sd": a["auroc"]["std"], |
| "paper_acc": PAPER[name].get("accuracy"), |
| "paper_f1": PAPER[name].get("f1"), |
| "paper_auroc": PAPER[name].get("auroc"), |
| "n_seeds": len(r[name]["runs"]), |
| "sec_per_seed": float(np.mean([x["seconds"] for x in r[name]["runs"]])), |
| "params": r[name]["runs"][0]["params"], |
| }) |
|
|
| print(f"{'config':<18}{'Acc (ours)':>16}{'Acc (paper)':>13}" |
| f"{'F1 (ours)':>16}{'F1 (paper)':>12}{'AUROC (ours)':>17}{'AUROC (paper)':>15}") |
| for x in rows: |
| pa = f"{x['paper_auroc']:.2f}" if x["paper_auroc"] else "-" |
| print(f"{x['config']:<18}{x['acc']:>9.2f}±{x['acc_sd']:<5.2f}" |
| f"{x['paper_acc']:>13.2f}{x['f1']:>9.2f}±{x['f1_sd']:<5.2f}" |
| f"{x['paper_f1']:>12.2f}{x['auroc']:>10.2f}±{x['auroc_sd']:<5.2f}{pa:>15}") |
|
|
| |
| def promo(scratch, pre, key="f1"): |
| if scratch not in r or pre not in r: |
| return None |
| s = r[scratch]["aggregate"][key]["mean"] |
| p = r[pre]["aggregate"][key]["mean"] |
| return {"scratch": s, "pretrained": p, "abs": p - s, "rel_pct": (p - s) / s * 100} |
|
|
| promos = { |
| "TS-Fingerprint": {"f1": promo("tsfp_scratch", "tsfp_rec_div", "f1"), |
| "accuracy": promo("tsfp_scratch", "tsfp_rec_div", "accuracy")}, |
| "Ti-MAE": {"f1": promo("timae_scratch", "timae_pretrain", "f1"), |
| "accuracy": promo("timae_scratch", "timae_pretrain", "accuracy")}, |
| "SimMTM": {"f1": promo("simmtm_scratch", "simmtm_pretrain", "f1"), |
| "accuracy": promo("simmtm_scratch", "simmtm_pretrain", "accuracy")}, |
| } |
| |
| div_effect = promo("tsfp_rec", "tsfp_rec_div", "f1") |
| print("\nrelative F1 promotion (scratch -> pre-trained):") |
| for k, v in promos.items(): |
| if v["f1"]: |
| print(f" {k:<16} {v['f1']['scratch']:.2f} -> {v['f1']['pretrained']:.2f} " |
| f"(+{v['f1']['rel_pct']:.2f}%)") |
| if div_effect: |
| print(f"\nL_div increment on top of L_rec: F1 {div_effect['scratch']:.2f} -> " |
| f"{div_effect['pretrained']:.2f} ({div_effect['abs']:+.2f} pts, " |
| f"{div_effect['rel_pct']:+.2f}%)") |
|
|
| |
| paired = None |
| if "tsfp_rec" in r and "tsfp_rec_div" in r: |
| a = {x["seed"]: x["test"]["f1"] for x in r["tsfp_rec"]["runs"]} |
| b = {x["seed"]: x["test"]["f1"] for x in r["tsfp_rec_div"]["runs"]} |
| seeds = sorted(set(a) & set(b)) |
| d = [b[s] - a[s] for s in seeds] |
| paired = {"seeds": seeds, "per_seed_delta_f1": d, |
| "mean_delta": float(np.mean(d)), |
| "n_positive": int(sum(x > 0 for x in d))} |
| print(f"\npaired per-seed F1 delta (rec+div minus rec): " |
| f"{[round(x, 2) for x in d]} mean {np.mean(d):+.2f}, " |
| f"{paired['n_positive']}/{len(d)} seeds positive") |
|
|
| json.dump({"rows": rows, "promotions": promos, "div_effect": div_effect, |
| "paired_div_test": paired}, |
| open(os.path.join(OUT, "adftd_summary.json"), "w"), indent=2) |
|
|
| |
| |
| order = ["tsfp_scratch", "tsfp_rec", "tsfp_rec_div"] |
| have = [x for x in order if x in r] |
| if have: |
| xs = ["Scratch", "Pre-trained\n(L_rec)", "Pre-trained\n(L_rec + L_div)"][:len(have)] |
| fig = go.Figure() |
| for m, col in (("f1", "#2F6F8F"), ("accuracy", "#B07C2B"), ("auroc", "#7E7E7E")): |
| ours = [r[c]["aggregate"][m]["mean"] for c in have] |
| sd = [r[c]["aggregate"][m]["std"] for c in have] |
| pap = [PAPER[c].get(m) for c in have] |
| fig.add_bar(name=f"ours · {m}", x=xs, y=ours, |
| error_y=dict(type="data", array=sd), marker_color=col) |
| fig.add_scatter(name=f"paper · {m}", x=xs, y=pap, mode="markers", |
| marker=dict(symbol="line-ew", size=26, line=dict( |
| width=3, color=col))) |
| fig.update_layout( |
| title="Claim 3 — ADFTD ablation: this reproduction (bars, 3 seeds) " |
| "vs the paper's Table 2 (ticks)", |
| yaxis_title="score (%)", barmode="group", template="plotly_white", |
| height=460, legend=dict(orientation="h", y=-0.18)) |
| fig.write_html(os.path.join(OUT, "claim3_ablation.html"), |
| include_plotlyjs="cdn") |
|
|
| |
| if all(promos[k]["f1"] for k in promos): |
| fig2 = go.Figure() |
| names = list(promos) |
| fig2.add_bar(name="this reproduction", x=names, |
| y=[promos[k]["f1"]["rel_pct"] for k in names], |
| marker_color="#2F6F8F", |
| text=[f"{promos[k]['f1']['rel_pct']:+.2f}%" for k in names], |
| textposition="outside") |
| paper_promo = {"TS-Fingerprint": 13.07, "Ti-MAE": -0.96, "SimMTM": 4.42} |
| fig2.add_bar(name="paper (Table 4)", x=names, |
| y=[paper_promo[k] for k in names], marker_color="#B07C2B", |
| text=[f"{paper_promo[k]:+.2f}%" for k in names], |
| textposition="outside") |
| fig2.update_layout( |
| title="Claim 4 — relative F1 promotion from pre-training on ADFTD", |
| yaxis_title="relative F1 gain (%)", barmode="group", |
| template="plotly_white", height=440, |
| legend=dict(orientation="h", y=-0.15)) |
| fig2.write_html(os.path.join(OUT, "claim4_promotion.html"), |
| include_plotlyjs="cdn") |
|
|
| |
| with open(os.path.join(OUT, "adftd_results.csv"), "w") as f: |
| f.write("config,label,n_seeds,acc,acc_sd,f1,f1_sd,auroc,auroc_sd," |
| "paper_acc,paper_f1,paper_auroc,sec_per_seed,params\n") |
| for x in rows: |
| f.write(f"{x['config']},{x['label']},{x['n_seeds']},{x['acc']:.3f}," |
| f"{x['acc_sd']:.3f},{x['f1']:.3f},{x['f1_sd']:.3f}," |
| f"{x['auroc']:.3f},{x['auroc_sd']:.3f},{x['paper_acc']}," |
| f"{x['paper_f1']},{x['paper_auroc']},{x['sec_per_seed']:.0f}," |
| f"{x['params']}\n") |
| print("\nwrote results/adftd_summary.json, adftd_results.csv, " |
| "claim3_ablation.html, claim4_promotion.html") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|