From 0aa1be5dc8c2712ef8fed3b87bb0d7d30897cffd Mon Sep 17 00:00:00 2001 From: Mateusz Praski Date: Thu, 17 Sep 2026 12:18:23 +0200 Subject: [PATCH 1/2] Added strong and weak posterior plots --- btk/tests/bbt/_types.py | 4 + btk/tests/bbt/bbt.py | 270 ++++++++++++++++++- btk/tests/bbt/plots/__init__.py | 4 +- btk/tests/bbt/plots/_strong_posterior.py | 184 +++++++++++++ btk/tests/bbt/plots/_weak_posterior.py | 246 ++++++++++++++++++ tests/tests/bbt/test_bbt.py | 315 +++++++++++++++++++++++ 6 files changed, 1021 insertions(+), 2 deletions(-) create mode 100644 btk/tests/bbt/plots/_strong_posterior.py create mode 100644 btk/tests/bbt/plots/_weak_posterior.py diff --git a/btk/tests/bbt/_types.py b/btk/tests/bbt/_types.py index d04abf7..7bb35a0 100644 --- a/btk/tests/bbt/_types.py +++ b/btk/tests/bbt/_types.py @@ -24,6 +24,10 @@ "strong_interpretation_raw", ] +# Figures dispatched through ``BBTTest.plot(kind=...)``. +PlotKindType = Literal["strong-posterior", "weak-posterior"] +PlotOrientationType = Literal["horizontal", "vertical"] + InterpretationTypes = Literal[ "weak", "strong", diff --git a/btk/tests/bbt/bbt.py b/btk/tests/bbt/bbt.py index 5ac6f79..c42aadc 100644 --- a/btk/tests/bbt/bbt.py +++ b/btk/tests/bbt/bbt.py @@ -12,12 +12,14 @@ ALL_PROPERTIES_COLUMNS, HyperPriorType, InterpretationTypes, + PlotKindType, + PlotOrientationType, ReportedPropertyColumnType, TieSolverType, ) from .alg import _construct_win_table, _get_pwin, _hdi from .model import _mcmcbbt_pymc -from .plots import plot_cdd_diagram +from .plots import plot_cdd_diagram, plot_strong_posterior, plot_weak_posterior class BBTTest(BaseBayesianTest): @@ -700,6 +702,272 @@ def rope_comparison_control_table( return pd.DataFrame.from_records(records) + @_validate_params + def plot( + self, + kind: PlotKindType = "strong-posterior", + control_model: str | None = None, + selected_pairs: Sequence[tuple[str, str]] | None = None, + selected_models: Iterable[str] | None = None, + hdi_prob: float = 0.89, + rope_value: tuple[float, float] = (0.45, 0.55), + orientation: PlotOrientationType = "horizontal", + ax: plt.Axes | Sequence[plt.Axes] | None = None, + **kwargs, + ) -> plt.Axes | np.ndarray: + """Plot the posterior of the fitted BBT model; shorthand for the ``plot_*`` methods. + + Arguments the chosen ``kind`` does not use are ignored, so one call + signature serves every kind. See the dedicated methods for what each + figure shows and what its arguments mean. + + Parameters + ---------- + kind : str, optional + The figure to draw. Defaults to `strong-posterior`. + + - `strong-posterior` - :meth:`plot_strong_posterior`; uses + ``hdi_prob``, ignores ``rope_value``. + - `weak-posterior` - :meth:`plot_weak_posterior`; uses + ``rope_value``, ignores ``hdi_prob``. + + control_model : str | None, optional + Compare every other model against this one. Defaults to None. + selected_pairs : Sequence[tuple[str, str]] | None, optional + Pairs to draw, each read in the given order. Defaults to None. + selected_models : Iterable[str] | None, optional + With ``control_model``, the subset of models to compare against it. + hdi_prob : float, optional + Probability mass of the HDIs. Defaults to 0.89. + rope_value : tuple[float, float], optional + Region of Practical Equivalence (ROPE). Defaults to (0.45, 0.55). + orientation : str, optional + `horizontal` or `vertical`. Defaults to `horizontal`. + ax : plt.Axes | Sequence[plt.Axes] | None, optional + One Axes for `strong-posterior`, two for `weak-posterior`. If None, + a new figure is created. + **kwargs + Additional keyword arguments passed to the underlying plotting function. + + Returns + ------- + plt.Axes | np.ndarray + A single Axes for `strong-posterior`, an array of two for `weak-posterior`. + + See Also + -------- + plot_strong_posterior : Posterior mean and HDI under the strong interpretation. + plot_weak_posterior : P(pi > 0.5) and P(pi in ROPE) under the weak interpretation. + """ + selection = { + "control_model": control_model, + "selected_pairs": selected_pairs, + "selected_models": selected_models, + "orientation": orientation, + "ax": ax, + } + if kind == "strong-posterior": + return self.plot_strong_posterior(**selection, hdi_prob=hdi_prob, **kwargs) + if kind == "weak-posterior": + return self.plot_weak_posterior( + **selection, rope_value=rope_value, **kwargs + ) + raise ValueError(f"Unsupported plot kind {kind!r}.") + + @_validate_params + def plot_strong_posterior( + self, + control_model: str | None = None, + selected_pairs: Sequence[tuple[str, str]] | None = None, + selected_models: Iterable[str] | None = None, + hdi_prob: float = 0.89, + orientation: PlotOrientationType = "horizontal", + ax: plt.Axes | None = None, + **kwargs, + ) -> plt.Axes: + r"""Plot each pairwise probability under the strong interpretation. + + Draws the posterior mean of :math:`\pi` with its HDI for each comparison, + coloured by the strong interpretation (Wainer 2023, sec. 8.3): `better` + if :math:`E[\pi] > 0.70`, `equivalent` if :math:`0.45 \leq E[\pi] \leq 0.55`, + `weaker` if :math:`E[\pi] < 0.30`, `no claim` otherwise. Shaded regions + mark the three claims. + + Exactly one of ``control_model`` or ``selected_pairs`` must be given: with + many models the full set of pairs is too large to read. + + Parameters + ---------- + control_model : str | None, optional + Compare every other model against this one, each read as + ``P(model > control_model)``. Defaults to None. + selected_pairs : Sequence[tuple[str, str]] | None, optional + Pairs to draw, each read in the given order, so ``("a", "b")`` is + ``P(a > b)``. Defaults to None. + selected_models : Iterable[str] | None, optional + With ``control_model``, the subset of models to compare against it. + Defaults to all fitted models. + hdi_prob : float, optional + Probability mass of the HDIs. Defaults to 0.89. + orientation : str, optional + `horizontal` lays the comparisons along the x axis, best on the left; + `vertical` lays them along the y axis, best on top. Defaults to + `horizontal`. + ax : plt.Axes | None, optional + Matplotlib Axes to plot on. If None, a new figure and axes will be created. + **kwargs + Additional keyword arguments passed to the ``scatter`` call of the means. + + Returns + ------- + plt.Axes + The Axes the figure was drawn on. + + See Also + -------- + plot_weak_posterior : The same comparisons under the weak interpretation. + """ + samples, labels, value_label, subtitle = self._plot_samples( + control_model, selected_pairs, selected_models + ) + hdi_values = _hdi(samples, hdi_prob) + return plot_strong_posterior( + labels=labels, + means=samples.mean(axis=0), + hdi_low=hdi_values[0], + hdi_high=hdi_values[1], + better_threshold=self._STRONG_INTERPRETATION_BETTER_THRESHOLD, + equal_threshold=self._STRONG_INTERPRETATION_EQUAL_THRESHOLD, + hdi_prob=hdi_prob, + value_label=value_label, + orientation=orientation, + subtitle=subtitle, + ax=ax, + **kwargs, + ) + + @_validate_params + def plot_weak_posterior( + self, + control_model: str | None = None, + selected_pairs: Sequence[tuple[str, str]] | None = None, + selected_models: Iterable[str] | None = None, + rope_value: tuple[float, float] = (0.45, 0.55), + orientation: PlotOrientationType = "horizontal", + ax: Sequence[plt.Axes] | None = None, + **kwargs, + ) -> np.ndarray: + r"""Plot each pairwise probability under the weak interpretation. + + Draws two panels per comparison: :math:`P(\pi > 0.5)` and + :math:`P(\pi \in \mathrm{ROPE})`, coloured by the weak interpretation + (Wainer 2023, sec. 8.2): `equivalent` if + :math:`P(\pi \in \mathrm{ROPE}) \geq 0.95`, otherwise `better` if + :math:`P(\pi > 0.5) \geq 0.95`, `weaker` if :math:`P(\pi < 0.5) \geq 0.95`, + `no claim` otherwise. Comparisons are ordered by :math:`E[\pi]`. + + Exactly one of ``control_model`` or ``selected_pairs`` must be given: with + many models the full set of pairs is too large to read. + + Parameters + ---------- + control_model : str | None, optional + Compare every other model against this one, each read as + ``P(model > control_model)``. Defaults to None. + selected_pairs : Sequence[tuple[str, str]] | None, optional + Pairs to draw, each read in the given order, so ``("a", "b")`` is + ``P(a > b)``. Defaults to None. + selected_models : Iterable[str] | None, optional + With ``control_model``, the subset of models to compare against it. + Defaults to all fitted models. + rope_value : tuple[float, float], optional + Region of Practical Equivalence (ROPE). Defaults to (0.45, 0.55). + orientation : str, optional + `horizontal` stacks the panels and lays the comparisons along the + x axis, best on the left; `vertical` puts the panels side by side + and lays the comparisons along the y axis, best on top. Defaults to + `horizontal`. + ax : Sequence[plt.Axes] | None, optional + Exactly two Axes: ``ax[0]`` for :math:`P(\pi > 0.5)`, ``ax[1]`` for + :math:`P(\pi \in \mathrm{ROPE})`. If None, a new figure is created. + **kwargs + Additional keyword arguments passed to both ``scatter`` calls. + + Returns + ------- + np.ndarray + The two Axes drawn on, in the order described for ``ax``. + + See Also + -------- + plot_strong_posterior : The same comparisons under the strong interpretation. + """ + samples, labels, value_label, subtitle = self._plot_samples( + control_model, selected_pairs, selected_models + ) + return plot_weak_posterior( + labels=labels, + means=samples.mean(axis=0), + above_50=np.mean(samples > 0.5, axis=0), + below_50=np.mean(samples < 0.5, axis=0), + in_rope=np.mean( + (samples >= rope_value[0]) & (samples <= rope_value[1]), axis=0 + ), + threshold=self._WEAK_INTERPRETATION_THRESHOLD, + rope_value=rope_value, + value_label=value_label, + orientation=orientation, + subtitle=subtitle, + ax=ax, + **kwargs, + ) + + def _plot_samples( + self, + control_model: str | None, + selected_pairs: Sequence[tuple[str, str]] | None, + selected_models: Iterable[str] | None, + ) -> tuple[np.ndarray, list[str], str, str | None]: + """Posterior draws of ``pi`` per plotted comparison, with labels. + + Returns the ``(draws, comparisons)`` sample matrix, one label per + comparison, the axis label defining ``pi`` and an optional subtitle. + """ + self._check_if_fitted() + if (control_model is None) == (selected_pairs is None): + raise ValueError("Pass exactly one of control_model or selected_pairs.") + if selected_pairs is not None and selected_models is not None: + raise ValueError( + "selected_models only applies with control_model; " + "list the pairs in selected_pairs instead." + ) + if selected_pairs is not None: + pairs = list(selected_pairs) + if not pairs: + raise ValueError("selected_pairs must contain at least one pair.") + draws = self.pairwise_samples(pairs) + value_label = r"$\pi = P(\mathrm{left} \succ \mathrm{right})$" + return draws.to_numpy(), list(draws.columns), value_label, None + + samples, names = _get_pwin( + bbt_result=self._fit_posterior, + alg_names=self._algorithms, + control=control_model, + selected=list(selected_models) if selected_models is not None else None, + ) + # ``_get_pwin`` puts the better model on the left; re-orient every pair + # as ``other > control`` so the control is the fixed reference. + labels = [] + for k, name in enumerate(names): + left, right = (part.strip() for part in name.split(">")) + if left == control_model: + samples[:, k] = 1.0 - samples[:, k] + left = right + labels.append(left) + value_label = rf"$\pi = P(\mathrm{{model}} \succ$ {control_model}$)$" + # Rows are named by the model alone, so say what they are compared to. + return samples, labels, value_label, f"Control model: {control_model}" + @_validate_params def plot_cdd_diagram( self, diff --git a/btk/tests/bbt/plots/__init__.py b/btk/tests/bbt/plots/__init__.py index e338db3..91f8e93 100644 --- a/btk/tests/bbt/plots/__init__.py +++ b/btk/tests/bbt/plots/__init__.py @@ -1,3 +1,5 @@ from ._critical_difference import plot_cdd_diagram +from ._strong_posterior import plot_strong_posterior +from ._weak_posterior import plot_weak_posterior -__all__ = ["plot_cdd_diagram"] +__all__ = ["plot_cdd_diagram", "plot_strong_posterior", "plot_weak_posterior"] diff --git a/btk/tests/bbt/plots/_strong_posterior.py b/btk/tests/bbt/plots/_strong_posterior.py new file mode 100644 index 0000000..538a8bb --- /dev/null +++ b/btk/tests/bbt/plots/_strong_posterior.py @@ -0,0 +1,184 @@ +"""Forest plot of the pairwise posterior under the strong interpretation (Wainer 2023, sec. 8.3).""" + +from collections.abc import Sequence + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.lines import Line2D +from matplotlib.patches import Patch + +STRONG_VERDICTS = ("better", "equivalent", "no claim", "weaker") +STRONG_COLOURS = { + "better": "#1b9e77", + "equivalent": "#7570b3", + "no claim": "#9e9e9e", + "weaker": "#d95f02", +} + + +def strong_verdicts( + means: np.ndarray, better_threshold: float, equal_threshold: float +) -> np.ndarray: + """Classify posterior means of ``P(left > right)`` under the strong reading. + + The rule is applied symmetrically around 0.5, so a pair oriented with the + weaker model on the left is reported as ``weaker`` rather than ``no claim``. + """ + means = np.asarray(means, dtype=float) + return np.select( + [ + means > better_threshold, + 1.0 - means > better_threshold, + (means <= equal_threshold) & (1.0 - means <= equal_threshold), + ], + ["better", "weaker", "equivalent"], + default="no claim", + ) + + +def plot_strong_posterior( + labels: Sequence[str], + means: np.ndarray, + hdi_low: np.ndarray, + hdi_high: np.ndarray, + better_threshold: float, + equal_threshold: float, + hdi_prob: float, + value_label: str, + orientation: str = "horizontal", + subtitle: str | None = None, + ax: plt.Axes | None = None, + **kwargs, +) -> plt.Axes: + """Draw ``E[pi]`` with its HDI per comparison, coloured by the strong verdict. + + Parameters + ---------- + labels : Sequence[str] + One label per comparison. + means : np.ndarray + Posterior means of ``pi = P(left > right)``. + hdi_low, hdi_high : np.ndarray + HDI bounds of ``pi``. + better_threshold : float + ``E[pi]`` above this is ``better`` (below ``1 - better_threshold`` is ``weaker``). + equal_threshold : float + ``E[pi]`` within ``[1 - equal_threshold, equal_threshold]`` is ``equivalent``. + hdi_prob : float + Mass of the HDI, used in the title. + value_label : str + Label of the probability axis. + orientation : {"horizontal", "vertical"}, default "horizontal" + ``horizontal`` lays the comparisons along the x axis, best on the left; + ``vertical`` lays them along the y axis, best on top. + subtitle : str | None, default None + Line drawn under the title, e.g. naming the control model. + ax : plt.Axes | None, default None + Axes to draw on; created if ``None``. + **kwargs + Extra keyword arguments forwarded to the ``scatter`` call of the means. + + Returns + ------- + plt.Axes + The axes drawn on. + """ + if ax is not None and not isinstance(ax, plt.Axes): + raise ValueError("ax must be a matplotlib Axes object or None.") + if orientation not in ("horizontal", "vertical"): + raise ValueError( + f"orientation must be 'horizontal' or 'vertical', got {orientation!r}." + ) + horizontal = orientation == "horizontal" + + # Horizontal reads best-first left to right; vertical puts the best on top, + # i.e. last along the y axis. + means = np.asarray(means, dtype=float) + order = np.argsort(-means if horizontal else means, kind="stable") + labels = [labels[i] for i in order] + means = means[order] + hdi_low = np.asarray(hdi_low, dtype=float)[order] + hdi_high = np.asarray(hdi_high, dtype=float)[order] + verdicts = strong_verdicts(means, better_threshold, equal_threshold) + n = len(labels) + + if ax is None: + figsize = (0.45 * n + 3.5, 5.5) if horizontal else (9, 0.32 * n + 1.6) + _, ax = plt.subplots(figsize=figsize) + + band = ax.axhspan if horizontal else ax.axvspan + midline = ax.axhline if horizontal else ax.axvline + for lo, hi, verdict in ( + (0.0, 1.0 - better_threshold, "weaker"), + (1.0 - equal_threshold, equal_threshold, "equivalent"), + (better_threshold, 1.0, "better"), + ): + band(lo, hi, color=STRONG_COLOURS[verdict], alpha=0.10, lw=0) + midline(0.5, color="k", ls=":", lw=1.2) + + positions = np.arange(n) + colours = [STRONG_COLOURS[v] for v in verdicts] + scatter_kwargs = {"s": 45, "zorder": 3, "edgecolor": "white", **kwargs} + value_limits = ( + min(1.0 - better_threshold - 0.05, float(hdi_low.min()) - 0.02), + max(better_threshold + 0.05, float(hdi_high.max()) + 0.02), + ) + if horizontal: + ax.vlines(positions, hdi_low, hdi_high, colors=colours, lw=2.2) + ax.scatter(positions, means, c=colours, **scatter_kwargs) + ax.set_xticks(positions, labels, rotation=60, ha="right") + ax.set_xlim(-0.7, n - 0.3) + ax.set_ylim(*value_limits) + ax.grid(axis="x", visible=False) + ax.set_ylabel(value_label) + else: + ax.hlines(positions, hdi_low, hdi_high, colors=colours, lw=2.2) + ax.scatter(means, positions, c=colours, **scatter_kwargs) + ax.set_yticks(positions, labels) + ax.set_ylim(-0.7, n - 0.3) + ax.set_xlim(*value_limits) + ax.grid(axis="y", visible=False) + ax.set_xlabel(value_label) + + _add_title_and_legend(ax, verdicts, hdi_prob, subtitle) + return ax + + +def _add_title_and_legend( + ax: plt.Axes, verdicts: np.ndarray, hdi_prob: float, subtitle: str | None +) -> None: + title = rf"Strong interpretation: $E[\pi]$ with {hdi_prob:.0%} HDI" + if subtitle is None: + ax.set_title(title, loc="left") + else: + ax.set_title(title, loc="left", pad=20) + ax.text( + 0.0, + 1.01, + subtitle, + transform=ax.transAxes, + ha="left", + va="bottom", + fontsize="small", + color="0.35", + ) + + present = set(verdicts.tolist()) + handles = [ + Line2D( + [], [], marker="o", ls="", color=STRONG_COLOURS[v], markersize=8, label=v + ) + for v in STRONG_VERDICTS + if v in present + ] + handles += [ + Patch(color=STRONG_COLOURS[v], alpha=0.25, label=f"{v} region") + for v in ("better", "equivalent", "weaker") + ] + ax.legend( + handles=handles, + loc="upper left", + bbox_to_anchor=(1.01, 1.0), + fontsize=9, + frameon=False, + ) diff --git a/btk/tests/bbt/plots/_weak_posterior.py b/btk/tests/bbt/plots/_weak_posterior.py new file mode 100644 index 0000000..89e3a3b --- /dev/null +++ b/btk/tests/bbt/plots/_weak_posterior.py @@ -0,0 +1,246 @@ +"""Point plots of the pairwise posterior under the weak interpretation (Wainer 2023, sec. 8.2).""" + +from collections.abc import Sequence + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import to_rgb +from matplotlib.lines import Line2D + +from ._strong_posterior import STRONG_COLOURS as WEAK_COLOURS +from ._strong_posterior import STRONG_VERDICTS as WEAK_VERDICTS + + +def weak_verdicts( + above_50: np.ndarray, + below_50: np.ndarray, + in_rope: np.ndarray, + threshold: float, +) -> np.ndarray: + """Classify comparisons of ``pi = P(left > right)`` under the weak reading. + + Equivalence is checked first, as in ``BBTTest.posterior_table``; a pair + oriented with the weaker model on the left is reported as ``weaker``. + """ + return np.select( + [ + np.asarray(in_rope) >= threshold, + np.asarray(above_50) >= threshold, + np.asarray(below_50) >= threshold, + ], + ["equivalent", "better", "weaker"], + default="no claim", + ) + + +def _lighten(colour: str, amount: float = 0.55) -> tuple[float, float, float]: + rgb = np.array(to_rgb(colour)) + return tuple(rgb + (1.0 - rgb) * amount) + + +def plot_weak_posterior( + labels: Sequence[str], + means: np.ndarray, + above_50: np.ndarray, + below_50: np.ndarray, + in_rope: np.ndarray, + threshold: float, + rope_value: tuple[float, float], + value_label: str, + orientation: str = "horizontal", + subtitle: str | None = None, + ax: Sequence[plt.Axes] | None = None, + **kwargs, +) -> np.ndarray: + """Draw ``P(pi > 0.5)`` and ``P(pi in ROPE)`` per comparison in two panels. + + Parameters + ---------- + labels : Sequence[str] + One label per comparison. + means : np.ndarray + Posterior means of ``pi = P(left > right)``, used only to order the comparisons. + above_50, below_50 : np.ndarray + Posterior probabilities that ``pi`` is above / below 0.5. + in_rope : np.ndarray + Posterior probability that ``pi`` lies in ``rope_value``. + threshold : float + Probability a quantity must reach for a claim. + rope_value : tuple[float, float] + The ROPE, shown in the axis label. + value_label : str + Definition of ``pi``, shown in the title. + orientation : {"horizontal", "vertical"}, default "horizontal" + ``horizontal`` stacks the panels and lays the comparisons along the x + axis, best on the left; ``vertical`` puts the panels side by side and + lays the comparisons along the y axis, best on top. + subtitle : str | None, default None + Line drawn under the title, e.g. naming the control model. + ax : Sequence[plt.Axes] | None, default None + Two Axes: ``P(pi > 0.5)`` is drawn on the first, ``P(pi in ROPE)`` on + the second. Created if ``None``. + **kwargs + Extra keyword arguments forwarded to both ``scatter`` calls. + + Returns + ------- + np.ndarray + The two Axes drawn on. + """ + if orientation not in ("horizontal", "vertical"): + raise ValueError( + f"orientation must be 'horizontal' or 'vertical', got {orientation!r}." + ) + horizontal = orientation == "horizontal" + + means = np.asarray(means, dtype=float) + order = np.argsort(-means if horizontal else means, kind="stable") + labels = [labels[i] for i in order] + above_50 = np.asarray(above_50, dtype=float)[order] + below_50 = np.asarray(below_50, dtype=float)[order] + in_rope = np.asarray(in_rope, dtype=float)[order] + verdicts = weak_verdicts(above_50, below_50, in_rope, threshold) + n = len(labels) + + axes = _resolve_axes(ax, n, horizontal) + positions = np.arange(n) + colours = [WEAK_COLOURS[v] for v in verdicts] + scatter_kwargs = {"s": 110, "zorder": 3, "edgecolor": "white", **kwargs} + rope_label = rf"$P(\pi \in$ ROPE$)$, ROPE = [{rope_value[0]}, {rope_value[1]}]" + + bands = (("better", threshold, 1.02), ("weaker", -0.02, 1.0 - threshold)) + _draw_panel( + axes[0], + positions, + above_50, + colours, + labels, + r"$P(\pi > 0.5)$", + bands, + horizontal, + scatter_kwargs, + ) + bands = (("equivalent", threshold, 1.02),) + _draw_panel( + axes[1], + positions, + in_rope, + colours, + labels, + rope_label, + bands, + horizontal, + scatter_kwargs, + ) + + if horizontal: + # Stacked panels share the model axis; only the bottom one names it. + axes[0].tick_params(axis="x", labelbottom=False) + else: + axes[1].tick_params(axis="y", labelleft=False) + + _add_title_and_legend(axes, verdicts, threshold, value_label, subtitle, horizontal) + return axes + + +def _resolve_axes( + ax: Sequence[plt.Axes] | None, n: int, horizontal: bool +) -> np.ndarray: + if ax is None: + if horizontal: + _, axes = plt.subplots(2, 1, figsize=(0.42 * n + 3.5, 8.5), sharex=True) + else: + _, axes = plt.subplots(1, 2, figsize=(11, 0.32 * n + 1.6), sharey=True) + return np.asarray(axes, dtype=object) + axes = ( + np.array([ax], dtype=object) + if isinstance(ax, plt.Axes) + else np.asarray(ax, dtype=object).ravel() + ) + if len(axes) != 2 or not all(isinstance(a, plt.Axes) for a in axes): + raise ValueError( + "The weak-posterior plot draws two panels and needs two Axes: " + "ax[0] for P(pi > 0.5) and ax[1] for P(pi in ROPE), " + f"got {len(axes)} object(s). Pass ax=None to create them, or e.g. " + "fig, ax = plt.subplots(2, 1, sharex=True)." + ) + return axes + + +def _draw_panel( + panel: plt.Axes, + positions: np.ndarray, + values: np.ndarray, + colours: list[str], + labels: list[str], + value_label: str, + bands: Sequence[tuple[str, float, float]], + horizontal: bool, + scatter_kwargs: dict, +) -> None: + """One lollipop panel: a stem from 0 to each value, lighter than its dot.""" + band = panel.axhspan if horizontal else panel.axvspan + line = panel.axhline if horizontal else panel.axvline + for verdict, lo, hi in bands: + band(lo, hi, color=WEAK_COLOURS[verdict], alpha=0.10, lw=0) + # Dash the claim threshold, i.e. the band edge inside [0, 1]. + line(hi if lo < 0 else lo, color="0.4", ls="--", lw=1) + + stems = [_lighten(c) for c in colours] + n = len(positions) + if horizontal: + panel.vlines(positions, 0.0, values, colors=stems, lw=2.5, zorder=2) + panel.scatter(positions, values, c=colours, **scatter_kwargs) + panel.set_xticks(positions, labels, rotation=60, ha="right") + panel.set_xlim(-0.7, n - 0.3) + panel.set_ylim(-0.02, 1.02) + panel.grid(axis="x", visible=False) + panel.set_ylabel(value_label) + else: + panel.hlines(positions, 0.0, values, colors=stems, lw=2.5, zorder=2) + panel.scatter(values, positions, c=colours, **scatter_kwargs) + panel.set_yticks(positions, labels) + panel.set_ylim(-0.7, n - 0.3) + panel.set_xlim(-0.02, 1.02) + panel.grid(axis="y", visible=False) + panel.set_xlabel(value_label) + + +def _add_title_and_legend( + axes: np.ndarray, + verdicts: np.ndarray, + threshold: float, + value_label: str, + subtitle: str | None, + horizontal: bool, +) -> None: + title = f"Weak interpretation: claim at {threshold}, {value_label}" + if subtitle is None: + axes[0].set_title(title, loc="left") + else: + axes[0].set_title(title, loc="left", pad=20) + axes[0].text( + 0.0, + 1.01, + subtitle, + transform=axes[0].transAxes, + ha="left", + va="bottom", + fontsize="small", + color="0.35", + ) + + present = set(verdicts.tolist()) + handles = [ + Line2D([], [], marker="o", ls="", color=WEAK_COLOURS[v], markersize=8, label=v) + for v in WEAK_VERDICTS + if v in present + ] + # Beside the top panel when stacked, beside the right one when side by side. + (axes[0] if horizontal else axes[1]).legend( + handles=handles, + loc="upper left", + bbox_to_anchor=(1.01, 1.0), + fontsize=9, + frameon=False, + ) diff --git a/tests/tests/bbt/test_bbt.py b/tests/tests/bbt/test_bbt.py index 761c633..ddfd86d 100644 --- a/tests/tests/bbt/test_bbt.py +++ b/tests/tests/bbt/test_bbt.py @@ -382,6 +382,321 @@ def test_best_model_is_rank_one(self, fitted_model): assert by_x == by_beta +class TestStrongPosteriorPlot: + """``plot(kind="strong-posterior")`` draws one interval per comparison.""" + + @pytest.fixture(autouse=True) + def _agg_backend(self): + import matplotlib as mpl + + mpl.use("Agg") + import matplotlib.pyplot as plt + + yield + plt.close("all") + + def test_selected_pairs_keep_their_order(self, fitted_model): + """Each pair is read as given, so a worse model on the left sits below 0.5.""" + pairs = [("model_a", "model_b"), ("model_c", "model_a")] + ax = fitted_model.plot( + kind="strong-posterior", selected_pairs=pairs, orientation="vertical" + ) + labels = [t.get_text() for t in ax.get_yticklabels()] + drawn = dict(zip(labels, ax.collections[-1].get_offsets()[:, 0], strict=False)) + assert set(drawn) == {"model_a > model_b", "model_c > model_a"} + + expected = fitted_model.pairwise_samples(pairs).mean() + for label, value in drawn.items(): + assert value == pytest.approx(expected[label]) + + @pytest.mark.parametrize( + "kwargs", + [ + {}, + {"control_model": "model_a", "selected_pairs": [("model_b", "model_c")]}, + ], + ids=["neither", "both"], + ) + def test_requires_exactly_one_of_control_or_pairs(self, fitted_model, kwargs): + """Every pair of a large benchmark is unreadable, so it is not a default.""" + with pytest.raises(ValueError, match="exactly one of control_model"): + fitted_model.plot(**kwargs) + + def test_selected_models_needs_a_control(self, fitted_model): + """selected_models filters rows against a control, not a pair list.""" + with pytest.raises(ValueError, match="selected_models only applies"): + fitted_model.plot( + selected_pairs=[("model_a", "model_b")], + selected_models=["model_a", "model_b"], + ) + + def test_unknown_pair_model_raises(self, fitted_model): + """A misspelled name in a pair fails loudly.""" + with pytest.raises(ValueError, match="Unknown algorithms"): + fitted_model.plot(selected_pairs=[("model_a", "model_z")]) + + def test_control_orients_every_pair_against_it(self, fitted_model): + """With a control, rows are the other models at P(model > control).""" + ax = fitted_model.plot( + kind="strong-posterior", control_model="model_b", orientation="vertical" + ) + labels = [t.get_text() for t in ax.get_yticklabels()] + drawn = dict(zip(labels, ax.collections[-1].get_offsets()[:, 0], strict=False)) + assert set(drawn) == {"model_a", "model_c"} + + expected = fitted_model.pairwise_samples( + [("model_a", "model_b"), ("model_c", "model_b")] + ).mean() + assert drawn["model_a"] == pytest.approx(expected["model_a > model_b"]) + assert drawn["model_c"] == pytest.approx(expected["model_c > model_b"]) + # Rows run from the lowest mean at the bottom to the highest at the top. + assert list(drawn.values()) == sorted(drawn.values()) + + def test_horizontal_is_the_default_best_on_the_left(self, fitted_model): + """By default comparisons run along the x axis, highest mean first.""" + ax = fitted_model.plot(control_model="model_b") + labels = [t.get_text() for t in ax.get_xticklabels()] + assert set(labels) == {"model_a", "model_c"} + + offsets = ax.collections[-1].get_offsets() + drawn = dict(zip(labels, offsets[:, 1], strict=False)) + expected = fitted_model.pairwise_samples( + [("model_a", "model_b"), ("model_c", "model_b")] + ).mean() + assert drawn["model_a"] == pytest.approx(expected["model_a > model_b"]) + assert list(offsets[:, 1]) == sorted(offsets[:, 1], reverse=True) + + def test_unknown_orientation_raises(self, fitted_model): + """Orientation is validated like the other literal options.""" + with pytest.raises(ValueError, match="Invalid value 'diagonal'"): + fitted_model.plot(control_model="model_a", orientation="diagonal") + + @pytest.mark.parametrize("orientation", ["horizontal", "vertical"]) + def test_control_model_is_named_in_a_subtitle(self, fitted_model, orientation): + """Labels drop the control, so the figure must say what they compare to.""" + ax = fitted_model.plot(control_model="model_b", orientation=orientation) + assert "Control model: model_b" in [t.get_text() for t in ax.texts] + + def test_selected_pairs_have_no_subtitle(self, fitted_model): + """Pair labels already name both models.""" + ax = fitted_model.plot(selected_pairs=[("model_a", "model_b")]) + assert not ax.texts + + def test_draws_on_given_axes(self, fitted_model): + """A passed Axes is used, not replaced.""" + import matplotlib.pyplot as plt + + _, ax = plt.subplots() + assert fitted_model.plot(control_model="model_a", ax=ax) is ax + + def test_plot_is_a_shorthand_for_the_dedicated_method(self, fitted_model): + """Both entry points draw the same values.""" + via_plot = fitted_model.plot(control_model="model_b", hdi_prob=0.8) + direct = fitted_model.plot_strong_posterior( + control_model="model_b", hdi_prob=0.8 + ) + assert ( + via_plot.collections[-1].get_offsets() + == direct.collections[-1].get_offsets() + ).all() + + def test_dedicated_method_applies_selection_rules(self, fitted_model): + """The checks live in the shared helper, not only in ``plot``.""" + with pytest.raises(ValueError, match="exactly one of control_model"): + fitted_model.plot_strong_posterior() + + def test_unknown_kind_raises(self, fitted_model): + """Only the advertised kinds are accepted.""" + with pytest.raises(ValueError, match="Invalid value 'density'"): + fitted_model.plot(kind="density", control_model="model_a") + + def test_unknown_control_raises(self, fitted_model): + """A misspelled control fails loudly, as in ``posterior_table``.""" + with pytest.raises(ValueError, match="Unknown control_model"): + fitted_model.plot(control_model="model_z") + + +class TestStrongVerdicts: + """The strong rule is symmetric around 0.5.""" + + def test_thresholds(self): + """Means on and around each threshold get the expected verdict.""" + from btk.tests.bbt.plots._strong_posterior import strong_verdicts + + means = np.array([0.1, 0.3, 0.4, 0.45, 0.5, 0.55, 0.6, 0.7, 0.9]) + assert strong_verdicts(means, 0.70, 0.55).tolist() == [ + "weaker", + "no claim", + "no claim", + "equivalent", + "equivalent", + "equivalent", + "no claim", + "no claim", + "better", + ] + + +class TestWeakPosteriorPlot: + """``plot(kind="weak-posterior")`` draws P(pi > 0.5) and P(pi in ROPE) in two panels.""" + + @pytest.fixture(autouse=True) + def _agg_backend(self): + import matplotlib as mpl + + mpl.use("Agg") + import matplotlib.pyplot as plt + + yield + plt.close("all") + + @staticmethod + def _expected(fitted_model, pairs, rope): + draws = fitted_model.pairwise_samples(pairs) + above = (draws > 0.5).mean() + in_rope = ((draws >= rope[0]) & (draws <= rope[1])).mean() + return above, in_rope + + def test_returns_two_axes(self, fitted_model): + """One panel per quantity.""" + axes = fitted_model.plot(kind="weak-posterior", control_model="model_a") + assert len(axes) == 2 + + def test_control_orients_every_pair_against_it(self, fitted_model): + """Horizontal: models on x, P(pi > 0.5) on top, P(pi in ROPE) below.""" + rope = (0.4, 0.6) + top, bottom = fitted_model.plot( + kind="weak-posterior", control_model="model_b", rope_value=rope + ) + labels = [t.get_text() for t in bottom.get_xticklabels()] + assert set(labels) == {"model_a", "model_c"} + + above, in_rope = self._expected( + fitted_model, [("model_a", "model_b"), ("model_c", "model_b")], rope + ) + drawn_above = dict( + zip(labels, top.collections[-1].get_offsets()[:, 1], strict=False) + ) + drawn_rope = dict( + zip(labels, bottom.collections[-1].get_offsets()[:, 1], strict=False) + ) + for model in ("model_a", "model_c"): + key = f"{model} > model_b" + assert drawn_above[model] == pytest.approx(above[key]) + assert drawn_rope[model] == pytest.approx(in_rope[key]) + + def test_best_first_by_posterior_mean(self, fitted_model): + """Horizontal runs best-left; vertical puts the best on top.""" + means = fitted_model.pairwise_samples( + [("model_a", "model_b"), ("model_c", "model_b")] + ).mean() + best = means.idxmax().split(" > ")[0] + + _, bottom = fitted_model.plot(kind="weak-posterior", control_model="model_b") + assert bottom.get_xticklabels()[0].get_text() == best + + left, _ = fitted_model.plot( + kind="weak-posterior", control_model="model_b", orientation="vertical" + ) + assert left.get_yticklabels()[-1].get_text() == best + + def test_selected_pairs_keep_their_order(self, fitted_model): + """Vertical: pairs on y, both quantities on x, each pair read as given.""" + pairs = [("model_a", "model_b"), ("model_c", "model_a")] + left, right = fitted_model.plot( + kind="weak-posterior", selected_pairs=pairs, orientation="vertical" + ) + labels = [t.get_text() for t in left.get_yticklabels()] + assert set(labels) == {"model_a > model_b", "model_c > model_a"} + + above, in_rope = self._expected(fitted_model, pairs, (0.45, 0.55)) + drawn_above = dict( + zip(labels, left.collections[-1].get_offsets()[:, 0], strict=False) + ) + drawn_rope = dict( + zip(labels, right.collections[-1].get_offsets()[:, 0], strict=False) + ) + for label in labels: + assert drawn_above[label] == pytest.approx(above[label]) + assert drawn_rope[label] == pytest.approx(in_rope[label]) + + def test_control_model_is_named_in_a_subtitle(self, fitted_model): + """Labels drop the control, so the figure must say what they compare to.""" + top, _ = fitted_model.plot(kind="weak-posterior", control_model="model_b") + assert "Control model: model_b" in [t.get_text() for t in top.texts] + + def test_draws_on_given_axes(self, fitted_model): + """Two passed Axes are used, not replaced.""" + import matplotlib.pyplot as plt + + _, given = plt.subplots(1, 2) + axes = fitted_model.plot( + kind="weak-posterior", control_model="model_a", ax=given + ) + assert list(axes) == list(given) + + @pytest.mark.parametrize("n_axes", [1, 3]) + def test_wrong_number_of_axes_says_which_panel_goes_where( + self, fitted_model, n_axes + ): + """Two panels need two Axes, and the error names what each one holds.""" + import matplotlib.pyplot as plt + + _, given = plt.subplots(1, n_axes) + with pytest.raises(ValueError, match=r"ax\[0\] for P\(pi > 0.5\) and ax\[1\]"): + fitted_model.plot_weak_posterior(control_model="model_a", ax=given) + + def test_unused_arguments_are_ignored(self, fitted_model): + """``plot`` takes every kind's options; each kind skips the ones it does not use.""" + axes = fitted_model.plot( + kind="weak-posterior", control_model="model_a", hdi_prob=0.5 + ) + assert len(axes) == 2 + ax = fitted_model.plot( + kind="strong-posterior", control_model="model_a", rope_value=(0.4, 0.6) + ) + assert ax.get_title(loc="left").startswith("Strong interpretation") + + def test_plot_is_a_shorthand_for_the_dedicated_method(self, fitted_model): + """Both entry points draw the same values.""" + rope = (0.4, 0.6) + via_plot = fitted_model.plot( + kind="weak-posterior", control_model="model_b", rope_value=rope + ) + direct = fitted_model.plot_weak_posterior( + control_model="model_b", rope_value=rope + ) + for a, b in zip(via_plot, direct, strict=False): + assert ( + a.collections[-1].get_offsets() == b.collections[-1].get_offsets() + ).all() + + def test_requires_exactly_one_of_control_or_pairs(self, fitted_model): + """Shares the selection rules of the strong plot.""" + with pytest.raises(ValueError, match="exactly one of control_model"): + fitted_model.plot(kind="weak-posterior") + + +class TestWeakVerdicts: + """The weak rule checks equivalence first and is symmetric around 0.5.""" + + def test_thresholds(self): + """Equivalence wins over better, and P(pi < 0.5) gives weaker.""" + from btk.tests.bbt.plots._weak_posterior import weak_verdicts + + above_50 = np.array([0.99, 0.99, 0.01, 0.50, 0.94, 0.03]) + below_50 = 1.0 - above_50 + in_rope = np.array([0.00, 0.96, 0.00, 0.95, 0.10, 0.94]) + assert weak_verdicts(above_50, below_50, in_rope, 0.95).tolist() == [ + "better", + "equivalent", + "weaker", + "equivalent", + "no claim", + "weaker", + ] + + class TestBBTTestInitialization: """Test BBTTest initialization and parameter validation.""" From b9e86ebdb2ee0841ae193068899827bf3df88c21 Mon Sep 17 00:00:00 2001 From: Mateusz Praski Date: Thu, 17 Sep 2026 12:33:02 +0200 Subject: [PATCH 2/2] Add CDD plot to explicit documentation of BBT --- btk/tests/bbt/_types.py | 2 +- btk/tests/bbt/bbt.py | 56 +++++++++++++++++++++++++++++++++---- tests/tests/bbt/test_bbt.py | 45 +++++++++++++++++++++++++++++ 3 files changed, 96 insertions(+), 7 deletions(-) diff --git a/btk/tests/bbt/_types.py b/btk/tests/bbt/_types.py index 7bb35a0..e02d9bd 100644 --- a/btk/tests/bbt/_types.py +++ b/btk/tests/bbt/_types.py @@ -25,7 +25,7 @@ ] # Figures dispatched through ``BBTTest.plot(kind=...)``. -PlotKindType = Literal["strong-posterior", "weak-posterior"] +PlotKindType = Literal["strong-posterior", "weak-posterior", "cdd"] PlotOrientationType = Literal["horizontal", "vertical"] InterpretationTypes = Literal[ diff --git a/btk/tests/bbt/bbt.py b/btk/tests/bbt/bbt.py index c42aadc..e11c6e0 100644 --- a/btk/tests/bbt/bbt.py +++ b/btk/tests/bbt/bbt.py @@ -711,6 +711,7 @@ def plot( selected_models: Iterable[str] | None = None, hdi_prob: float = 0.89, rope_value: tuple[float, float] = (0.45, 0.55), + interpretation: InterpretationTypes = "weak", orientation: PlotOrientationType = "horizontal", ax: plt.Axes | Sequence[plt.Axes] | None = None, **kwargs, @@ -727,9 +728,13 @@ def plot( The figure to draw. Defaults to `strong-posterior`. - `strong-posterior` - :meth:`plot_strong_posterior`; uses - ``hdi_prob``, ignores ``rope_value``. + ``control_model`` or ``selected_pairs``, ``selected_models``, + ``hdi_prob``, ``orientation`` and ``ax``. - `weak-posterior` - :meth:`plot_weak_posterior`; uses - ``rope_value``, ignores ``hdi_prob``. + ``control_model`` or ``selected_pairs``, ``selected_models``, + ``rope_value``, ``orientation`` and ``ax``. + - `cdd` - :meth:`plot_cdd_diagram`; uses ``rope_value``, + ``interpretation`` and ``ax``. control_model : str | None, optional Compare every other model against this one. Defaults to None. @@ -741,24 +746,33 @@ def plot( Probability mass of the HDIs. Defaults to 0.89. rope_value : tuple[float, float], optional Region of Practical Equivalence (ROPE). Defaults to (0.45, 0.55). + interpretation : {"weak", "strong"}, optional + Reading rule that decides which models the CDD bars join. Defaults + to "weak". orientation : str, optional `horizontal` or `vertical`. Defaults to `horizontal`. ax : plt.Axes | Sequence[plt.Axes] | None, optional - One Axes for `strong-posterior`, two for `weak-posterior`. If None, - a new figure is created. + One Axes for `strong-posterior` and `cdd`, two for `weak-posterior`. + If None, a new figure is created. **kwargs Additional keyword arguments passed to the underlying plotting function. Returns ------- plt.Axes | np.ndarray - A single Axes for `strong-posterior`, an array of two for `weak-posterior`. + A single Axes for `strong-posterior` and `cdd`, an array of two for + `weak-posterior`. See Also -------- plot_strong_posterior : Posterior mean and HDI under the strong interpretation. plot_weak_posterior : P(pi > 0.5) and P(pi in ROPE) under the weak interpretation. + plot_cdd_diagram : Critical difference diagram of the whole ranking. """ + if kind == "cdd": + return self.plot_cdd_diagram( + rope_value=rope_value, interpretation=interpretation, ax=ax, **kwargs + ) selection = { "control_model": control_model, "selected_pairs": selected_pairs, @@ -976,7 +990,37 @@ def plot_cdd_diagram( ax: plt.Axes | None = None, **kwargs, ) -> plt.Axes: - """Plot critical difference diagram for the fitted BBT model.""" + r"""Plot a critical difference diagram for the fitted BBT model. + + Models are placed on a ruler by their rank in mean :math:`\beta`, best at + the "better" end. A bar spans each maximal group of models that the chosen + reading rule calls pairwise equivalent. Models not joined by a bar are + not claimed equivalent, which is not the same as claimed different. + + Parameters + ---------- + rope_value : tuple[float, float], optional + Region of Practical Equivalence (ROPE). Defaults to (0.45, 0.55). + interpretation : {"weak", "strong"}, optional + Reading rule applied to every pair, see :meth:`posterior_table`. + Defaults to "weak". + ax : plt.Axes | None, optional + Matplotlib Axes to plot on. If None, a new figure and axes will be created. + **kwargs + Additional keyword arguments passed to the underlying plotting + function: ``bar_y_spacing``, ``xlabel_spacing`` and + ``draw_equivalence_lines_to_axis``. + + Returns + ------- + plt.Axes + The Axes the figure was drawn on. + + See Also + -------- + plot_strong_posterior : Pairwise posteriors against a control, strong reading. + plot_weak_posterior : Pairwise posteriors against a control, weak reading. + """ self._check_if_fitted() posterior_df = self.posterior_table( rope_value=rope_value, diff --git a/tests/tests/bbt/test_bbt.py b/tests/tests/bbt/test_bbt.py index ddfd86d..23a6875 100644 --- a/tests/tests/bbt/test_bbt.py +++ b/tests/tests/bbt/test_bbt.py @@ -697,6 +697,51 @@ def test_thresholds(self): ] +@pytest.mark.filterwarnings("ignore:No groups of equivalent algorithms") +class TestCddViaPlot: + """``plot(kind="cdd")`` is a shorthand for ``plot_cdd_diagram``.""" + + @pytest.fixture(autouse=True) + def _agg_backend(self): + import matplotlib as mpl + + mpl.use("Agg") + import matplotlib.pyplot as plt + + yield + plt.close("all") + + def test_draws_on_given_axes_without_a_control(self, fitted_model): + """The CDD ranks every model, so control and pair arguments are not required.""" + import matplotlib.pyplot as plt + + _, ax = plt.subplots() + assert fitted_model.plot(kind="cdd", ax=ax) is ax + + def test_matches_dedicated_method(self, fitted_model): + """Same ROPE and interpretation give the same labels and lines.""" + import matplotlib.pyplot as plt + + _, (a, b) = plt.subplots(1, 2) + options = {"rope_value": (0.4, 0.6), "interpretation": "strong"} + fitted_model.plot(kind="cdd", ax=a, **options) + fitted_model.plot_cdd_diagram(ax=b, **options) + assert [t.get_text() for t in a.texts] == [t.get_text() for t in b.texts] + assert len(a.lines) == len(b.lines) + + def test_ignores_posterior_plot_arguments(self, fitted_model): + """Arguments of the other kinds are ignored rather than rejected.""" + ax = fitted_model.plot( + kind="cdd", control_model="model_a", hdi_prob=0.5, orientation="vertical" + ) + assert ax is not None + + def test_unknown_interpretation_raises(self, fitted_model): + """``interpretation`` is validated on the shorthand too.""" + with pytest.raises(ValueError, match="Invalid value 'medium'"): + fitted_model.plot(kind="cdd", interpretation="medium") + + class TestBBTTestInitialization: """Test BBTTest initialization and parameter validation."""