Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 19 additions & 4 deletions btk/tests/bbt/bbt.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,12 @@
)
from .alg import _construct_win_table, _get_pwin, _hdi
from .model import _mcmcbbt_pymc
from .plots import plot_cdd_diagram, plot_strong_posterior, plot_weak_posterior
from .plots import (
describe_interpretation,
plot_cdd_diagram,
plot_strong_posterior,
plot_weak_posterior,
)


class BBTTest(BaseBayesianTest):
Expand Down Expand Up @@ -1008,8 +1013,10 @@ def plot_cdd_diagram(
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``.
function: ``bar_y_spacing``, ``xlabel_spacing``,
``draw_equivalence_lines_to_axis`` and ``interpretation_note``.
The note under the diagram spells out the reading rule that joins
models with a bar; pass ``interpretation_note=""`` to drop it.

Returns
-------
Expand All @@ -1021,6 +1028,15 @@ def plot_cdd_diagram(
plot_strong_posterior : Pairwise posteriors against a control, strong reading.
plot_weak_posterior : Pairwise posteriors against a control, weak reading.
"""
interpretation_col = self._get_interpretation_columns(interpretation)
if "interpretation_note" not in kwargs:
kwargs["interpretation_note"] = describe_interpretation(
interpretation_col,
rope_value=rope_value,
weak_threshold=self._WEAK_INTERPRETATION_THRESHOLD,
equal_threshold=self._STRONG_INTERPRETATION_EQUAL_THRESHOLD,
)

self._check_if_fitted()
posterior_df = self.posterior_table(
rope_value=rope_value,
Expand All @@ -1032,7 +1048,6 @@ def plot_cdd_diagram(
),
round_ndigits=None,
)
interpretation_col = self._get_interpretation_columns(interpretation)
# ``pos`` is the aggregated rank: 1 is the best algorithm, i.e. the
# highest mean beta. ``_plot_cdd_diagram`` draws ``pos = 1`` at the
# "better" end of the ruler, so the sort must be descending.
Expand Down
9 changes: 7 additions & 2 deletions btk/tests/bbt/plots/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,10 @@
from ._critical_difference import plot_cdd_diagram
from ._critical_difference import describe_interpretation, plot_cdd_diagram
from ._strong_posterior import plot_strong_posterior
from ._weak_posterior import plot_weak_posterior

__all__ = ["plot_cdd_diagram", "plot_strong_posterior", "plot_weak_posterior"]
__all__ = [
"describe_interpretation",
"plot_cdd_diagram",
"plot_strong_posterior",
"plot_weak_posterior",
]
64 changes: 63 additions & 1 deletion btk/tests/bbt/plots/_critical_difference.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,52 @@
CDD plot will not contain any equivalence bars."""


def describe_interpretation(
interpretation_col: str,
rope_value: tuple[float, float] | None = None,
weak_threshold: float | None = None,
equal_threshold: float | None = None,
) -> str:
"""Describe in words the rule that joins two models with an equivalence bar.

Parameters
----------
interpretation_col : str
Column the bars were derived from, e.g. ``"weak_interpretation_raw"``
or ``"strong_interpretation_raw"``.
rope_value : tuple[float, float], optional
ROPE used by the weak reading. Defaults to (0.45, 0.55).
weak_threshold : float, optional
Posterior mass the weak reading requires. Defaults to 0.95.
equal_threshold : float, optional
Upper bound on the posterior mean used by the strong reading. Defaults
to 0.55.

Returns
-------
str
A single line naming the reading rule, drawn under the diagram.
"""
if interpretation_col.startswith("weak"):
if rope_value is None:
raise ValueError("rope_value must be provided for weak interpretation")
return (
f"Weak interpretation: a bar joins models with "
f"P(\u03c0 \u2208 [{rope_value[0]:g}, {rope_value[1]:g}]) "
f"\u2265 {weak_threshold:g}"
)
if interpretation_col.startswith("strong"):
if equal_threshold is None:
raise ValueError(
"equal_threshold must be provided for strong interpretation"
)
return (
f"Strong interpretation: a bar joins models with "
f"{1 - equal_threshold:g} \u2264 E[\u03c0] \u2264 {equal_threshold:g}"
)
return f"A bar joins models called equivalent by {interpretation_col!r}"


def get_bars_for_cdd(
posterior_df: pd.DataFrame,
models_df: pd.DataFrame,
Expand Down Expand Up @@ -89,6 +135,7 @@ def _plot_cdd_diagram(
ax: plt.Axes | None = None,
xlabel_spacing: int = 5,
draw_equivalence_lines_to_axis: bool = True,
interpretation_note: str | None = None,
) -> plt.Axes:
"""Plot a critical difference diagram."""
if ax is None:
Expand Down Expand Up @@ -166,7 +213,7 @@ def _plot_cdd_diagram(
# Clip axes
min_bar_y = ruler_y - 0.4 - max_bar_pos * bar_y_spacing
ax.set_xlim(0, n_models + 1)
ax.set_ylim(min_bar_y - 0.3, 2.5)
ax.set_ylim(min_bar_y - (0.45 if interpretation_note else 0.3), 2.5)
ax.axis("off")

# Legend
Expand All @@ -178,6 +225,15 @@ def _plot_cdd_diagram(
style="italic",
)

# Reading rule behind the equivalence bars
if interpretation_note:
ax.text(
0.5,
min_bar_y - 0.27,
interpretation_note,
fontsize=7,
)

return ax


Expand All @@ -189,6 +245,7 @@ def plot_cdd_diagram(
bar_y_spacing: float = 0.12,
xlabel_spacing: int = 5,
draw_equivalence_lines_to_axis: bool = True,
interpretation_note: str | None = None,
) -> plt.Axes:
"""Plot a critical difference diagram.

Expand All @@ -211,6 +268,10 @@ def plot_cdd_diagram(
Whether to draw equivalence lines to extend equivalence bars up to the axis.
If False, equivalence bars will not have vertical lines connecting them to
the axis. Default is True.
interpretation_note : str | None, optional
Line drawn under the diagram spelling out the rule that joins models
with a bar. Defaults to :func:`describe_interpretation` applied to
``interpretation_col``; pass ``""`` to draw nothing.
"""
if ax is not None and not isinstance(ax, plt.Axes):
raise ValueError("ax must be a matplotlib Axes object or None.")
Expand All @@ -229,4 +290,5 @@ def plot_cdd_diagram(
bar_y_spacing=bar_y_spacing,
xlabel_spacing=xlabel_spacing,
draw_equivalence_lines_to_axis=draw_equivalence_lines_to_axis,
interpretation_note=interpretation_note,
)
Loading