Source code for proteopy.pl.upset

"""UpSet plots of feature membership across annotation categories.

Importing this module also installs compatibility patches for
``upsetplot`` 0.9.0 (the version range pinned in ``pyproject.toml``),
which is otherwise unusable on current pandas/numpy:

1. Per-dot style defaults of the UpSet matrix. upsetplot fills them
   with ``styles["linewidth"].fillna(1, inplace=True)`` and three
   sibling calls; under pandas' Copy-on-Write (mandatory in pandas 3)
   these fills silently do nothing, and every ``UpSet.plot()`` fails in
   ``Axes.scatter`` with ``ValueError: Invalid RGBA argument: nan``.
   The patch restores the defaults where the scatter call consumes
   them and silences the accompanying chained-assignment warning
   (``FutureWarning`` on pandas 2, ``ChainedAssignmentError`` on
   pandas 3). It is harmless on pandas 2.
2. Aggregation of a single-category input, which loses its
   ``MultiIndex`` under pandas' groupby and then fails in
   ``Series.reorder_levels``.
3. Count labels, positioned with a one-element array that numpy >= 2.4
   refuses to convert to a scalar when the figure is drawn.
4. The totals axis of an input whose categories are all empty, which
   triggers matplotlib's "identical xlims" ``UserWarning``; matplotlib
   expands the limits itself, so the warning is silenced.

Each patch is installed at most once. Remove them once ``upsetplot``
ships a release fixing these defects.
"""

from __future__ import annotations

import math
import warnings
from pathlib import Path

import anndata as ad
import matplotlib.axes
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from matplotlib.axes import Axes
from scipy import sparse
from upsetplot import UpSet, reformat

from proteopy.utils.anndata import check_proteodata

_NO_CATEGORY_LABEL = "No category"
_COUNT_NAME = "n_features"


# -- upsetplot 0.9.0 compatibility patches (see the module docstring)

_DOT_STYLE_PATCHED_FLAG = "_proteopy_pandas3_dot_style_patch"
_AGG_PATCHED_FLAG = "_proteopy_single_category_agg_patch"
_LABEL_PATCHED_FLAG = "_proteopy_numpy_label_position_patch"
_TOTALS_PATCHED_FLAG = "_proteopy_empty_totals_xlim_patch"

# Defaults that upsetplot intends to fill in, per scatter keyword.
_LITERAL_DOT_DEFAULTS = {
    "linewidths": 1,
    "linestyles": "solid",
}


def _is_missing(value) -> bool:
    return value is None or (isinstance(value, float) and math.isnan(value))


def _filled(values, defaults) -> list:
    """Replace missing entries of ``values`` using ``defaults``.

    ``defaults`` is either a single value used for every gap, or a
    sequence of the same length supplying a per-entry replacement.
    """
    items = list(values)
    if not any(_is_missing(item) for item in items):
        return items
    if isinstance(defaults, list) and len(defaults) == len(items):
        per_entry = defaults
    else:
        per_entry = [defaults] * len(items)
    return [
        per_entry[i] if _is_missing(item) else item
        for i, item in enumerate(items)
    ]


def _restore_dot_style_defaults(kwargs: dict, facecolor) -> None:
    for keyword, default in _LITERAL_DOT_DEFAULTS.items():
        if keyword in kwargs:
            kwargs[keyword] = _filled(kwargs[keyword], default)

    faces = None
    if "facecolors" in kwargs:
        faces = _filled(kwargs["facecolors"], facecolor)
        kwargs["facecolors"] = faces
    if "edgecolors" in kwargs:
        fallback = faces if faces is not None else facecolor
        kwargs["edgecolors"] = _filled(kwargs["edgecolors"], fallback)


def _patch_dot_styles() -> None:
    """Restore per-dot style defaults in ``UpSet.plot_matrix``, once."""
    if getattr(UpSet, _DOT_STYLE_PATCHED_FLAG, False):
        return

    original_plot_matrix = UpSet.plot_matrix

    def plot_matrix(self, ax):
        real_scatter = matplotlib.axes.Axes.scatter

        def scatter(axes, *args, **kwargs):
            _restore_dot_style_defaults(kwargs, self._facecolor)
            return real_scatter(axes, *args, **kwargs)

        matplotlib.axes.Axes.scatter = scatter
        try:
            with warnings.catch_warnings():
                # pandas 2: FutureWarning; pandas 3: ChainedAssignmentError
                warnings.filterwarnings(
                    "ignore",
                    message=("A value is (trying to be|being) set on a copy"),
                    category=Warning,
                )
                return original_plot_matrix(self, ax)
        finally:
            matplotlib.axes.Axes.scatter = real_scatter

    plot_matrix.__doc__ = original_plot_matrix.__doc__
    UpSet.plot_matrix = plot_matrix
    setattr(UpSet, _DOT_STYLE_PATCHED_FLAG, True)


def _patch_single_category_aggregation() -> None:
    """Keep a one-level ``MultiIndex`` through aggregation, once."""
    if getattr(reformat, _AGG_PATCHED_FLAG, False):
        return

    original_aggregate_data = reformat._aggregate_data

    def _aggregate_data(df, subset_size, sum_over):
        data, aggregated = original_aggregate_data(
            df,
            subset_size,
            sum_over,
        )
        if (
            isinstance(data.index, pd.MultiIndex)
            and data.index.nlevels == 1
            and not isinstance(aggregated.index, pd.MultiIndex)
        ):
            aggregated.index = pd.MultiIndex.from_arrays(
                [aggregated.index],
                names=[aggregated.index.name],
            )
        return data, aggregated

    _aggregate_data.__doc__ = original_aggregate_data.__doc__
    reformat._aggregate_data = _aggregate_data
    setattr(reformat, _AGG_PATCHED_FLAG, True)


def _as_scalar(value):
    """Unwrap a one-element array; return anything else unchanged."""
    array = np.asarray(value)
    if array.size == 1:
        return array.item()
    return value


def _patch_label_positions() -> None:
    """Make count-label positions scalars after ``_label_sizes``, once."""
    if getattr(UpSet, _LABEL_PATCHED_FLAG, False):
        return

    original_label_sizes = UpSet._label_sizes

    def _label_sizes(self, ax, rects, where):
        n_texts_before = len(ax.texts)
        # upsetplot's ``_label_sizes`` returns nothing.
        original_label_sizes(self, ax, rects, where)
        for text in list(ax.texts)[n_texts_before:]:
            x, y = text.get_position()
            text.set_position((_as_scalar(x), _as_scalar(y)))

    _label_sizes.__doc__ = original_label_sizes.__doc__
    UpSet._label_sizes = _label_sizes
    setattr(UpSet, _LABEL_PATCHED_FLAG, True)


def _patch_empty_totals() -> None:
    """Silence the singular-xlim warning of all-zero totals, once."""
    if getattr(UpSet, _TOTALS_PATCHED_FLAG, False):
        return

    original_plot_totals = UpSet.plot_totals

    def plot_totals(self, ax):
        with warnings.catch_warnings():
            warnings.filterwarnings(
                "ignore",
                message="Attempting to set identical low and high xlims",
                category=UserWarning,
            )
            return original_plot_totals(self, ax)

    plot_totals.__doc__ = original_plot_totals.__doc__
    UpSet.plot_totals = plot_totals
    setattr(UpSet, _TOTALS_PATCHED_FLAG, True)


_patch_dot_styles()
_patch_single_category_aggregation()
_patch_label_positions()
_patch_empty_totals()


# -- var_detected_by_cat_upset


def _check_flag(value, name: str) -> None:
    """Reject anything that is not an exact Python ``bool``."""
    if not isinstance(value, bool):
        raise TypeError(
            f"`{name}` must be a bool, got " f"{type(value).__name__}."
        )


def _resolve_threshold(
    min_count: int | None,
    min_fraction: float | None,
) -> tuple[str, int | float]:
    """Validate the thresholds and return the active one."""
    if min_count is not None and (
        isinstance(min_count, bool) or not isinstance(min_count, int)
    ):
        raise TypeError(
            "`min_count` must be a non-boolean int or None, got "
            f"{type(min_count).__name__}."
        )
    if min_fraction is not None and (
        isinstance(min_fraction, bool)
        or not isinstance(min_fraction, (int, float))
    ):
        raise TypeError(
            "`min_fraction` must be a non-boolean number or None, "
            f"got {type(min_fraction).__name__}."
        )

    if min_count is not None and min_fraction is not None:
        raise ValueError(
            "`min_count` and `min_fraction` are mutually exclusive. "
            "Provide one or neither."
        )

    if min_count is not None:
        if min_count < 0:
            raise ValueError("`min_count` must be greater than or equal to 0.")
        return "min_count", min_count

    if min_fraction is not None:
        # A chained comparison rather than math.isfinite: it also
        # rejects NaN and infinities, and cannot overflow on huge ints.
        if not 0 <= min_fraction <= 1:
            raise ValueError(
                "`min_fraction` must be a finite number between 0 " "and 1."
            )
        return "min_fraction", min_fraction

    return "min_count", 1


def _validate_upset_args(
    adata: ad.AnnData,
    cat_key,
    min_count,
    min_fraction,
    zero_to_na,
    print_stats,
    verbose,
    show,
    save,
) -> tuple[str, int | float]:
    """Validate all inputs and return the active threshold."""
    check_proteodata(adata)

    for value, name in (
        (zero_to_na, "zero_to_na"),
        (print_stats, "print_stats"),
        (verbose, "verbose"),
        (show, "show"),
    ):
        _check_flag(value, name)

    if save is not None and not isinstance(save, (str, Path)):
        raise TypeError(
            "`save` must be a str, Path or None, got "
            f"{type(save).__name__}."
        )

    if not isinstance(cat_key, str):
        raise TypeError(
            "`cat_key` must be a str, got " f"{type(cat_key).__name__}."
        )
    if cat_key == "":
        raise ValueError("`cat_key` must not be an empty string.")

    threshold = _resolve_threshold(min_count, min_fraction)

    if cat_key not in adata.obs.columns:
        raise KeyError(f"'{cat_key}' is not a column of `adata.obs`.")

    if adata.n_obs == 0 or adata.n_vars == 0:
        raise ValueError(
            "Cannot build an UpSet plot from an AnnData with an "
            "empty observation or variable axis."
        )

    if adata.obs[cat_key].isna().any():
        raise ValueError(
            f"`adata.obs['{cat_key}']` must not contain missing " "values."
        )

    return threshold


def _resolve_categories(
    series: pd.Series,
) -> tuple[list[str], list[np.ndarray]]:
    """Return the category names in spec order and their obs masks.

    Categorical columns keep their category order, including
    categories without observations; any other dtype is ordered by the
    lexicographic order of the ``str``-coerced unique values.
    """
    if isinstance(series.dtype, pd.CategoricalDtype):
        raw_values = list(series.cat.categories)
    else:
        raw_values = list(pd.unique(series))
        raw_values.sort(key=str)

    names = [str(value) for value in raw_values]
    if len(set(names)) != len(names):
        raise ValueError(
            "Category values collide after coercion to str; make "
            "the categories unique as strings."
        )

    coerced = series.astype(object).map(str).to_numpy()
    masks = [coerced == name for name in names]
    return names, masks


def _detection_matrix(
    adata: ad.AnnData,
    zero_to_na: bool,
) -> np.ndarray:
    """Return the boolean detection matrix of ``adata.X``."""
    matrix = adata.X
    if sparse.issparse(matrix):
        warnings.warn(
            "`adata.X` is sparse and is being densified to build "
            "the UpSet plot.",
            UserWarning,
            stacklevel=2,
        )
        matrix = matrix.toarray()

    values = np.asarray(matrix)
    if values.dtype.kind not in "fiu":
        values = values.astype(float)

    detected = ~np.isnan(values)
    if zero_to_na:
        detected &= values != 0
    return detected


def _membership_matrix(
    detected: np.ndarray,
    masks: list[np.ndarray],
    threshold_name: str,
    threshold_value: int | float,
) -> np.ndarray:
    """Return the (features x categories) boolean membership matrix."""
    n_vars = detected.shape[1]
    membership = np.zeros((n_vars, len(masks)), dtype=bool)

    for index, mask in enumerate(masks):
        n_obs_cat = int(mask.sum())
        if n_obs_cat == 0:
            # An empty category has no members, for any threshold.
            continue
        counts = detected[mask, :].sum(axis=0)
        if threshold_name == "min_count":
            membership[:, index] = counts >= threshold_value
        else:
            membership[:, index] = counts / n_obs_cat >= threshold_value
    return membership


def _intersection_counts(
    membership: np.ndarray,
    names: list[str],
) -> pd.Series:
    """Count features per membership vector, plus ``No category``."""
    counts: dict[tuple[bool, ...], int] = {}
    for row in membership:
        key = tuple(bool(value) for value in row)
        counts[key] = counts.get(key, 0) + 1

    all_false = (False,) * len(names)
    if all_false not in counts:
        counts[all_false] = 0

    index = pd.MultiIndex.from_tuples(
        list(counts.keys()),
        names=names,
    )
    return pd.Series(
        list(counts.values()),
        index=index,
        dtype=int,
        name=_COUNT_NAME,
    )


def _print_stats_df(df: pd.DataFrame) -> None:
    """Print a DataFrame with one-decimal formatting."""
    print(df.to_string(index=False, float_format="%.1f"))


def _global_stats_df(counts: pd.Series) -> pd.DataFrame:
    """One-row summary of the intersection counts."""
    return pd.DataFrame(
        {
            "count": [counts.count()],
            "mean": [counts.mean()],
            "median": [counts.median()],
            "std": [counts.std()],
            "min": [counts.min()],
            "max": [counts.max()],
        }
    )


def _intersection_label(
    vector: tuple[bool, ...],
    names: list[str],
) -> str:
    """Human-readable label of one membership vector."""
    members = [name for name, flag in zip(names, vector) if flag]
    if not members:
        return _NO_CATEGORY_LABEL
    return " & ".join(members)


def _intersections_df(
    counts: pd.Series,
    names: list[str],
) -> pd.DataFrame:
    """One row per intersection, sorted by size then label."""
    rows = []
    for vector, count in counts.items():
        vector = tuple(vector)
        rows.append([*vector, int(count), _intersection_label(vector, names)])

    # Built and sorted by position, then named: a category may itself
    # be named "n_features" or "label", so the headers can repeat.
    count_pos = len(names)
    df = pd.DataFrame(rows).sort_values(
        [count_pos, count_pos + 1],
        ascending=[False, True],
        kind="mergesort",
    )
    df.columns = [*names, _COUNT_NAME, "label"]
    return df


def _per_category_df(
    membership: np.ndarray,
    names: list[str],
    cat_key: str,
    n_vars: int,
) -> pd.DataFrame:
    """One row per category with its member-feature count."""
    n_features = membership.sum(axis=0).astype(int)
    # Built by position: ``cat_key`` may be "n_features" or "percent".
    df = pd.DataFrame(
        {
            0: names,
            1: n_features,
            2: 100 * n_features / n_vars,
        }
    )
    df.columns = [cat_key, _COUNT_NAME, "percent"]
    return df


def _print_upset_stats(
    counts: pd.Series,
    membership: np.ndarray,
    names: list[str],
    cat_key: str,
    n_vars: int,
) -> None:
    """Print the three tables underlying the plot."""
    print("Global:")
    _print_stats_df(_global_stats_df(counts))
    print("Intersections:")
    _print_stats_df(_intersections_df(counts, names))
    print(f"Per {cat_key}:")
    _print_stats_df(_per_category_df(membership, names, cat_key, n_vars))


def _print_verbose_report(
    cat_key: str,
    threshold_name: str,
    threshold_value: int | float,
    n_vars: int,
    n_categories: int,
) -> None:
    """Print a short report about the input of the plot."""
    print(
        f"Using the .X matrix.\n"
        f"Categories from .obs['{cat_key}'].\n"
        f"Threshold: {threshold_name} = {str(threshold_value)}\n"
        f"Features: {n_vars}\n"
        f"Categories: {n_categories}"
    )


def _plot_upset(counts: pd.Series) -> dict[str, Axes]:
    """Render the UpSet plot of the intersection counts."""
    upset = UpSet(
        counts,
        subset_size="sum",
        sort_by="degree",
        sort_categories_by="input",
        show_counts=True,
        include_empty_subsets=False,
    )
    return upset.plot()


[docs] def var_detected_by_cat_upset( adata: ad.AnnData, cat_key: str, min_count: int | None = None, min_fraction: float | None = None, zero_to_na: bool = False, print_stats: bool = False, verbose: bool = False, show: bool = True, save: str | Path | None = None, ) -> dict[str, Axes]: """ UpSet plot of feature membership across categories of an .obs column. A feature (a variable, i.e. a peptide or a protein) is a *member* of a category when it is detected in enough observations of that category. Detection is read from ``adata.X`` only: a value counts as detected when it is not NaN, and -- with ``zero_to_na=True`` -- not zero. The plot shows how many features share each combination of category memberships. Features that are a member of no category form their own intersection, shown with no filled matrix dots and labelled ``"No category"`` in the ``print_stats`` tables. Category order follows the default ProteoPy rule: the category order of ``adata.obs[cat_key]`` when it is a Categorical (store it as an ordered Categorical to control the order), otherwise the lexicographic order of the ``str``-coerced unique values. Parameters ---------- adata : AnnData ProteoPy AnnData (peptide- or protein-level). Intensities are read from ``adata.X`` only. cat_key : str Column in ``adata.obs`` defining the categories. min_count : int | None Minimum number of detected observations within a category for a feature to be a member of it. Non-boolean int >= 0. If both ``min_count`` and ``min_fraction`` are None, a threshold of ``min_count=1`` is used. min_fraction : float | None Minimum fraction of a category's observations in which a feature must be detected to be a member of it. Finite, non-boolean number in [0, 1]. Set at most one of ``min_count`` and ``min_fraction``; setting both raises ``ValueError``. zero_to_na : bool If True, zeros in ``.X`` count as missing. print_stats : bool If True, print the statistics underlying the plot. verbose : bool If True, print status messages about the input. show : bool Call ``plt.show()`` at the end. save : str | Path | None Path to save the figure to; None skips saving. Returns ------- dict[str, Axes] The axes returned by ``upsetplot.UpSet.plot()``, with keys ``"matrix"``, ``"intersections"``, ``"totals"`` and ``"shading"``. Raises ------ TypeError When an argument has the wrong type: ``save`` not a str, Path or None; a flag not a bool; ``cat_key`` not a str; ``min_count`` not a non-boolean int; ``min_fraction`` not a non-boolean number. KeyError When ``cat_key`` is not a column of ``adata.obs``. ValueError When both ``min_count`` and ``min_fraction`` are not None; when a value is invalid: ``cat_key == ""``; ``min_count < 0``; ``min_fraction`` non-finite or outside [0, 1]; missing values in ``adata.obs[cat_key]``; category values that collide after ``str`` coercion; an empty observation or variable axis. Warns ----- UserWarning When ``adata.X`` is sparse and is densified. Examples -------- Build a protein-level AnnData with six samples from three tissues. P4 is never measured; P3 is measured in only one lung sample. >>> import numpy as np >>> import pandas as pd >>> import anndata as ad >>> import proteopy as pr >>> samples = ["S1", "S2", "S3", "S4", "S5", "S6"] >>> proteins = ["P1", "P2", "P3", "P4"] >>> obs = pd.DataFrame( ... { ... "sample_id": samples, ... "tissue": [ ... "liver", "liver", "lung", "lung", "brain", "brain", ... ], ... }, ... index=samples, ... ) >>> var = pd.DataFrame({"protein_id": proteins}, index=proteins) >>> nan = np.nan >>> X = np.array([ ... [5.0, 2.0, 1.0, nan], ... [4.0, 3.0, 2.0, nan], ... [6.0, nan, 1.5, nan], ... [5.5, nan, nan, nan], ... [4.5, nan, nan, nan], ... [5.0, nan, nan, nan], ... ]) >>> adata = ad.AnnData(X=X, obs=obs, var=var) By default, one detection makes a protein a member of a tissue. >>> axes = pr.pl.var_detected_by_cat_upset(adata, cat_key="tissue") >>> sorted(axes) ['intersections', 'matrix', 'shading', 'totals'] Require detection in every sample of a tissue and print the counts behind the plot. >>> axes = pr.pl.var_detected_by_cat_upset( ... adata, ... cat_key="tissue", ... min_fraction=1.0, ... print_stats=True, ... ) Global: count mean median std min max 3 1.3 1.0 0.6 1 2 Intersections: brain liver lung n_features label False True False 2 liver False False False 1 No category True True True 1 brain & liver & lung Per tissue: tissue n_features percent brain 1 25.0 liver 3 75.0 lung 1 25.0 """ # -- Validate inputs threshold_name, threshold_value = _validate_upset_args( adata, cat_key, min_count, min_fraction, zero_to_na, print_stats, verbose, show, save, ) # -- Derive categories and feature memberships names, masks = _resolve_categories(adata.obs[cat_key]) detected = _detection_matrix(adata, zero_to_na) membership = _membership_matrix( detected, masks, threshold_name, threshold_value, ) counts = _intersection_counts(membership, names) # -- Report if verbose: _print_verbose_report( cat_key, threshold_name, threshold_value, adata.n_vars, len(names), ) if print_stats: _print_upset_stats( counts, membership, names, cat_key, adata.n_vars, ) # -- Plot axes = _plot_upset(counts) if save is not None: axes["matrix"].figure.savefig(save) if show: plt.show() return axes