File size: 10,884 Bytes
ef53368 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 | """Broken-axis accuracy-vs-cost Pareto panel used by Figure 2A and SI Figure S1.
Both figures draw the same panel (fitting RMSE vs total compute time, colored by method type, shaped by
geometry source, sized by basis) and differ only in their config dicts and which panels they show, so
that shared drawing engine lives here instead of being duplicated in each notebook. Call
`plot_pareto_panel(...)`; the notebook supplies the points table and the styling dicts.
"""
import numpy as np
import seaborn as sns
import matplotlib.pyplot as plt
from matplotlib.lines import Line2D
from matplotlib.gridspec import GridSpec
# methods colored as double hybrids; dlpno_mp2 is classified ab initio instead (checked first)
DOUBLE_HYBRID_METHODS = ["dsd_pbep86", "B2GP_PLYP", "B2PLYP", "mPW2PLYP", "revdsd_pbep86", "dlpno_mp2"]
def _format_seconds(seconds):
if not np.isfinite(seconds) or seconds <= 0:
return ""
nice = lambda v: str(int(round(v)))
if seconds < 60:
return f"{nice(seconds)} s"
if seconds < 3600:
return f"{nice(seconds / 60)} min"
if seconds < 86400:
return f"{nice(seconds / 3600)} h"
return f"{nice(seconds / 86400)} day"
def _apply_time_xticklabels(ax, xticks):
ax.set_xticks(xticks)
ax.set_xticklabels([_format_seconds(10 ** float(t)) for t in xticks], fontsize=10)
def _classify_method(method):
# returns (label_text, category); category drives color, label_text is the annotation
if method in ["hf", "mp2", "dlpno_mp2"]:
return method, "ab initio"
if method in DOUBLE_HYBRID_METHODS:
return method.lower(), "double hybrids"
if method.startswith("MagNET"):
return "MagNET-Zero", "MagNET-Zero"
if method.startswith(("wp", "wc")):
return method, "NMR-specific"
return method.split("_")[0], "DFT"
def _classify_and_plot_points(df, x_range1, ax1, ax2, color_map, geometry_marker_map,
marker_alpha, manual_label_positions, labels_per_category=6):
size_dict = {"pcSseg1": 20, "pcSseg2": 50, "pcSseg3": 100, "N/A": 110}
color_counts = {k: 0 for k in ["ab initio", "double hybrids", "MagNET-Zero", "NMR-specific", "DFT"]}
basis_counts = {k: 0 for k in ["pcSseg1", "pcSseg2", "pcSseg3"]}
geometry_counts = {}
x_all = np.log10(df["total_time"].astype(float).values)
x_leftmost = np.nanmin(x_all) if len(x_all) else np.nan
# label selection: the min-RMSE point per category plus a few spread across time
tmp = df.copy()
tmp["_log10_time"] = np.log10(tmp["total_time"].astype(float))
tmp["_rmse"] = tmp["fitting_RMSE"].astype(float)
tmp["_category"] = tmp["nmr_method"].astype(str).apply(lambda m: _classify_method(m)[1])
label_indices = set()
for cat in ["MagNET-Zero", "ab initio", "double hybrids", "NMR-specific", "DFT"]:
sub = tmp[tmp["_category"] == cat]
if sub.empty:
continue
picks = [int(sub["_rmse"].idxmin())]
n_extra = max(0, labels_per_category - 1)
if n_extra > 0 and len(sub) > 1:
tvals = sub["_log10_time"].to_numpy()
for q in np.linspace(0.05, 0.95, num=n_extra):
picks.append(int((sub["_log10_time"] - np.quantile(tvals, q)).abs().idxmin()))
seen, unique = set(), []
for p in picks:
if p not in seen:
unique.append(p); seen.add(p)
label_indices.update(unique[:labels_per_category])
for idx, row in df.iterrows():
x_val = float(np.log10(float(row.total_time)))
y_val = float(row["fitting_RMSE"])
geometry_type = str(row["geometry_type"]); method_raw = str(row["nmr_method"]); basis = str(row["basis"])
label_text, category = _classify_method(method_raw)
color_counts[category] += 1
point_color = color_map[category]
marker = geometry_marker_map.get(geometry_type, "o")
geometry_counts[geometry_type] = geometry_counts.get(geometry_type, 0) + 1
if basis in basis_counts:
basis_counts[basis] += 1
base_size = size_dict.get(basis, size_dict.get("N/A", 30))
size = base_size * (2.8 if abs(x_val - x_leftmost) <= 1e-12 else 1.0)
ax = ax1 if (x_range1[0] <= x_val <= x_range1[1]) else ax2
unfilled = marker in ["x", "+", "1", "2", "3", "4", "|", "_"]
kws = dict(x=[x_val], y=[y_val], ax=ax, color=point_color, s=size, marker=marker,
legend=False, alpha=marker_alpha, zorder=3)
kws.update(dict(linewidth=1.2) if unfilled else dict(edgecolor="white", linewidth=0.6))
sns.scatterplot(**kws)
override = None
if manual_label_positions:
for key in [label_text, method_raw, (method_raw, basis),
(method_raw, basis, geometry_type), (label_text, basis, geometry_type)]:
if key in manual_label_positions:
override = manual_label_positions[key]; break
if (idx in label_indices) or (override is not None):
xytext = tuple(override["xytext"]) if isinstance(override, dict) else (4, 3)
ax.annotate(label_text, xy=(x_val, y_val), xytext=xytext, textcoords="offset points",
ha="left", va="bottom", fontsize=8, color=point_color, zorder=4,
annotation_clip=False)
return color_counts, basis_counts, geometry_counts
def _add_legends(ax1, ax2, color_counts, basis_counts, geometry_counts, color_map,
geometry_marker_map, method_legend_label_map, geometry_legend_label_map):
method_order = ["MagNET-Zero", "ab initio", "double hybrids", "NMR-specific", "DFT"]
method_elems = [Line2D([0], [0], marker="o", color="w", markerfacecolor=color_map[c],
markeredgecolor="white", markeredgewidth=0.6, markersize=8, linestyle="None",
label=method_legend_label_map.get(c, c))
for c in method_order if color_counts.get(c, 0) > 0]
basis_elems = [Line2D([0], [0], marker="o", color="none", markerfacecolor="k", markersize=ms, label=b)
for b, ms in [("pcSseg1", 3), ("pcSseg2", 6), ("pcSseg3", 8)] if basis_counts.get(b, 0) > 0]
geometry_elems = []
for g in ["aimnet2", "pbe0_tz"]:
if geometry_counts.get(g, 0) > 0:
gm = geometry_marker_map.get(g, "o")
geometry_elems.append(Line2D([0], [0], marker=gm, color="k",
markerfacecolor="none" if gm == "x" else "k",
markeredgecolor="k", markersize=7, linestyle="None",
label=geometry_legend_label_map.get(g, g)))
ax1.add_artist(ax1.legend(handles=method_elems, loc="upper left", bbox_to_anchor=(0.005, 0.995),
frameon=True, fontsize=9, title="Method Type", title_fontsize=9,
borderaxespad=0.0, handletextpad=0.4, labelspacing=0.3))
ax2.add_artist(ax2.legend(handles=basis_elems, loc="upper right", bbox_to_anchor=(0.995, 0.995),
frameon=True, fontsize=9, title="Basis Set", title_fontsize=9,
borderaxespad=0.0, handletextpad=0.4, labelspacing=0.3))
if geometry_elems:
ax2.add_artist(ax2.legend(handles=geometry_elems, loc="upper right", bbox_to_anchor=(0.80, 0.995),
frameon=True, fontsize=9, title="Geometries", title_fontsize=9,
borderaxespad=0.0, handletextpad=0.4, labelspacing=0.3))
def plot_pareto_panel(pareto_df, query_str, nucleus, method_color_map, method_legend_label_map,
geometry_legend_label_map, geometry_marker_map, xlim_left, xlim_right, ylim,
figsize, marker_alpha, manual_label_positions, save_png=None,
suptitle="MagNET-Zero Extends a Flat Pareto Frontier",
supxlabel="Total delta-22 calculation time"):
"""Draw the broken-axis accuracy-vs-cost Pareto panel and optionally save it as a PNG."""
sns.set_theme(style="ticks", context="paper", rc={"axes.grid": False})
query_df = pareto_df.query(query_str.format(nucleus=nucleus), engine="python").reset_index(drop=True)
default_color_map = {"ab initio": "black", "double hybrids": "#ede6ba", "MagNET-Zero": "#61a89a",
"NMR-specific": "#a72608", "DFT": "black"}
color_map = {**default_color_map, **method_color_map}
fig = plt.figure(figsize=figsize)
gs = GridSpec(1, 2, width_ratios=[1, 5], wspace=0.05)
ax1 = fig.add_subplot(gs[0]); ax2 = fig.add_subplot(gs[1], sharey=ax1)
cc, bc, gc = _classify_and_plot_points(query_df, xlim_left, ax1, ax2, color_map,
geometry_marker_map, marker_alpha, manual_label_positions)
ax1.set_xlim(xlim_left); ax2.set_xlim(xlim_right)
ax1.set_ylim(ylim); ax2.set_ylim(ylim)
y_ticks = np.linspace(ylim[0], ylim[1], num=6)
ax1.set_yticks(y_ticks); ax1.set_yticklabels([f"{t:.2f}" for t in y_ticks], fontsize=10)
_apply_time_xticklabels(ax1, np.array([xlim_left[0], (xlim_left[0] + xlim_left[1]) / 2, xlim_left[1]]))
_apply_time_xticklabels(ax2, np.arange(xlim_right[0] + 0.25, xlim_right[1] + 0.25, 0.25))
ax1.tick_params(axis="both", which="both", length=1.5, width=0.6, labelsize=10)
ax2.tick_params(axis="x", which="both", length=1.5, width=0.6, labelsize=10)
ax2.tick_params(axis="y", which="both", left=False, right=False, labelleft=False)
ax1.grid(False); ax2.grid(False)
for ax in (ax1, ax2):
for side in ("left", "right", "top", "bottom"):
ax.spines[side].set_visible(True); ax.spines[side].set_linewidth(0.6)
ax1.spines["right"].set_visible(False); ax2.spines["left"].set_visible(False)
# broken-axis break marks
width_ratio = ax2.get_position().width / ax1.get_position().width
d1 = 0.01; d2 = d1 / width_ratio; slope2 = width_ratio
kw = dict(color="k", clip_on=False, lw=0.6, transform=ax1.transAxes)
ax1.plot((1 - d1, 1 + d1), (-d1, d1), **kw); ax1.plot((1 - d1, 1 + d1), (1 - d1, 1 + d1), **kw)
kw.update(transform=ax2.transAxes)
ax2.plot((-d2, d2), (-d2 * slope2, d2 * slope2), **kw); ax2.plot((-d2, d2), (1 - d2 * slope2, 1 + d2 * slope2), **kw)
fig.subplots_adjust(left=0.15, right=0.95, top=0.92, bottom=0.12)
fig.suptitle(suptitle, fontsize=14)
fig.supxlabel(supxlabel, fontsize=13)
ax1.set_ylabel(r"RMSE ($^{1}$H ppm, CDCl$_3$)" if nucleus == "H" else r"RMSE ($^{13}$C ppm, CDCl$_3$)",
fontsize=13)
_add_legends(ax1, ax2, cc, bc, gc, color_map, geometry_marker_map,
method_legend_label_map, geometry_legend_label_map)
if save_png:
fig.savefig(save_png, dpi=300, bbox_inches="tight")
plt.show()
|