"""Matplotlib figures for the SPARK v4 outer-sphere ET demo. All plots take the dicts produced by inference.SPARKPredictor.predict and return a matplotlib Figure for Gradio to render. """ import matplotlib matplotlib.use("Agg") import numpy as np import matplotlib.pyplot as plt plt.rcParams.update({ "font.size": 13, "axes.titlesize": 13, "axes.labelsize": 13, "xtick.labelsize": 11, "ytick.labelsize": 11, "legend.fontsize": 11, "axes.linewidth": 0.9, "axes.edgecolor": "#94a3b8", "figure.autolayout": False, "axes.grid": True, "grid.color": "#e2e8f0", "grid.linewidth": 0.8, "font.family": "sans-serif", "savefig.facecolor": "white", "figure.facecolor": "white", }) # palette matched to the app theme (indigo / emerald / amber / rose) _INDIGO = "#4f46e5" _EMERALD = "#059669" _AMBER = "#d97706" _ROSE = "#e11d48" _GREEN = _EMERALD _GREY = "#cbd5e1" _BLUE = _INDIGO _RED = _ROSE def plot_mechanism_posterior(mechanisms, tol=None): """Horizontal bar chart of the categorical posterior over mechanisms.""" fig, ax = plt.subplots(figsize=(6.6, 4.3), constrained_layout=True) labels = [m["label"] for m in mechanisms][::-1] probs = [m["prob"] for m in mechanisms][::-1] within = [m.get("within_tol") for m in mechanisms][::-1] colors = [_EMERALD if w else (_GREY if w is not None else _INDIGO) for w in within] y = np.arange(len(labels)) ax.barh(y, probs, color=colors) for yi, p in zip(y, probs): ax.text(p + 0.012, yi, f"{p*100:.0f}%", va="center", fontsize=11, color="#334155") ax.set_yticks(y); ax.set_yticklabels(labels, fontsize=11) ax.set_xlim(0, min(1.0, max(probs) * 1.28 + 0.05)) ax.set_xlabel("posterior probability") ax.grid(axis="y", visible=False) return fig def plot_alternatives(alternatives): """Observed vs posterior-predictive reconstruction for the top alternative mechanisms (slowest scan shown, which is the most mechanism-discriminating).""" cands = [c for c in alternatives if c.get("panels")] if not cands: fig, ax = plt.subplots(figsize=(6, 3)) ax.text(0.5, 0.5, "Reconstruction unavailable", ha="center", va="center") ax.axis("off") return fig k = min(len(cands), 3) fig, axes = plt.subplots(1, k, figsize=(4.6 * k, 4.0), squeeze=False) cmap = [_BLUE, _RED, "#9467bd"] for j, c in enumerate(cands[:k]): ax = axes[0][j] sg, pot, obs, rec = c["panels"][0] # slowest scan ax.plot(pot, obs, "k", lw=1.6, label="observed") ax.plot(pot, rec, "--", color=cmap[j % 3], lw=1.5, label="posterior-pred.") ax.invert_xaxis() ax.set_xlabel("dimensionless potential") if j == 0: ax.set_ylabel("dimensionless current") badge = " \u2713 within tol." if c.get("within_tol") else "" ax.set_title(f"{c['label']}\np={c['prob']*100:.0f}%, NRMSE={c['nrmse']:.3f}{badge}", fontsize=9.5) ax.legend(fontsize=8, frameon=False) fig.suptitle("Alternative mechanisms consistent with this voltammogram", fontsize=12, fontweight="bold") fig.tight_layout(rect=[0, 0, 1, 0.92]) return fig def plot_presence(presence, gate_pretty): """Bar chart of per-elementary-step presence probability.""" items = sorted(presence.items(), key=lambda kv: -kv[1]) items = [it for it in items if it[1] > 0.02][:10] if not items: items = sorted(presence.items(), key=lambda kv: -kv[1])[:6] labels = [gate_pretty.get(g, g.replace("log10_", "")) for g, _ in items][::-1] probs = [p for _, p in items][::-1] fig, ax = plt.subplots(figsize=(6.6, 4.3), constrained_layout=True) y = np.arange(len(labels)) ax.barh(y, probs, color=[_EMERALD if p > 0.5 else _GREY for p in probs]) ax.axvline(0.5, color="#64748b", ls="--", lw=1.0) ax.set_yticks(y); ax.set_yticklabels(labels, fontsize=11) ax.set_xlim(0, 1); ax.set_xlabel("presence probability P(step active)") ax.grid(axis="y", visible=False) return fig def plot_parameter_posteriors(samples, params, slot_idx): """Histogram grid of key parameter posteriors with 90% CI shading.""" names = list(params.keys()) k = len(names) ncol = min(3, k) nrow = int(np.ceil(k / ncol)) fig, axes = plt.subplots(nrow, ncol, figsize=(3.2 * ncol, 2.3 * nrow), squeeze=False, constrained_layout=True) for a in axes.ravel()[k:]: a.axis("off") for i, nm in enumerate(names): ax = axes[i // ncol][i % ncol] j = slot_idx[nm] col = samples[:, j] ax.hist(col, bins=40, color=_INDIGO, alpha=0.55, density=True) lo, hi = params[nm]["lo"], params[nm]["hi"] ax.axvspan(lo, hi, color=_INDIGO, alpha=0.12) ax.axvline(params[nm]["median"], color=_ROSE, lw=1.8) ax.set_yticks([]); ax.grid(visible=False) ax.set_title(params[nm]["label"], fontsize=9.5) fig.suptitle("Parameter posteriors (rose = median, band = 90% CI)", fontsize=11.5) return fig def plot_reconstruction(panels, label, nrmse): """Observed vs posterior-predictive reconstruction across scan rates for one mechanism.""" if not panels: fig, ax = plt.subplots(figsize=(6, 3)) ax.text(0.5, 0.5, "Reconstruction unavailable", ha="center", va="center") ax.axis("off") return fig k = len(panels) fig, axes = plt.subplots(1, k, figsize=(3.4 * k, 3.6), squeeze=False, constrained_layout=True) for j, (sg, pot, obs, rec) in enumerate(panels): ax = axes[0][j] ax.plot(pot, obs, color="#0f172a", lw=2.0, label="observed") ax.plot(pot, rec, "--", color=_EMERALD, lw=2.0, label="reconstruction") ax.invert_xaxis() ax.set_xlabel("dimensionless potential") if j == 0: ax.set_ylabel("dimensionless current") ax.set_title(f"$\\sigma$={sg:.2g}", fontsize=11) if j == 0: ax.legend(frameon=False, loc="best") fig.suptitle(f"{label} (median NRMSE {nrmse:.3f})", fontsize=12, fontweight="bold") return fig def parameter_table(params): """Markdown table of parameter posterior summaries for the UI.""" rows = ["| Parameter | Median | 90% CI |", "|---|---|---|"] for nm, d in params.items(): rows.append(f"| {d['label']} | {d['median']:+.2f} | [{d['lo']:+.2f}, {d['hi']:+.2f}] |") return "\n".join(rows)