spark / plotting.py
bingyan user
Cursor
Frontend redesign: branded hero header, Soft theme + CSS cards, decluttered two-column layout, polished plots, grouped examples
827f36a
Raw History Blame Contribute Delete
6.46 kB
"""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)