| """Render the poster's PNG figures from the logged reproduction results. |
| |
| Every number here comes from results/*.json produced by the runs; nothing is |
| hand-entered except the paper's own reported values, which are labelled as such. |
| """ |
|
|
| import json |
| import os |
|
|
| import numpy as np |
| import plotly.graph_objects as go |
| from plotly.subplots import make_subplots |
|
|
| OUT = "poster/images" |
| os.makedirs(OUT, exist_ok=True) |
| ACCENT, GOLD, GREY, RED = "#17697B", "#B07C2B", "#8A8A8A", "#A33B3B" |
| FONT = dict(family="Helvetica, Arial, sans-serif", size=30, color="#1A1A1A") |
| SCALE = 3 |
|
|
|
|
| def save(fig, name, w, h): |
| fig.update_layout(template="plotly_white", font=FONT, |
| margin=dict(l=90, r=30, t=80, b=70)) |
| fig.write_image(f"{OUT}/{name}.png", width=w, height=h, scale=SCALE) |
| print(f"wrote {OUT}/{name}.png ({w*SCALE}x{h*SCALE} px)") |
|
|
|
|
| |
| def fig_ablation(): |
| s = json.load(open("results/adftd_summary.json")) |
| rows = {r["config"]: r for r in s["rows"]} |
| order = ["tsfp_scratch", "tsfp_rec", "tsfp_rec_div"] |
| order = [o for o in order if o in rows] |
| if len(order) < 3: |
| print("skip ablation fig (incomplete)") |
| return |
| xs = ["Scratch", "Pre-trained<br>(ℒ_rec)", "Pre-trained<br>(ℒ_rec + ℒ_div)"] |
| fig = go.Figure() |
| fig.add_bar(name="this reproduction (3 seeds)", x=xs, |
| y=[rows[o]["f1"] for o in order], |
| error_y=dict(type="data", array=[rows[o]["f1_sd"] for o in order], |
| thickness=3, width=14), |
| marker_color=ACCENT, |
| text=[f"{rows[o]['f1']:.2f}" for o in order], |
| textposition="outside", textfont=dict(size=32)) |
| fig.add_bar(name="paper (Table 2)", x=xs, |
| y=[rows[o]["paper_f1"] for o in order], marker_color=GOLD, |
| text=[f"{rows[o]['paper_f1']:.2f}" for o in order], |
| textposition="outside", textfont=dict(size=32)) |
| fig.update_layout(barmode="group", yaxis_title="ADFTD macro F1 (%)", |
| yaxis_range=[38, 58], height=620, width=1000, |
| legend=dict(orientation="h", y=1.13, x=0)) |
| save(fig, "claim3_ablation", 1000, 620) |
|
|
|
|
| |
| def fig_promotion(): |
| s = json.load(open("results/adftd_summary.json")) |
| p = s["promotions"] |
| names, ours = [], [] |
| for k in ("TS-Fingerprint", "SimMTM", "Ti-MAE"): |
| if p[k]["f1"]: |
| names.append(k) |
| ours.append(p[k]["f1"]["rel_pct"]) |
| if not names: |
| print("skip promotion fig") |
| return |
| paper = {"TS-Fingerprint": 13.07, "SimMTM": 4.42, "Ti-MAE": -0.96} |
| fig = go.Figure() |
| fig.add_bar(name="this reproduction", x=names, y=ours, marker_color=ACCENT, |
| text=[f"{v:+.1f}%" for v in ours], textposition="outside", |
| textfont=dict(size=32)) |
| fig.add_bar(name="paper (Table 4)", x=names, y=[paper[n] for n in names], |
| marker_color=GOLD, text=[f"{paper[n]:+.1f}%" for n in names], |
| textposition="outside", textfont=dict(size=32)) |
| lo = min(ours + [paper[n] for n in names]) |
| hi = max(ours + [paper[n] for n in names]) |
| fig.add_hline(y=0, line=dict(color="#555", width=2)) |
| fig.update_layout(barmode="group", |
| yaxis_title="relative F1 gain from pre-training (%)", |
| yaxis_range=[min(lo * 1.5, -3), hi * 1.35], |
| height=620, width=1000, |
| legend=dict(orientation="h", y=1.13, x=0)) |
| save(fig, "claim4_promotion", 1000, 620) |
|
|
|
|
| |
| def fig_sweep(): |
| p = "results/sweep_summary.json" |
| if not os.path.exists(p): |
| print("skip sweep fig (not run yet)") |
| return |
| d = json.load(open(p)) |
| ks, rs = [6, 8, 10], [0.5, 0.6, 0.7, 0.8] |
| ours = np.full((3, 4), np.nan) |
| for row in d["rows"]: |
| ours[ks.index(row["k"]), rs.index(row["r"])] = row["f1"] |
| paper = np.array([[54.81, 56.68, 55.41, 54.32], |
| [62.60, 63.51, 62.10, 59.19], |
| [51.76, 50.70, 56.33, 55.70]]) |
| fig = make_subplots(rows=1, cols=2, horizontal_spacing=0.14, |
| subplot_titles=("this reproduction (PTB-XL, 1 seed)", |
| "paper (Table 3)")) |
| for j, (m, lbl) in enumerate(((ours, "ours"), (paper, "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, ACCENT]], showscale=False, |
| text=[[("" if np.isnan(v) else f"{v:.2f}") for v in row] |
| for row in m], |
| texttemplate="%{text}", textfont=dict(size=30), |
| row=1, col=j + 1) |
| fig.update_layout(height=560, width=1300, |
| title_text="macro F1 (%) across bottleneck size k and mask ratio r") |
| save(fig, "claim6_sweep", 1300, 560) |
|
|
|
|
| |
| def fig_theorem(): |
| rng = np.random.default_rng(0) |
| LOG2PI = np.log(2 * np.pi) |
| n, dd, s2, b = 20000, 6, 0.7, 0.6 |
| x = rng.normal(size=(n, dd)) |
| mses, gll, lll = [], [], [] |
| for s in np.linspace(0.05, 2.0, 25): |
| xh = x + s * rng.normal(size=(n, dd)) |
| mses.append(float(((x - xh) ** 2).mean())) |
| sq = ((x - xh) ** 2).sum(1) |
| gll.append(float((-0.5 * dd * (LOG2PI + np.log(s2)) - sq / (2 * s2)).mean())) |
| lll.append(float((-dd * np.log(2 * b) - np.abs(x - xh).sum(1) / b).mean())) |
| mses, gll, lll = np.array(mses), np.array(gll), np.array(lll) |
|
|
| def resid(xv, yv): |
| A = np.vstack([xv, np.ones_like(xv)]).T |
| c, *_ = np.linalg.lstsq(A, yv, rcond=None) |
| return np.abs(A @ c - yv) |
|
|
| fig = go.Figure() |
| fig.add_scatter(x=mses, y=np.maximum(resid(mses, gll), 1e-16), mode="lines+markers", |
| name="Gaussian decoder (theorem's assumption)", |
| line=dict(color=ACCENT, width=4), marker=dict(size=11)) |
| fig.add_scatter(x=mses, y=resid(mses, lll), mode="lines+markers", |
| name="Laplace decoder — CONTROL", |
| line=dict(color=RED, width=4, dash="dash"), marker=dict(size=11)) |
| fig.update_yaxes(type="log", title_text="|residual| from affine fit [nats]", |
| exponentformat="power") |
| fig.update_xaxes(title_text="reconstruction MSE (ℒ_rec)") |
| fig.update_layout(height=560, width=1000, |
| legend=dict(orientation="h", y=1.16, x=0)) |
| save(fig, "claim2_control", 1000, 560) |
|
|
|
|
| |
| def fig_rank(): |
| d = json.load(open("results/table1_rank_audit.json")) |
| ranks = d["avg_rank_all_methods_5metrics"] |
| items = sorted(ranks.items(), key=lambda kv: kv[1]) |
| names = [k if k != "Ours" else "TS-Fingerprint" for k, _ in items] |
| vals = [v for _, v in items] |
| cols = [GOLD if n == "TS-Fingerprint" else |
| (ACCENT if n == "Medformer" else GREY) for n in names] |
| fig = go.Figure() |
| fig.add_bar(x=vals, y=names, orientation="h", marker_color=cols, |
| text=[f"{v:.2f}" for v in vals], textposition="outside", |
| textfont=dict(size=28)) |
| fig.update_layout(xaxis_title="average rank over 7 datasets × 5 metrics (lower better)", |
| height=760, width=1000, yaxis=dict(autorange="reversed"), |
| xaxis_range=[0, 12]) |
| save(fig, "claim5_rank", 1000, 760) |
|
|
|
|
| if __name__ == "__main__": |
| for f in (fig_ablation, fig_promotion, fig_sweep, fig_theorem, fig_rank): |
| try: |
| f() |
| except Exception as e: |
| print(f"skip {f.__name__}: {type(e).__name__}: {e}") |
|
|