"""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)") # ------------------------------------------------------------------ Fig 1: Claim 3 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
(ℒ_rec)", "Pre-trained
(ℒ_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) # ------------------------------------------------------------------ Fig 2: Claim 4 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) # ------------------------------------------------------------------ Fig 3: Claim 6 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) # ------------------------------------------------------------------ Fig 4: Claim 2 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) # ------------------------------------------------------------------ Fig 5: Claim 5 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: # noqa: BLE001 print(f"skip {f.__name__}: {type(e).__name__}: {e}")