From bc745b3bc92b0ef527856457fa21c08d3880c49d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:35:43 +0200 Subject: [PATCH 1/6] tests/conftest.py Tests: add spec-guided TDD support conftest Add tests/conftest.py from the spec-guided-tdd skill, formatted to ProteoPy's black/flake8 standards. It only affects test modules marked `pytestmark = pytest.mark.spec_guided` (e.g. tests/pl/test_upset.py); all other tests run unchanged. - `pytest` (default, CI) runs the holdout and randomized tests and deselects the `_IMP` mirrors - `pytest --implementer` runs the `_IMP` and randomized tests with obscured failure output, for implementation sessions - `--seed N` reproduces randomized inputs (seed printed in the header) - provides the `rng` and `report_input` fixtures - checks test naming and holdout/_IMP pairing --- tests/conftest.py | 304 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 304 insertions(+) create mode 100644 tests/conftest.py diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..abc4318 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,304 @@ +""" +Spec-guided TDD support for pytest (soft workflow). + +Scope: only test modules marked with ``pytestmark = pytest.mark.spec_guided`` +are affected. All other tests in the suite run exactly as without this file, +in both modes, so the support conftest can live next to an existing suite. + +Naming is the only link between tests and the (unpushed) specs: + +* Every test is named ``test__``, where ```` is the test ID + from the test spec (e.g. T10, P1). No spec files or requirement markers are + referenced from code. +* Holdout tests carry the clean default names and live in ``Test``, first + in the module. Each has an implementer counterpart with the same name plus + the suffix ``_IMP`` in the class ``TestIMP`` after it: the mirror with + different inputs. +* Randomized tests (those using the ``rng`` fixture, i.e. property and + metamorphic tests) have no ``_IMP`` counterpart: fresh random inputs cannot + be memorized. They live in ``Test`` and run in both modes. + +Behavior in marked modules: + +* ``pytest`` human mode (default), behaves like any pytest run: + checks that holdout and ``_IMP`` tests pair up, + then runs the holdout, property and metamorphic + tests with full pytest detail. +* ``pytest --implementer`` implementer mode: runs the ``_IMP``, property and + metamorphic tests with obscured errors: test ID, + kind of failure, and for randomized tests the + failing input registered via ``report_input``. No + tracebacks, source, captured output or warning + locations are shown; ``--pdb`` / ``--trace`` are + refused. +* ``--seed N`` reproduces the random inputs (seed in the header). + +Soft workflow: nothing here prevents an agent from reading the tests. +""" + +from __future__ import annotations + +import random +import re + +import pytest +from _pytest.outcomes import Failed + +_MARKER = "spec_guided" +_SEED_KEY = pytest.StashKey[int]() +_SUFFIX = "_IMP" +_CLASS_SUFFIX = "IMP" +_NAME = re.compile(rf"^test_([A-Z]+\d+)_\w+?({_SUFFIX})?$") +_INPUT_ATTR = "_sgtdd_failing_input" + + +# -------------------------------------------------------------------------- +# Options, configuration and header +# -------------------------------------------------------------------------- + + +def pytest_addoption(parser): + group = parser.getgroup("spec-guided-tdd") + group.addoption( + "--implementer", + action="store_true", + default=False, + help=( + "Implementer mode: run _IMP and randomized tests with " + "obscured errors (no holdout tests)." + ), + ) + group.addoption( + "--seed", + type=int, + default=None, + help="Session seed for the rng fixture (default: random).", + ) + + +def pytest_configure(config): + config.addinivalue_line( + "markers", + f"{_MARKER}: module follows the spec-guided TDD layout " + "(holdout/_IMP pairs).", + ) + seed = config.getoption("--seed") + config.stash[_SEED_KEY] = ( + seed if seed is not None else random.SystemRandom().randrange(2**32) + ) + + if config.getoption("--implementer"): + if config.getoption("usepdb", False) or config.getoption( + "trace", False + ): + raise pytest.UsageError( + "--pdb and --trace are not available with --implementer" + ) + # The warnings summary prints the test source line that triggered a + # warning, which would reveal test inputs. pytest.warns still works. + config.addinivalue_line("filterwarnings", "ignore") + + +def pytest_report_header(config): + mode = "implementer" if config.getoption("--implementer") else "human" + return f"spec-guided-tdd: mode={mode} seed={config.stash[_SEED_KEY]}" + + +# -------------------------------------------------------------------------- +# Helpers +# -------------------------------------------------------------------------- + + +def _base_name(item) -> str: + return getattr(item, "originalname", item.name) + + +def _test_id(item) -> str: + m = _NAME.match(_base_name(item)) + return m.group(1) if m else _base_name(item) + + +def _is_spec_guided(item) -> bool: + return item.get_closest_marker(_MARKER) is not None + + +def _is_imp(item) -> bool: + return _base_name(item).endswith(_SUFFIX) + + +def _looks_like_imp(item) -> bool: + cls = getattr(item, "cls", None) + return _is_imp(item) or ( + cls is not None and cls.__name__.endswith(_CLASS_SUFFIX) + ) + + +def _is_randomized(item) -> bool: + return "rng" in getattr(item, "fixturenames", ()) + + +# -------------------------------------------------------------------------- +# Collection: scope, naming check, pairing check, selection +# -------------------------------------------------------------------------- + + +def pytest_collection_modifyitems(config, items): + implementer = config.getoption("--implementer") + guided = [it for it in items if _is_spec_guided(it)] + _check_unmarked(items) + _check_names(guided) + if not implementer: + _check_pairs(guided) + + def wanted(it) -> bool: + if implementer: + return _is_imp(it) or _is_randomized(it) + return not _is_imp(it) + + keep = [it for it in items if not _is_spec_guided(it) or wanted(it)] + drop = [it for it in items if _is_spec_guided(it) and not wanted(it)] + if drop: + config.hook.pytest_deselected(items=drop) + items[:] = keep + + +def _check_unmarked(items) -> None: + # A forgotten marker would run holdout tests unobscured in implementer + # mode; the _IMP counterparts in the same module give it away. + errors = [ + f"{it.nodeid}: {_SUFFIX}-style test outside a module marked " + f"'pytestmark = pytest.mark.{_MARKER}'" + for it in items + if not _is_spec_guided(it) and _looks_like_imp(it) + ] + if errors: + raise pytest.UsageError( + "Spec-guided scope errors:\n " + "\n ".join(errors) + ) + + +def _check_names(items) -> None: + errors: list[str] = [] + for item in items: + name = _base_name(item) + if not _NAME.match(name): + errors.append( + f"{item.nodeid}: name must be " + f"test__[{_SUFFIX}], e.g. test_T10_empty" + ) + continue + if item.cls is None: + errors.append( + f"{item.nodeid}: tests must live in a Test or " + f"Test{_CLASS_SUFFIX} class" + ) + continue + in_imp_class = item.cls.__name__.endswith(_CLASS_SUFFIX) + if _is_imp(item) != in_imp_class: + errors.append( + f"{item.nodeid}: *{_SUFFIX} tests belong in " + f"Test{_CLASS_SUFFIX}, all others in Test" + ) + if _is_imp(item) and _is_randomized(item): + errors.append( + f"{item.nodeid}: randomized tests (using rng) have no " + f"{_SUFFIX} counterpart" + ) + if errors: + raise pytest.UsageError( + "Test naming errors:\n " + "\n ".join(errors) + ) + + +def _check_pairs(items) -> None: + hold: set[tuple] = set() + imp: set[tuple] = set() + for item in items: + cls = item.cls.__name__ + name = _base_name(item) + if _is_imp(item): + imp.add( + ( + str(item.path), + cls[: -len(_CLASS_SUFFIX)], + name[: -len(_SUFFIX)], + ) + ) + elif not _is_randomized(item): + hold.add((str(item.path), cls, name)) + + errors = [ + f"{c}.{n}: no counterpart {c}{_CLASS_SUFFIX}.{n}{_SUFFIX}" + for _, c, n in sorted(hold - imp) + ] + errors += [ + f"{c}{_CLASS_SUFFIX}.{n}{_SUFFIX}: no holdout counterpart {c}.{n}" + for _, c, n in sorted(imp - hold) + ] + if errors: + raise pytest.UsageError( + "Holdout pairing errors:\n " + "\n ".join(errors) + ) + + +# -------------------------------------------------------------------------- +# Obscured errors in implementer mode +# -------------------------------------------------------------------------- + + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + report = outcome.get_result() + if not item.config.getoption("--implementer") or not _is_spec_guided(item): + return + # Never show captured output of guided tests, passed or failed. + report.sections = [] + if not report.failed: + return + + test_id = _test_id(item) + excinfo = call.excinfo + if report.when != "call": + line = f"{test_id}: error outside the check (test {report.when})" + elif excinfo is not None and not excinfo.errisinstance( + (AssertionError, Failed) + ): + line = f"{test_id}: unexpected exception ({excinfo.type.__name__})" + else: + line = f"{test_id}: wrong result" + + if hasattr(item, _INPUT_ATTR): + line += f"; failing input: {getattr(item, _INPUT_ATTR)}" + + report.longrepr = line + + +# -------------------------------------------------------------------------- +# Fixtures +# -------------------------------------------------------------------------- + + +@pytest.fixture +def rng(request) -> random.Random: + """Per-test random generator derived from the session seed. + + Using this fixture marks a test as randomized: it needs no _IMP + counterpart. + """ + seed = request.config.stash[_SEED_KEY] + return random.Random(f"{seed}:{request.node.nodeid}") + + +@pytest.fixture +def report_input(request): + """Register the current generated input. Use ONLY in randomized tests. + + Call it right before checking each generated input; on failure the last + registered input is shown. + """ + + def _register(value) -> None: + setattr(request.node, _INPUT_ATTR, repr(value)) + + return _register From c2299b95b55d930eb88a2598f1a3e0f35b9ae6e1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:38:43 +0200 Subject: [PATCH 2/6] tests/upset.py Tests: pl.var_detected_by_cat_upset() Add a spec-guided test suite for var_detected_by_cat_upset() (module marked `spec_guided`, run by tests/conftest.py): - TestVarDetectedByCatUpset: 75 holdout tests plus 9 randomized property and metamorphic tests - TestVarDetectedByCatUpsetIMP: 75 _IMP mirrors on different inputs, run only with `pytest --implementer` Coverage: intersection counts and thresholds (min_count / min_fraction, inclusivity, per-category denominators, empty categories, float edge cases), zero/NaN detection, category order and str coercion, sparse input, non-mutation (incl. views), UpSet construction and rendered axes, print_stats tables (incl. column-name collisions), verbose output, save/show, and argument validation. UpSet calls are observed through spies on the upsetplot class; randomized tests use the seeded rng fixture (reproduce with `--seed N`). --- tests/pl/test_upset.py | 2246 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 2246 insertions(+) create mode 100644 tests/pl/test_upset.py diff --git a/tests/pl/test_upset.py b/tests/pl/test_upset.py new file mode 100644 index 0000000..17c4d89 --- /dev/null +++ b/tests/pl/test_upset.py @@ -0,0 +1,2246 @@ +"""Tests for ``pr.pl.var_detected_by_cat_upset``. + +Sections +-------- +* Fixtures and helpers +* Core intersection counts and thresholds +* Category order and names +* Feature counting +* Degenerate structures +* Sparse input and mutation +* Plot construction +* print_stats output +* verbose output +* save and show +* Threshold selection +* Negative cases +* Property and metamorphic relations +""" + +from __future__ import annotations + +import inspect +import math + +import anndata as ad +import matplotlib +import numpy as np +import pandas as pd +import pytest +from scipy import sparse +from upsetplot import UpSet + +from proteopy.pl import var_detected_by_cat_upset + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 + +pytestmark = pytest.mark.spec_guided + +T = True +F = False + + +# --------------------------------------------------------------------------- +# AnnData builders +# --------------------------------------------------------------------------- + + +def _adata_from_columns( + cat_key, + cat_values, + columns, + *, + level="protein", + protein_ids=None, + categories=None, +): + names = list(columns) + n_obs = len(cat_values) + X = np.empty((n_obs, len(names)), dtype=float) + for j, name in enumerate(names): + X[:, j] = columns[name] + + sample_ids = [f"s{i + 1}" for i in range(n_obs)] + obs = pd.DataFrame({"sample_id": sample_ids}, index=pd.Index(sample_ids)) + if categories is None: + obs[cat_key] = pd.Series(cat_values, index=obs.index) + else: + obs[cat_key] = pd.Categorical(cat_values, categories=categories) + + var = pd.DataFrame(index=pd.Index(names)) + if level == "protein": + var["protein_id"] = names + else: + var["peptide_id"] = names + var["protein_id"] = [protein_ids[name] for name in names] + + return ad.AnnData(X=X, obs=obs, var=var) + + +def _adata( + cat_key, + cat_values, + detections, + value, + *, + level="protein", + protein_ids=None, + categories=None, +): + """Build an AnnData from per-feature detected-observation counts.""" + positions = {} + for i, cat in enumerate(cat_values): + positions.setdefault(cat, []).append(i) + + columns = {} + for name, per_cat in detections.items(): + column = [np.nan] * len(cat_values) + for cat, n_detected in per_cat.items(): + for i in positions[cat][:n_detected]: + column[i] = value + columns[name] = column + + return _adata_from_columns( + cat_key, + cat_values, + columns, + level=level, + protein_ids=protein_ids, + categories=categories, + ) + + +def _f1(): + proteins = { + "p1": "PA", + "p2": "PA", + "p3": "PA", + "p4": "PB", + "p5": "PB", + "p6": "PB", + "p7": "PC", + "p8": "PC", + "p9": "PC", + } + return _adata( + "tissue", + ["A", "B", "C", "A", "B", "C"], + { + "p1": {"A": 2}, + "p2": {"B": 2}, + "p3": {"C": 2}, + "p4": {"A": 2, "B": 2}, + "p5": {"A": 2, "C": 2}, + "p6": {"B": 2, "C": 2}, + "p7": {"A": 2, "B": 2, "C": 2}, + "p8": {}, + "p9": {"A": 1}, + }, + 10.0, + level="peptide", + protein_ids=proteins, + ) + + +def _f3(): + return _adata( + "group", + ["A"] * 4 + ["B"] * 4, + { + "f1": {"B": 4}, + "f2": {"A": 1, "B": 3}, + "f3": {"A": 2, "B": 2}, + "f4": {"A": 3, "B": 1}, + "f5": {"A": 4}, + }, + 2.0, + ) + + +def _f4(): + return _adata( + "group", + ["A"] * 4 + ["B"] * 2, + { + "f1": {"A": 2}, + "f2": {"A": 1, "B": 1}, + "f3": {"A": 3, "B": 2}, + "f4": {"A": 1}, + }, + 5.0, + ) + + +def _f5(): + return _adata_from_columns( + "group", + ["A", "A", "B", "B"], + { + "f1": [0.0, 0.0, np.nan, np.nan], + "f2": [np.nan, np.nan, 3.0, 0.0], + }, + ) + + +def _f6(): + return _adata( + "tissue", + ["A", "B", "C"], + {"f1": {"A": 1, "B": 1, "C": 1}, "f2": {"A": 1}}, + 1.0, + categories=["C", "A", "B", "D"], + ) + + +def _f7(): + return _adata("tissue", ["B", "C", "A"], {"f1": {"B": 1}}, 1.0) + + +def _h1(): + return _adata( + "organ", + ["lung", "kidney", "spleen"] * 3, + { + "g01": {"kidney": 3}, + "g02": {"lung": 3}, + "g03": {"spleen": 3}, + "g04": {"kidney": 3, "lung": 3}, + "g05": {"kidney": 3, "spleen": 3}, + "g06": {"lung": 3, "spleen": 3}, + "g07": {"kidney": 3, "lung": 3, "spleen": 3}, + "g08": {"kidney": 3, "lung": 3, "spleen": 3}, + "g09": {}, + "g10": {"kidney": 1, "lung": 2}, + "g11": {"kidney": 3}, + }, + 3.5, + ) + + +def _h3(): + names = [f"g{i}" for i in range(1, 7)] + return _adata( + "site", + ["north"] * 5 + ["south"] * 5, + { + "g1": {"south": 5}, + "g2": {"north": 1, "south": 4}, + "g3": {"north": 2, "south": 3}, + "g4": {"north": 3, "south": 2}, + "g5": {"north": 4, "south": 1}, + "g6": {"north": 5}, + }, + 7.0, + level="peptide", + protein_ids={name: "Q1" for name in names}, + ) + + +def _h4(): + return _adata( + "arm", + ["p"] * 5 + ["q"] * 3, + { + "g1": {"p": 3, "q": 1}, + "g2": {"p": 2, "q": 2}, + "g3": {"p": 5, "q": 3}, + "g4": {"p": 2, "q": 1}, + "g5": {"p": 4}, + }, + 8.0, + ) + + +def _h5(): + return _adata_from_columns( + "arm", + ["u"] * 3 + ["v"] * 3, + { + "g1": [0.0, 0.0, 7.5, np.nan, np.nan, np.nan], + "g2": [np.nan, np.nan, np.nan, 0.0, 0.0, 0.0], + "g3": [1.0, np.nan, 0.0, 0.0, np.nan, np.nan], + }, + ) + + +def _h6(): + return _adata( + "batch", + ["mid", "alpha", "zeta", "mid", "alpha", "zeta"], + { + "v1": {"zeta": 2, "alpha": 2, "mid": 2}, + "v2": {"mid": 1}, + "v3": {}, + }, + 4.0, + categories=["zeta", "beta", "alpha", "mid"], + ) + + +def _h7(): + return _adata( + "fruit", + ["kiwi", "apple", "fig", "apple"], + {"w1": {"fig": 1}, "w2": {"apple": 2}}, + 1.0, + ) + + +# --------------------------------------------------------------------------- +# Spies on the plotting library +# --------------------------------------------------------------------------- + + +def _bind(func, args, kwargs): + bound = inspect.signature(func).bind(None, *args, **kwargs) + bound.apply_defaults() + arguments = dict(bound.arguments) + arguments.pop("self", None) + return arguments + + +class _Spy: + def __init__(self): + self.init = [] + self.style = [] + self.plot_returns = [] + + @property + def data(self): + return self.init[0]["data"] + + +@pytest.fixture(autouse=True) +def show_calls(monkeypatch): + matplotlib.use("Agg") + recorded = [] + monkeypatch.setattr(plt, "show", lambda *a, **k: recorded.append(1)) + yield recorded + plt.close("all") + + +@pytest.fixture +def spy(monkeypatch): + recorder = _Spy() + original_init = UpSet.__init__ + original_style = UpSet.style_subsets + original_plot = UpSet.plot + + def init(self, *args, **kwargs): + recorder.init.append(_bind(original_init, args, kwargs)) + return original_init(self, *args, **kwargs) + + def style_subsets(self, *args, **kwargs): + recorder.style.append(_bind(original_style, args, kwargs)) + return original_style(self, *args, **kwargs) + + def plot(self, *args, **kwargs): + result = original_plot(self, *args, **kwargs) + recorder.plot_returns.append(result) + return result + + monkeypatch.setattr(UpSet, "__init__", init) + monkeypatch.setattr(UpSet, "style_subsets", style_subsets) + monkeypatch.setattr(UpSet, "plot", plot) + return recorder + + +def _series(spy, adata, cat_key, **kwargs): + kwargs.setdefault("show", False) + var_detected_by_cat_upset(adata, cat_key, **kwargs) + return spy.data + + +# --------------------------------------------------------------------------- +# Assertion helpers +# --------------------------------------------------------------------------- + + +def _check_series(series, levels, expected): + assert isinstance(series, pd.Series) + assert series.name == "n_features" + assert pd.api.types.is_integer_dtype(series.dtype) + assert isinstance(series.index, pd.MultiIndex) + assert list(series.index.names) == list(levels) + for level_values in series.index.levels: + assert level_values.dtype == np.dtype(bool) + observed = {tuple(key): int(value) for key, value in series.items()} + assert observed == {tuple(k): v for k, v in expected.items()} + + +def _vec(n_levels, members): + return tuple(i in members for i in range(n_levels)) + + +def _per_category(series, levels): + return { + name: int(sum(v for key, v in series.items() if key[position])) + for position, name in enumerate(levels) + } + + +def _matrix_labels(axes): + matrix = axes["matrix"] + matrix.figure.canvas.draw() + texts = [] + for axis in (matrix.xaxis, matrix.yaxis): + if axis.get_visible(): + texts += [label.get_text() for label in axis.get_ticklabels()] + return sorted(text for text in texts if text) + + +def _legend_texts(axes): + figure = axes["matrix"].figure + texts = [] + for legend in figure.legends: + texts += [entry.get_text() for entry in legend.get_texts()] + for axis in axes.values(): + legend = axis.get_legend() + if legend is not None: + texts += [entry.get_text() for entry in legend.get_texts()] + return texts + + +def _bar_heights(axes): + return sorted( + patch.get_height() for patch in axes["intersections"].patches + ) + + +def _bar_widths(axes): + return sorted(patch.get_width() for patch in axes["totals"].patches) + + +_TITLE_PREFIXES = ("Global:", "Intersections:", "Per ") + + +def _is_title(line): + stripped = line.strip() + return stripped.endswith(":") and stripped.startswith(_TITLE_PREFIXES) + + +def _section(out, title): + lines = out.splitlines() + start = next(i for i, line in enumerate(lines) if line.strip() == title) + body = [] + for line in lines[start + 1 :]: + if _is_title(line): + break + if line.strip(): + body.append(line) + return body + + +def _global_table(out): + body = _section(out, "Global:") + return body[0].split(), body[1].split() + + +def _intersections_table(out, n_categories): + body = _section(out, "Intersections:") + header = body[0].split() + rows = [ + [field.strip() for field in line.split(None, n_categories + 1)] + for line in body[1:] + ] + return header, rows + + +def _per_table(out, cat_key): + body = _section(out, f"Per {cat_key}:") + header = body[0].split() + rows = [ + [field.strip() for field in line.rsplit(None, 2)] for line in body[1:] + ] + return header, rows + + +# --------------------------------------------------------------------------- +# Randomized-input generation and reference computation +# --------------------------------------------------------------------------- + +_FRACTIONS = [0.0, 0.25, 1 / 3, 0.5, 2 / 3, 1.0] +_ROUNDS = 12 + + +def _generate(rng): + n_categories = rng.randint(1, 5) + categories = [f"c{i}" for i in range(n_categories)] + cat_values = [] + for category in categories: + cat_values += [category] * rng.randint(1, 4) + rng.shuffle(cat_values) + + columns = {} + for j in range(rng.randint(1, 30)): + column = [] + for _ in cat_values: + draw = rng.random() + if draw < 0.4: + column.append(float("nan")) + elif draw < 0.6: + column.append(0.0) + else: + column.append(rng.uniform(0.1, 100.0)) + columns[f"q{j:02d}"] = column + + if rng.random() < 0.5: + threshold = ("min_count", rng.randint(0, 5)) + else: + threshold = ("min_fraction", rng.choice(_FRACTIONS)) + + return { + "categories": categories, + "cat_values": cat_values, + "columns": columns, + "threshold": threshold, + "zero_to_na": rng.choice([True, False]), + } + + +def _compact(case): + """Readable form of a generated case for failure reporting.""" + return { + "cat_values": case["cat_values"], + "columns": { + name: [ + "nan" if math.isnan(entry) else round(entry, 2) + for entry in column + ] + for name, column in case["columns"].items() + }, + "threshold": case["threshold"], + "zero_to_na": case["zero_to_na"], + } + + +def _case_adata(case): + return _adata_from_columns("cond", case["cat_values"], case["columns"]) + + +def _case_kwargs(case): + kind, value = case["threshold"] + return {kind: value, "zero_to_na": case["zero_to_na"], "show": False} + + +def _reference(case): + categories = case["categories"] + cat_values = case["cat_values"] + names = list(case["columns"]) + kind, bound = case["threshold"] + + members = {} + for category in categories: + rows = [i for i, value in enumerate(cat_values) if value == category] + flags = [] + for name in names: + column = case["columns"][name] + n_detected = 0 + for i in rows: + entry = column[i] + if math.isnan(entry): + continue + if case["zero_to_na"] and entry == 0.0: + continue + n_detected += 1 + if not rows: + flags.append(False) + elif kind == "min_count": + flags.append(n_detected >= bound) + else: + flags.append(n_detected / len(rows) >= bound) + members[category] = flags + + counts = {} + for j in range(len(names)): + vector = tuple(members[category][j] for category in categories) + counts[vector] = counts.get(vector, 0) + 1 + counts.setdefault(tuple(False for _ in categories), 0) + return counts + + +def _mapping(series): + return {tuple(key): int(value) for key, value in series.items()} + + +def _last_series(spy, adata, case): + var_detected_by_cat_upset(adata, "cond", **_case_kwargs(case)) + return spy.init[-1]["data"] + + +def _dose_ho(): + return _adata("dose", [3, 20, 100], {"x1": {20: 1}}, 1.0) + + +def _dose_imp(): + return _adata("dose", [10, 2, 1], {"x1": {2: 1}}, 1.0) + + +class TestVarDetectedByCatUpset: + # -- Core intersection counts and thresholds + + def test_T1_full_detection_membership(self, spy): + series = _series(spy, _h1(), "organ", min_fraction=1.0) + _check_series( + series, + ["kidney", "lung", "spleen"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 2, + (F, F, F): 2, + }, + ) + + def test_T2_default_threshold_membership(self, spy): + series = _series(spy, _h1(), "organ") + _check_series( + series, + ["kidney", "lung", "spleen"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 2, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 2, + (F, F, F): 1, + }, + ) + + def test_T3_count_threshold_inclusive(self, spy): + series = _series(spy, _h3(), "site", min_count=3) + _check_series( + series, + ["north", "south"], + {(F, T): 3, (T, F): 3, (F, F): 0}, + ) + + def test_T4_fraction_threshold_per_category(self, spy): + series = _series(spy, _h4(), "arm", min_fraction=0.6) + _check_series( + series, + ["p", "q"], + {(T, F): 2, (F, T): 1, (T, T): 1, (F, F): 1}, + ) + + def test_T5_zero_counts_as_detected(self, spy): + series = _series(spy, _h5(), "arm", min_count=3) + _check_series( + series, + ["u", "v"], + {(T, F): 1, (F, T): 1, (F, F): 1}, + ) + + def test_T6_zero_to_na_hides_zeros(self, spy): + series = _series(spy, _h5(), "arm", zero_to_na=True) + _check_series(series, ["u", "v"], {(T, F): 2, (F, F): 1}) + + def test_T7_only_all_false_zero_entry(self, spy): + series = _series(spy, _h3(), "site", min_count=3) + assert len(series) == 3 + assert int(series[(F, F)]) == 0 + + def test_T8_category_without_observations_has_no_members(self, spy): + series = _series(spy, _h6(), "batch", min_fraction=0.0) + _check_series( + series, + ["zeta", "beta", "alpha", "mid"], + {(T, F, T, T): 3, (F, F, F, F): 0}, + ) + + def test_T9_imbalanced_category_sizes(self, spy): + adata = _adata( + "grp", + ["small"] * 2 + ["large"] * 20, + { + "r1": {"small": 2}, + "r2": {"small": 1, "large": 20}, + "r3": {"small": 2, "large": 2}, + "r4": {"large": 1}, + "r5": {"small": 2, "large": 19}, + }, + 1.0, + ) + series = _series(spy, adata, "grp", min_count=2) + _check_series( + series, + ["large", "small"], + {(F, T): 1, (T, F): 1, (T, T): 2, (F, F): 1}, + ) + + def test_T73_fraction_uses_division(self, spy): + adata = _adata( + "grp", + ["B"] * 25, + {"e1": {"B": 14}, "e2": {"B": 13}, "e3": {"B": 25}}, + 1.0, + ) + series = _series(spy, adata, "grp", min_fraction=0.56) + _check_series(series, ["B"], {(T,): 2, (F,): 1}) + + # -- Category order and names + + def test_T10_categorical_level_order(self, spy): + series = _series(spy, _h6(), "batch") + _check_series( + series, + ["zeta", "beta", "alpha", "mid"], + { + (T, F, T, T): 1, + (F, F, F, T): 1, + (F, F, F, F): 1, + }, + ) + + def test_T11_plain_level_order(self, spy): + series = _series(spy, _h7(), "fruit") + _check_series( + series, + ["apple", "fig", "kiwi"], + {(F, T, F): 1, (T, F, F): 1, (F, F, F): 0}, + ) + + def test_T12_plain_values_str_coerced(self, spy): + series = _series(spy, _dose_ho(), "dose") + _check_series( + series, + ["100", "20", "3"], + {(F, T, F): 1, (F, F, F): 0}, + ) + + def test_T13_categorical_values_str_coerced(self, spy): + adata = _adata( + "grade", + [100, 4, 30], + {"x1": {4: 1}}, + 1.0, + categories=[30, 4, 100], + ) + series = _series(spy, adata, "grade") + _check_series( + series, + ["30", "4", "100"], + {(F, T, F): 1, (F, F, F): 0}, + ) + + def test_T14_matrix_labels_from_categorical(self): + axes = var_detected_by_cat_upset(_h6(), "batch", show=False) + assert _matrix_labels(axes) == ["alpha", "beta", "mid", "zeta"] + + def test_T15_matrix_labels_str_coerced(self): + axes = var_detected_by_cat_upset(_dose_ho(), "dose", show=False) + assert _matrix_labels(axes) == ["100", "20", "3"] + + def test_T16_special_and_long_names(self, spy): + values = ["Größe µ", "a b c", "_underscore"] + adata = _adata( + "site", + values, + {"f1": {value: 1 for value in values}}, + 1.0, + ) + series = _series(spy, adata, "site") + _check_series( + series, + ["Größe µ", "_underscore", "a b c"], + {(T, T, T): 1, (F, F, F): 0}, + ) + + # -- Feature counting + + def test_T17_feature_ids_case_sensitive(self, spy): + adata = _adata( + "group", + ["A", "A", "B", "B"], + { + "Prot1": {"B": 2}, + "PROT1": {"B": 2}, + "prot1": {"B": 2}, + }, + 1.0, + ) + series = _series(spy, adata, "group") + _check_series(series, ["A", "B"], {(F, T): 3, (F, F): 0}) + + def test_T18_repeated_detections_counted_once(self, spy): + adata = _adata( + "group", + ["A"] * 4 + ["B"] * 4, + {"f1": {"A": 4, "B": 4}, "f2": {"A": 4, "B": 4}}, + 1.0, + ) + series = _series(spy, adata, "group") + _check_series(series, ["A", "B"], {(T, T): 2, (F, F): 0}) + + # -- Degenerate structures + + def test_T19_identical_membership_vectors(self, spy): + adata = _adata( + "lane", + ["x"] * 3 + ["y"] * 3, + { + "f1": {"x": 3, "y": 3}, + "f2": {"x": 3, "y": 3}, + "f3": {"x": 3, "y": 3}, + "f4": {}, + }, + 1.0, + ) + series = _series(spy, adata, "lane") + _check_series(series, ["x", "y"], {(T, T): 3, (F, F): 1}) + + def test_T20_category_without_members(self, spy): + adata = _adata( + "part", + ["h", "h", "m", "m", "t", "t"], + { + "e1": {"h": 2, "t": 2}, + "e2": {"t": 1}, + "e3": {}, + }, + 1.0, + ) + series = _series(spy, adata, "part") + _check_series( + series, + ["h", "m", "t"], + {(T, F, T): 1, (F, F, T): 1, (F, F, F): 1}, + ) + + def test_T21_single_category(self, spy): + adata = _adata( + "only", + ["solo"] * 2, + {"f1": {"solo": 2}, "f2": {"solo": 2}, "f3": {"solo": 2}}, + 1.0, + ) + axes = var_detected_by_cat_upset(adata, "only", show=False) + _check_series(spy.data, ["solo"], {(T,): 3, (F,): 0}) + assert set(axes) == { + "matrix", + "intersections", + "totals", + "shading", + } + + def test_T22_many_categories(self, spy): + categories = [f"k{i:02d}" for i in range(1, 13)] + adata = _adata( + "plate", + categories, + { + "v1": {categories[0]: 1, categories[1]: 1}, + "v2": {categories[11]: 1}, + "v3": {}, + }, + 1.0, + ) + series = _series(spy, adata, "plate") + _check_series( + series, + categories, + { + _vec(12, {0, 1}): 1, + _vec(12, {11}): 1, + _vec(12, set()): 1, + }, + ) + + # -- Sparse input and mutation + + def test_T23_sparse_warns_and_densifies(self, spy): + adata = _h1() + adata.X = sparse.csc_matrix(adata.X) + with pytest.warns(UserWarning): + series = _series(spy, adata, "organ", min_fraction=1.0) + assert not isinstance(series.dtype, pd.SparseDtype) + _check_series( + series, + ["kidney", "lung", "spleen"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 2, + (F, F, F): 2, + }, + ) + + def test_T24_input_not_mutated(self): + view = _h1()[:, :7] + before_x = np.array(view.X, dtype=float) + before_obs = view.obs.copy(deep=True) + before_var = view.var.copy(deep=True) + var_detected_by_cat_upset( + view, + "organ", + zero_to_na=True, + print_stats=True, + show=False, + ) + np.testing.assert_array_equal(np.array(view.X, dtype=float), before_x) + pd.testing.assert_frame_equal(view.obs, before_obs) + pd.testing.assert_frame_equal(view.var, before_var) + assert view.is_view + + def test_T25_sparse_arrays_not_mutated(self): + adata = _h1() + adata.X = sparse.csc_matrix(adata.X) + before = { + name: getattr(adata.X, name).copy() + for name in ("data", "indices", "indptr") + } + with pytest.warns(UserWarning): + var_detected_by_cat_upset(adata, "organ", show=False) + for name, values in before.items(): + np.testing.assert_array_equal(getattr(adata.X, name), values) + + # -- Plot construction + + def test_T27_upset_constructor_options(self, spy): + var_detected_by_cat_upset(_h1(), "organ", show=False) + assert len(spy.init) == 1 + call = spy.init[0] + assert call["subset_size"] == "sum" + assert call["sort_by"] == "degree" + assert call["sort_categories_by"] == "input" + assert call["show_counts"] is True + assert call["include_empty_subsets"] is False + + def test_T28_no_category_styling(self, spy): + var_detected_by_cat_upset(_h6(), "batch", show=False) + assert spy.style + call = spy.style[0] + assert set(call["absent"]) == {"zeta", "beta", "alpha", "mid"} + assert call["label"] == "No category" + + def test_T29_returns_plot_result(self, spy): + axes = var_detected_by_cat_upset(_h1(), "organ", show=False) + assert axes is spy.plot_returns[0] + + def test_T30_axes_keys_share_figure(self): + axes = var_detected_by_cat_upset(_h3(), "site", show=False) + assert { + "matrix", + "intersections", + "totals", + "shading", + } <= set(axes) + figure = axes["matrix"].figure + assert all(ax.figure is figure for ax in axes.values()) + + def test_T31_intersection_bar_heights(self): + axes = var_detected_by_cat_upset( + _h1(), "organ", min_fraction=1.0, show=False + ) + assert _bar_heights(axes) == [1, 1, 1, 1, 1, 2, 2, 2] + + def test_T32_totals_bar_widths(self): + axes = var_detected_by_cat_upset( + _h1(), "organ", min_fraction=1.0, show=False + ) + assert _bar_widths(axes) == [5, 5, 6] + + def test_T33_no_category_legend_entry(self): + axes = var_detected_by_cat_upset(_h3(), "site", show=False) + assert "No category" in _legend_texts(axes) + + # -- print_stats output + + def test_T34_stats_table_order(self, capsys): + var_detected_by_cat_upset(_h1(), "organ", print_stats=True, show=False) + out = capsys.readouterr().out + assert ( + out.index("Global:") + < out.index("Intersections:") + < out.index("Per organ:") + ) + + def test_T35_global_stats_table(self, capsys): + var_detected_by_cat_upset( + _h1(), + "organ", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, values = _global_table(capsys.readouterr().out) + assert header == [ + "count", + "mean", + "median", + "std", + "min", + "max", + ] + assert float(values[0]) == 8 + assert values[1] == "1.4" + assert float(values[2]) == 1 + assert values[3] == "0.5" + assert float(values[4]) == 1 + assert float(values[5]) == 2 + + def test_T36_global_std_single_entry(self, capsys): + adata = _adata( + "axis", + ["x", "y", "z"], + {f"f{i}": {} for i in range(1, 6)}, + 1.0, + ) + var_detected_by_cat_upset(adata, "axis", print_stats=True, show=False) + header, values = _global_table(capsys.readouterr().out) + assert header == [ + "count", + "mean", + "median", + "std", + "min", + "max", + ] + assert float(values[0]) == 1 + assert values[1] == "5.0" + assert float(values[2]) == 5 + assert math.isnan(float(values[3])) + assert float(values[4]) == 5 + assert float(values[5]) == 5 + + def test_T37_intersections_stats_table(self, capsys): + var_detected_by_cat_upset( + _h1(), + "organ", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, rows = _intersections_table(capsys.readouterr().out, 3) + assert header == [ + "kidney", + "lung", + "spleen", + "n_features", + "label", + ] + assert rows == [ + ["False", "False", "False", "2", "No category"], + ["True", "False", "False", "2", "kidney"], + ["True", "True", "True", "2", "kidney & lung & spleen"], + ["True", "True", "False", "1", "kidney & lung"], + ["True", "False", "True", "1", "kidney & spleen"], + ["False", "True", "False", "1", "lung"], + ["False", "True", "True", "1", "lung & spleen"], + ["False", "False", "True", "1", "spleen"], + ] + + def test_T38_per_category_stats_table(self, capsys): + var_detected_by_cat_upset( + _h1(), + "organ", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, rows = _per_table(capsys.readouterr().out, "organ") + assert header == ["organ", "n_features", "percent"] + assert rows == [ + ["kidney", "6", "54.5"], + ["lung", "5", "45.5"], + ["spleen", "5", "45.5"], + ] + + def test_T39_stats_printed_before_show(self, monkeypatch, capsys): + seen = {} + + def show(*args, **kwargs): + seen.setdefault("out", capsys.readouterr().out) + + monkeypatch.setattr(plt, "show", show) + var_detected_by_cat_upset(_h1(), "organ", print_stats=True, show=True) + assert "Global:" in seen["out"] + + def test_T40_labels_follow_category_order(self, capsys): + var_detected_by_cat_upset(_h6(), "batch", print_stats=True, show=False) + _, rows = _intersections_table(capsys.readouterr().out, 4) + labels = [row[-1] for row in rows] + assert "zeta & alpha & mid" in labels + + def test_T75_category_named_like_table_column(self, spy, capsys): + adata = _adata( + "site", + ["n_features"] * 3 + ["alpha"] * 3, + { + "e1": {"n_features": 3}, + "e2": {"alpha": 1, "n_features": 2}, + "e3": {"alpha": 3}, + "e4": {"alpha": 2}, + }, + 2.0, + ) + series = _series(spy, adata, "site", min_count=2, print_stats=True) + _check_series( + series, + ["alpha", "n_features"], + {(T, F): 2, (F, T): 2, (F, F): 0}, + ) + out = capsys.readouterr().out + header, rows = _intersections_table(out, 2) + assert header == ["alpha", "n_features", "n_features", "label"] + assert rows == [ + ["True", "False", "2", "alpha"], + ["False", "True", "2", "n_features"], + ["False", "False", "0", "No category"], + ] + header, rows = _per_table(out, "site") + assert header == ["site", "n_features", "percent"] + assert rows == [ + ["alpha", "2", "50.0"], + ["n_features", "2", "50.0"], + ] + + def test_T76_cat_key_named_like_table_column(self, spy, capsys): + adata = _adata( + "n_features", + ["x"] * 3 + ["y"] * 2, + { + "g1": {"x": 3, "y": 2}, + "g2": {"y": 1}, + "g3": {"x": 1}, + "g4": {}, + }, + 3.0, + ) + series = _series( + spy, adata, "n_features", min_fraction=0.5, print_stats=True + ) + _check_series( + series, + ["x", "y"], + {(T, T): 1, (F, T): 1, (F, F): 2}, + ) + header, rows = _per_table(capsys.readouterr().out, "n_features") + assert header == ["n_features", "n_features", "percent"] + assert rows == [ + ["x", "1", "25.0"], + ["y", "2", "50.0"], + ] + + # -- verbose output + + def test_T41_verbose_report_contents(self, capsys): + var_detected_by_cat_upset( + _h3(), + "site", + min_fraction=0.75, + verbose=True, + show=False, + ) + out = capsys.readouterr().out + for fragment in (".X", "site", "6", "2", "min_fraction", "0.75"): + assert fragment in out + + def test_T42_quiet_by_default(self, capsys, tmp_path): + var_detected_by_cat_upset( + _h1(), + "organ", + show=False, + save=str(tmp_path / "quiet.png"), + ) + assert capsys.readouterr().out == "" + + def test_T43_verbose_precedes_stats(self, capsys): + var_detected_by_cat_upset( + _h1(), + "organ", + verbose=True, + print_stats=True, + show=False, + ) + out = capsys.readouterr().out + assert out.index(".X") < out.index("Global:") + + # -- save and show + + def test_T44_save_writes_file(self, tmp_path): + target = tmp_path / "fig.pdf" + var_detected_by_cat_upset(_h1(), "organ", show=False, save=target) + assert target.exists() + assert target.stat().st_size > 0 + + def test_T45_show_calls_pyplot_show(self, show_calls): + var_detected_by_cat_upset(_h3(), "site", show=True) + assert len(show_calls) == 1 + + def test_T46_no_show_when_disabled(self, show_calls, tmp_path): + var_detected_by_cat_upset( + _h3(), + "site", + show=False, + save=str(tmp_path / "quiet.png"), + ) + assert show_calls == [] + + def test_T47_new_figure_per_call(self): + first = var_detected_by_cat_upset(_h1(), "organ", show=False) + second = var_detected_by_cat_upset(_h1(), "organ", show=False) + assert first["matrix"].figure is not second["matrix"].figure + + def test_T48_figure_left_open(self): + axes = var_detected_by_cat_upset(_h1(), "organ", show=False) + assert plt.fignum_exists(axes["matrix"].figure.number) + + # -- Threshold selection + + def test_T56_thresholds_none_explicitly(self, spy): + series = _series( + spy, _h1(), "organ", min_count=None, min_fraction=None + ) + _check_series( + series, + ["kidney", "lung", "spleen"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 2, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 2, + (F, F, F): 1, + }, + ) + + def test_T65_min_fraction_integer_accepted(self, spy): + series = _series(spy, _h1(), "organ", min_fraction=0) + _check_series( + series, + ["kidney", "lung", "spleen"], + {(T, T, T): 11, (F, F, F): 0}, + ) + + def test_T66_min_fraction_numpy_float_accepted(self, spy): + series = _series(spy, _h4(), "arm", min_fraction=np.float64(0.5)) + _check_series( + series, + ["p", "q"], + {(T, F): 2, (F, T): 1, (T, T): 1, (F, F): 1}, + ) + + # -- Negative cases + + def test_T50_save_invalid_type(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", show=False, save=["plot.png"] + ) + + def test_T51_save_unsupported_type(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_h1(), "organ", show=False, save=42) + + def test_T52_flag_not_bool(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_h1(), "organ", show="yes") + + def test_T53_flag_bool_like_rejected(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", show=False, zero_to_na=np.bool_(False) + ) + + def test_T54_flag_non_bool_value(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_h1(), "organ", show=0.5) + + def test_T55_both_thresholds_set(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _h1(), + "organ", + min_count=1, + min_fraction=1.0, + show=False, + ) + + def test_T57_min_count_not_int(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", min_count=np.int32(3), show=False + ) + + def test_T58_min_count_wrong_type(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", min_count=[1], show=False + ) + + def test_T59_min_count_bool_rejected(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", min_count=False, show=False + ) + + def test_T60_min_count_negative(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset(_h1(), "organ", min_count=-7, show=False) + + def test_T61_min_fraction_wrong_type(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), + "organ", + min_fraction=np.float16(0.25), + show=False, + ) + + def test_T62_min_fraction_bool_rejected(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _h1(), "organ", min_fraction=False, show=False + ) + + def test_T63_min_fraction_out_of_range(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _h1(), "organ", min_fraction=-0.01, show=False + ) + + def test_T64_min_fraction_not_finite(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _h1(), "organ", min_fraction=float("-inf"), show=False + ) + + def test_T74_fraction_beyond_float_range_rejected(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _h1(), "organ", min_fraction=-(2**1100), show=False + ) + + def test_T67_cat_key_missing(self): + with pytest.raises(KeyError): + var_detected_by_cat_upset(_h1(), "Organ", show=False) + + def test_T68_cat_key_with_missing_values(self): + adata = _h6() + adata.obs.loc["s1", "batch"] = np.nan + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "batch", show=False) + + def test_T69_empty_axis(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset(_h1()[:, :0], "organ", show=False) + + def test_T70_cat_key_not_str(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_h1(), ("organ",), show=False) + + def test_T71_cat_key_empty_string(self): + adata = _h1() + adata.obs[""] = ["valid"] * adata.n_obs + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "", show=False) + + def test_T72_category_name_collision(self): + adata = _adata("flag", [True, "True"], {"x1": {True: 1}}, 1.0) + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "flag", show=False) + + def test_T26_proteodata_validated(self): + adata = _h1() + X = np.array(adata.X, dtype=float) + X[0, 0] = np.inf + adata.X = X + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "organ", show=False) + + # -- Property relations + + def test_P1_series_matches_reference(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + series = _last_series(spy, _case_adata(case), case) + _check_series(series, case["categories"], _reference(case)) + plt.close("all") + + def test_P2_counts_sum_to_n_vars(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + adata = _case_adata(case) + series = _last_series(spy, adata, case) + values = [int(value) for value in series.to_numpy()] + assert all(value >= 0 for value in values) + assert sum(values) == adata.n_vars + plt.close("all") + + def test_P3_deterministic_series(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + first = _last_series(spy, _case_adata(case), case) + second = _last_series(spy, _case_adata(case), case) + pd.testing.assert_series_equal(first, second) + plt.close("all") + + def test_P4_sparse_matches_dense(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + dense = _last_series(spy, _case_adata(case), case) + adata = _case_adata(case) + adata.X = sparse.csr_matrix(adata.X) + with pytest.warns(UserWarning): + sparse_series = _last_series(spy, adata, case) + assert not isinstance(sparse_series.dtype, pd.SparseDtype) + assert _mapping(sparse_series) == _mapping(dense) + plt.close("all") + + def test_P5_input_unchanged(self, rng, report_input): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + adata = _case_adata(case) + before_x = np.array(adata.X, dtype=float) + before_obs = adata.obs.copy(deep=True) + before_var = adata.var.copy(deep=True) + var_detected_by_cat_upset(adata, "cond", **_case_kwargs(case)) + np.testing.assert_array_equal( + np.array(adata.X, dtype=float), before_x + ) + pd.testing.assert_frame_equal(adata.obs, before_obs) + pd.testing.assert_frame_equal(adata.var, before_var) + plt.close("all") + + # -- Metamorphic relations + + def test_M1_higher_count_threshold_shrinks_membership( + self, rng, report_input, spy + ): + for _ in range(_ROUNDS): + case = _generate(rng) + case["threshold"] = ("min_count", rng.randint(0, 4)) + report_input(_compact(case)) + lower = _last_series(spy, _case_adata(case), case) + raised = dict(case) + raised["threshold"] = ( + "min_count", + case["threshold"][1] + 1, + ) + higher = _last_series(spy, _case_adata(raised), raised) + low_counts = _per_category(lower, case["categories"]) + high_counts = _per_category(higher, case["categories"]) + for category in case["categories"]: + assert high_counts[category] <= low_counts[category] + plt.close("all") + + def test_M2_obs_permutation_invariant(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + original = _last_series(spy, _case_adata(case), case) + order = list(range(len(case["cat_values"]))) + rng.shuffle(order) + shuffled = dict(case) + shuffled["cat_values"] = [case["cat_values"][i] for i in order] + shuffled["columns"] = { + name: [column[i] for i in order] + for name, column in case["columns"].items() + } + permuted = _last_series(spy, _case_adata(shuffled), shuffled) + assert list(permuted.index.names) == list(original.index.names) + assert _mapping(permuted) == _mapping(original) + plt.close("all") + + def test_M3_var_permutation_invariant(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + report_input(_compact(case)) + original = _last_series(spy, _case_adata(case), case) + names = list(case["columns"]) + rng.shuffle(names) + shuffled = dict(case) + shuffled["columns"] = { + name: case["columns"][name] for name in names + } + permuted = _last_series(spy, _case_adata(shuffled), shuffled) + assert _mapping(permuted) == _mapping(original) + plt.close("all") + + def test_M4_detected_value_irrelevant(self, rng, report_input, spy): + for _ in range(_ROUNDS): + case = _generate(rng) + case["zero_to_na"] = False + report_input(_compact(case)) + original = _last_series(spy, _case_adata(case), case) + zeroed = dict(case) + zeroed["columns"] = { + name: [entry if math.isnan(entry) else 0.0 for entry in column] + for name, column in case["columns"].items() + } + replaced = _last_series(spy, _case_adata(zeroed), zeroed) + assert _mapping(replaced) == _mapping(original) + plt.close("all") + + +class TestVarDetectedByCatUpsetIMP: + # -- Core intersection counts and thresholds + + def test_T1_full_detection_membership_IMP(self, spy): + series = _series(spy, _f1(), "tissue", min_fraction=1.0) + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 1, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 2, + }, + ) + + def test_T2_default_threshold_membership_IMP(self, spy): + series = _series(spy, _f1(), "tissue") + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 1, + }, + ) + + def test_T3_count_threshold_inclusive_IMP(self, spy): + series = _series(spy, _f3(), "group", min_count=2) + _check_series( + series, + ["A", "B"], + {(F, T): 2, (T, T): 1, (T, F): 2, (F, F): 0}, + ) + + def test_T4_fraction_threshold_per_category_IMP(self, spy): + series = _series(spy, _f4(), "group", min_fraction=0.5) + _check_series( + series, + ["A", "B"], + {(T, F): 1, (F, T): 1, (T, T): 1, (F, F): 1}, + ) + + def test_T5_zero_counts_as_detected_IMP(self, spy): + series = _series(spy, _f5(), "group", min_count=2) + _check_series( + series, + ["A", "B"], + {(T, F): 1, (F, T): 1, (F, F): 0}, + ) + + def test_T6_zero_to_na_hides_zeros_IMP(self, spy): + series = _series(spy, _f5(), "group", zero_to_na=True) + _check_series(series, ["A", "B"], {(F, T): 1, (F, F): 1}) + + def test_T7_only_all_false_zero_entry_IMP(self, spy): + series = _series(spy, _f3(), "group", min_count=2) + assert len(series) == 4 + assert int(series[(F, F)]) == 0 + + def test_T8_category_without_observations_has_no_members_IMP(self, spy): + series = _series(spy, _f6(), "tissue", min_count=0) + _check_series( + series, + ["C", "A", "B", "D"], + {(T, T, T, F): 2, (F, F, F, F): 0}, + ) + + def test_T9_imbalanced_category_sizes_IMP(self, spy): + adata = _adata( + "grp", + ["A"] + ["B"] * 12, + { + "q1": {"A": 1, "B": 12}, + "q2": {"A": 1, "B": 11}, + "q3": {"B": 12}, + "q4": {}, + }, + 1.0, + ) + series = _series(spy, adata, "grp", min_fraction=1.0) + _check_series( + series, + ["A", "B"], + {(T, T): 1, (T, F): 1, (F, T): 1, (F, F): 1}, + ) + + def test_T73_fraction_uses_division_IMP(self, spy): + adata = _adata( + "grp", + ["A"] * 25, + {"f1": {"A": 7}, "f2": {"A": 6}}, + 1.0, + ) + series = _series(spy, adata, "grp", min_fraction=0.28) + _check_series(series, ["A"], {(T,): 1, (F,): 1}) + + # -- Category order and names + + def test_T10_categorical_level_order_IMP(self, spy): + series = _series(spy, _f6(), "tissue") + _check_series( + series, + ["C", "A", "B", "D"], + { + (T, T, T, F): 1, + (F, T, F, F): 1, + (F, F, F, F): 0, + }, + ) + + def test_T11_plain_level_order_IMP(self, spy): + series = _series(spy, _f7(), "tissue") + _check_series( + series, + ["A", "B", "C"], + {(F, T, F): 1, (F, F, F): 0}, + ) + + def test_T12_plain_values_str_coerced_IMP(self, spy): + series = _series(spy, _dose_imp(), "dose") + _check_series( + series, + ["1", "10", "2"], + {(F, F, T): 1, (F, F, F): 0}, + ) + + def test_T13_categorical_values_str_coerced_IMP(self, spy): + adata = _adata( + "grade", + [1, 2], + {"x1": {1: 1}}, + 1.0, + categories=[2, 1], + ) + series = _series(spy, adata, "grade") + _check_series(series, ["2", "1"], {(F, T): 1, (F, F): 0}) + + def test_T14_matrix_labels_from_categorical_IMP(self): + axes = var_detected_by_cat_upset(_f6(), "tissue", show=False) + assert _matrix_labels(axes) == ["A", "B", "C", "D"] + + def test_T15_matrix_labels_str_coerced_IMP(self): + axes = var_detected_by_cat_upset(_dose_imp(), "dose", show=False) + assert _matrix_labels(axes) == ["1", "10", "2"] + + def test_T16_special_and_long_names_IMP(self, spy): + values = ["tumor (T1)", "normal/adjacent", "L" * 60] + adata = _adata( + "site", + values, + {"f1": {value: 1 for value in values}}, + 1.0, + ) + series = _series(spy, adata, "site") + _check_series( + series, + ["L" * 60, "normal/adjacent", "tumor (T1)"], + {(T, T, T): 1, (F, F, F): 0}, + ) + + # -- Feature counting + + def test_T17_feature_ids_case_sensitive_IMP(self, spy): + adata = _adata( + "group", + ["A", "A", "B", "B"], + {"PEP": {"A": 2}, "pep": {"A": 2}}, + 1.0, + level="peptide", + protein_ids={"PEP": "PX", "pep": "PY"}, + ) + series = _series(spy, adata, "group") + _check_series(series, ["A", "B"], {(T, F): 2, (F, F): 0}) + + def test_T18_repeated_detections_counted_once_IMP(self, spy): + adata = _adata( + "group", + ["A"] * 5, + { + "f1": {"A": 5}, + "f2": {"A": 5}, + "f3": {"A": 5}, + }, + 1.0, + ) + series = _series(spy, adata, "group") + _check_series(series, ["A"], {(T,): 3, (F,): 0}) + + # -- Degenerate structures + + def test_T19_identical_membership_vectors_IMP(self, spy): + adata = _adata( + "group", + ["A", "A", "B", "B", "C", "C"], + {f"f{i}": {"A": 2, "B": 2, "C": 2} for i in range(1, 5)}, + 1.0, + ) + series = _series(spy, adata, "group") + _check_series( + series, + ["A", "B", "C"], + {(T, T, T): 4, (F, F, F): 0}, + ) + + def test_T20_category_without_members_IMP(self, spy): + adata = _adata( + "group", + ["A", "A", "B", "B"], + {"f1": {"A": 2}, "f2": {"A": 1}}, + 1.0, + ) + series = _series(spy, adata, "group") + _check_series(series, ["A", "B"], {(T, F): 2, (F, F): 0}) + + def test_T21_single_category_IMP(self, spy): + adata = _adata( + "group", + ["A"] * 3, + {"f1": {"A": 2}, "f2": {}}, + 1.0, + ) + axes = var_detected_by_cat_upset(adata, "group", show=False) + _check_series(spy.data, ["A"], {(T,): 1, (F,): 1}) + assert set(axes) == { + "matrix", + "intersections", + "totals", + "shading", + } + + def test_T22_many_categories_IMP(self, spy): + categories = [f"c{i:02d}" for i in range(1, 12)] + detections = {f"u{i + 1:02d}": {categories[i]: 1} for i in range(11)} + detections["u12"] = {category: 1 for category in categories} + adata = _adata("plate", categories, detections, 1.0) + series = _series(spy, adata, "plate") + expected = {_vec(11, {i}): 1 for i in range(11)} + expected[_vec(11, set(range(11)))] = 1 + expected[_vec(11, set())] = 0 + _check_series(series, categories, expected) + + # -- Sparse input and mutation + + def test_T23_sparse_warns_and_densifies_IMP(self, spy): + adata = _f1() + adata.X = sparse.csr_matrix(adata.X) + with pytest.warns(UserWarning): + series = _series(spy, adata, "tissue", min_fraction=1.0) + assert not isinstance(series.dtype, pd.SparseDtype) + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 1, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 2, + }, + ) + + def test_T24_input_not_mutated_IMP(self): + adata = _f1() + before_x = np.array(adata.X, dtype=float) + before_obs = adata.obs.copy(deep=True) + before_var = adata.var.copy(deep=True) + var_detected_by_cat_upset(adata, "tissue", zero_to_na=True, show=False) + np.testing.assert_array_equal(np.array(adata.X, dtype=float), before_x) + pd.testing.assert_frame_equal(adata.obs, before_obs) + pd.testing.assert_frame_equal(adata.var, before_var) + + def test_T25_sparse_arrays_not_mutated_IMP(self): + adata = _f1() + adata.X = sparse.csr_matrix(adata.X) + before = { + name: getattr(adata.X, name).copy() + for name in ("data", "indices", "indptr") + } + with pytest.warns(UserWarning): + var_detected_by_cat_upset(adata, "tissue", show=False) + for name, values in before.items(): + np.testing.assert_array_equal(getattr(adata.X, name), values) + + # -- Plot construction + + def test_T27_upset_constructor_options_IMP(self, spy): + var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert len(spy.init) == 1 + call = spy.init[0] + assert call["subset_size"] == "sum" + assert call["sort_by"] == "degree" + assert call["sort_categories_by"] == "input" + assert call["show_counts"] is True + assert call["include_empty_subsets"] is False + + def test_T28_no_category_styling_IMP(self, spy): + var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert spy.style + call = spy.style[0] + assert set(call["absent"]) == {"A", "B", "C"} + assert call["label"] == "No category" + + def test_T29_returns_plot_result_IMP(self, spy): + axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert axes is spy.plot_returns[0] + + def test_T30_axes_keys_share_figure_IMP(self): + axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert { + "matrix", + "intersections", + "totals", + "shading", + } <= set(axes) + figure = axes["matrix"].figure + assert all(ax.figure is figure for ax in axes.values()) + + def test_T31_intersection_bar_heights_IMP(self): + axes = var_detected_by_cat_upset( + _f1(), "tissue", min_fraction=1.0, show=False + ) + assert _bar_heights(axes) == [1, 1, 1, 1, 1, 1, 1, 2] + + def test_T32_totals_bar_widths_IMP(self): + axes = var_detected_by_cat_upset( + _f1(), "tissue", min_fraction=1.0, show=False + ) + assert _bar_widths(axes) == [4, 4, 4] + + def test_T33_no_category_legend_entry_IMP(self): + axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert "No category" in _legend_texts(axes) + + # -- print_stats output + + def test_T34_stats_table_order_IMP(self, capsys): + var_detected_by_cat_upset( + _f1(), "tissue", print_stats=True, show=False + ) + out = capsys.readouterr().out + assert ( + out.index("Global:") + < out.index("Intersections:") + < out.index("Per tissue:") + ) + + def test_T35_global_stats_table_IMP(self, capsys): + var_detected_by_cat_upset( + _f1(), + "tissue", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, values = _global_table(capsys.readouterr().out) + assert header == [ + "count", + "mean", + "median", + "std", + "min", + "max", + ] + assert float(values[0]) == 8 + assert values[1] == "1.1" + assert float(values[2]) == 1 + assert values[3] == "0.4" + assert float(values[4]) == 1 + assert float(values[5]) == 2 + + def test_T36_global_std_single_entry_IMP(self, capsys): + adata = _adata( + "group", + ["A", "A", "B", "B"], + {f"f{i}": {} for i in range(1, 4)}, + 1.0, + ) + var_detected_by_cat_upset(adata, "group", print_stats=True, show=False) + header, values = _global_table(capsys.readouterr().out) + assert header == [ + "count", + "mean", + "median", + "std", + "min", + "max", + ] + assert float(values[0]) == 1 + assert values[1] == "3.0" + assert float(values[2]) == 3 + assert math.isnan(float(values[3])) + assert float(values[4]) == 3 + assert float(values[5]) == 3 + + def test_T37_intersections_stats_table_IMP(self, capsys): + var_detected_by_cat_upset( + _f1(), + "tissue", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, rows = _intersections_table(capsys.readouterr().out, 3) + assert header == ["A", "B", "C", "n_features", "label"] + assert rows == [ + ["False", "False", "False", "2", "No category"], + ["True", "False", "False", "1", "A"], + ["True", "True", "False", "1", "A & B"], + ["True", "True", "True", "1", "A & B & C"], + ["True", "False", "True", "1", "A & C"], + ["False", "True", "False", "1", "B"], + ["False", "True", "True", "1", "B & C"], + ["False", "False", "True", "1", "C"], + ] + + def test_T38_per_category_stats_table_IMP(self, capsys): + var_detected_by_cat_upset( + _f1(), + "tissue", + min_fraction=1.0, + print_stats=True, + show=False, + ) + header, rows = _per_table(capsys.readouterr().out, "tissue") + assert header == ["tissue", "n_features", "percent"] + assert rows == [ + ["A", "4", "44.4"], + ["B", "4", "44.4"], + ["C", "4", "44.4"], + ] + + def test_T39_stats_printed_before_show_IMP(self, monkeypatch, capsys): + seen = {} + + def show(*args, **kwargs): + seen.setdefault("out", capsys.readouterr().out) + + monkeypatch.setattr(plt, "show", show) + var_detected_by_cat_upset(_f1(), "tissue", print_stats=True, show=True) + assert "Global:" in seen["out"] + + def test_T40_labels_follow_category_order_IMP(self, capsys): + var_detected_by_cat_upset( + _f6(), "tissue", print_stats=True, show=False + ) + _, rows = _intersections_table(capsys.readouterr().out, 4) + labels = [row[-1] for row in rows] + assert "C & A & B" in labels + + def test_T75_category_named_like_table_column_IMP(self, spy, capsys): + adata = _adata( + "grp", + ["label"] * 2 + ["B"] * 2, + { + "f1": {"label": 2, "B": 2}, + "f2": {"B": 1}, + "f3": {}, + }, + 1.0, + ) + series = _series(spy, adata, "grp", print_stats=True) + _check_series( + series, + ["B", "label"], + {(T, T): 1, (T, F): 1, (F, F): 1}, + ) + out = capsys.readouterr().out + header, rows = _intersections_table(out, 2) + assert header == ["B", "label", "n_features", "label"] + assert rows == [ + ["True", "False", "1", "B"], + ["True", "True", "1", "B & label"], + ["False", "False", "1", "No category"], + ] + header, rows = _per_table(out, "grp") + assert header == ["grp", "n_features", "percent"] + assert rows == [ + ["B", "2", "66.7"], + ["label", "1", "33.3"], + ] + + def test_T76_cat_key_named_like_table_column_IMP(self, capsys): + adata = _adata( + "percent", + ["A", "A", "B", "B"], + { + "f1": {"A": 2}, + "f2": {"A": 1, "B": 1}, + "f3": {}, + }, + 1.0, + ) + var_detected_by_cat_upset( + adata, "percent", print_stats=True, show=False + ) + header, rows = _per_table(capsys.readouterr().out, "percent") + assert header == ["percent", "n_features", "percent"] + assert rows == [ + ["A", "2", "66.7"], + ["B", "1", "33.3"], + ] + + # -- verbose output + + def test_T41_verbose_report_contents_IMP(self, capsys): + var_detected_by_cat_upset(_f1(), "tissue", verbose=True, show=False) + out = capsys.readouterr().out + for fragment in (".X", "tissue", "9", "3", "min_count", "1"): + assert fragment in out + + def test_T42_quiet_by_default_IMP(self, capsys): + var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert capsys.readouterr().out == "" + + def test_T43_verbose_precedes_stats_IMP(self, capsys): + var_detected_by_cat_upset( + _f1(), + "tissue", + verbose=True, + print_stats=True, + show=False, + ) + out = capsys.readouterr().out + assert out.index(".X") < out.index("Global:") + + # -- save and show + + def test_T44_save_writes_file_IMP(self, tmp_path): + target = tmp_path / "upset.png" + var_detected_by_cat_upset( + _f1(), "tissue", show=False, save=str(target) + ) + assert target.exists() + assert target.stat().st_size > 0 + + def test_T45_show_calls_pyplot_show_IMP(self, show_calls): + var_detected_by_cat_upset(_f1(), "tissue", show=True) + assert len(show_calls) == 1 + + def test_T46_no_show_when_disabled_IMP(self, show_calls): + var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert show_calls == [] + + def test_T47_new_figure_per_call_IMP(self): + first = var_detected_by_cat_upset(_f1(), "tissue", show=False) + second = var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert first["matrix"].figure is not second["matrix"].figure + + def test_T48_figure_left_open_IMP(self): + axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) + assert plt.fignum_exists(axes["matrix"].figure.number) + + # -- Threshold selection + + def test_T56_thresholds_none_explicitly_IMP(self, spy): + series = _series( + spy, _f1(), "tissue", min_count=None, min_fraction=None + ) + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 2, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 1, + }, + ) + + def test_T65_min_fraction_integer_accepted_IMP(self, spy): + series = _series(spy, _f1(), "tissue", min_fraction=1) + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 1, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 2, + }, + ) + + def test_T66_min_fraction_numpy_float_accepted_IMP(self, spy): + series = _series(spy, _f1(), "tissue", min_fraction=np.float64(1.0)) + _check_series( + series, + ["A", "B", "C"], + { + (T, F, F): 1, + (F, T, F): 1, + (F, F, T): 1, + (T, T, F): 1, + (T, F, T): 1, + (F, T, T): 1, + (T, T, T): 1, + (F, F, F): 2, + }, + ) + + # -- Negative cases + + def test_T50_save_invalid_type_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_f1(), "tissue", show=False, save=True) + + def test_T51_save_unsupported_type_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_f1(), "tissue", show=False, save=False) + + def test_T52_flag_not_bool_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", show=False, zero_to_na=1 + ) + + def test_T53_flag_bool_like_rejected_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), + "tissue", + show=False, + print_stats=np.bool_(True), + ) + + def test_T54_flag_non_bool_value_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", show=False, verbose=None + ) + + def test_T55_both_thresholds_set_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _f1(), + "tissue", + min_count=2, + min_fraction=0.5, + show=False, + ) + + def test_T57_min_count_not_int_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", min_count=2.0, show=False + ) + + def test_T58_min_count_wrong_type_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", min_count="2", show=False + ) + + def test_T59_min_count_bool_rejected_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", min_count=True, show=False + ) + + def test_T60_min_count_negative_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _f1(), "tissue", min_count=-1, show=False + ) + + def test_T61_min_fraction_wrong_type_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", min_fraction="0.5", show=False + ) + + def test_T62_min_fraction_bool_rejected_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset( + _f1(), "tissue", min_fraction=True, show=False + ) + + def test_T63_min_fraction_out_of_range_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _f1(), "tissue", min_fraction=1.5, show=False + ) + + def test_T64_min_fraction_not_finite_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _f1(), + "tissue", + min_fraction=float("nan"), + show=False, + ) + + def test_T74_fraction_beyond_float_range_rejected_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset( + _f1(), "tissue", min_fraction=10**400, show=False + ) + + def test_T67_cat_key_missing_IMP(self): + with pytest.raises(KeyError): + var_detected_by_cat_upset(_f1(), "condition", show=False) + + def test_T68_cat_key_with_missing_values_IMP(self): + adata = _f1() + adata.obs["tissue"] = adata.obs["tissue"].astype(object) + adata.obs.loc["s2", "tissue"] = np.nan + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "tissue", show=False) + + def test_T69_empty_axis_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset(_f1()[:0], "tissue", show=False) + + def test_T70_cat_key_not_str_IMP(self): + with pytest.raises(TypeError): + var_detected_by_cat_upset(_f1(), 5, show=False) + + def test_T71_cat_key_empty_string_IMP(self): + with pytest.raises(ValueError): + var_detected_by_cat_upset(_f1(), "", show=False) + + def test_T72_category_name_collision_IMP(self): + adata = _adata("flag", [1, "1"], {"x1": {1: 1}}, 1.0) + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "flag", show=False) + + def test_T26_proteodata_validated_IMP(self): + adata = _f1() + adata.obs = adata.obs.drop(columns=["sample_id"]) + with pytest.raises(ValueError): + var_detected_by_cat_upset(adata, "tissue", show=False) From 9a6947c1a8665c6d90f356a44f58a04e8ca36973 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:40:23 +0200 Subject: [PATCH 3/6] feature/var_detected_by_cat_upset Feature: pl.var_detected_by_cat_upset() Add an UpSet plot of which features (.var) are detected in which categories of an .obs column, e.g. which proteins are found in which tissues. - A feature is a member of a category when it is detected (non-NaN in .X; with zero_to_na=True also non-zero) in enough of that category's observations: min_count (default 1) or min_fraction, inclusive, per category. Setting both raises ValueError. - Features that are members of no category form a "No category" set, always included, also with a count of 0. - Categories follow the default pl order: category order for a Categorical column (unused categories kept), else lexicographic order of the str-coerced values. Values colliding after str coercion raise ValueError. - print_stats prints global, per-intersection and per-category tables; verbose reports the input matrix, threshold and sizes. - Sparse .X is densified with a UserWarning; adata is never modified. - Returns the axes dict of upsetplot's UpSet.plot(). Add upsetplot>=0.9.0,<0.10 as a core dependency. upsetplot 0.9.0 is broken on pandas 3 / numpy >= 2.4, so proteopy.pl.upset patches it on import (dot style defaults, single-category aggregation, count-label positions, all-empty totals warning). The upper bound keeps these patches from hitting an untested upsetplot release. Export via pr.pl, list it under Quality Control in the plotting API docs, and add a HISTORY.md entry. --- HISTORY.md | 4 + proteopy/pl/__init__.py | 4 + proteopy/pl/upset.py | 744 ++++++++++++++++++++++++++++++++++++++++ pyproject.toml | 1 + 4 files changed, 753 insertions(+) create mode 100644 proteopy/pl/upset.py diff --git a/HISTORY.md b/HISTORY.md index b139c35..a91f95f 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -14,6 +14,10 @@ and this project adheres to - `peptide_intensities()`, `proteoform_intensities()`: new `facet_by` parameter splitting the samples (`.obs`) across a grid of subplots +- `var_detected_by_cat_upset()`: UpSet plot of which features (`.var`) + are detected in which categories of an `.obs` column, with + per-category `min_count` / `min_fraction` detection thresholds and a + `No category` set for features detected in none of them **Preprocessing** (`pr.pp`) diff --git a/proteopy/pl/__init__.py b/proteopy/pl/__init__.py index 8f0fba2..50cb8ea 100644 --- a/proteopy/pl/__init__.py +++ b/proteopy/pl/__init__.py @@ -35,6 +35,10 @@ hclustv_profile_intensities, ) +from .upset import ( + var_detected_by_cat_upset, +) + from .sequence import ( peptides_on_sequence, peptides_on_prot_sequence, diff --git a/proteopy/pl/upset.py b/proteopy/pl/upset.py new file mode 100644 index 0000000..20eb227 --- /dev/null +++ b/proteopy/pl/upset.py @@ -0,0 +1,744 @@ +"""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, + names: list[str], +) -> 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, + ) + upset.style_subsets(absent=names, label=_NO_CATEGORY_LABEL) + return upset.plot() + + +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 + are shown as ``"No category"``. + + 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, names) + + if save is not None: + axes["matrix"].figure.savefig(save) + if show: + plt.show() + return axes diff --git a/pyproject.toml b/pyproject.toml index 269bbb0..f959f2b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,6 +46,7 @@ dependencies = [ "scipy", "seaborn", "statsmodels", + "upsetplot>=0.9.0,<0.10", ] [project.optional-dependencies] From 4d51cf712f0cf034e5807a4b5e27bd8c8b65af2b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Tue, 29 Sep 2026 01:09:35 +0200 Subject: [PATCH 4/6] fixed/removed legend from var_detected_by_cat_upset --- proteopy/pl/upset.py | 13 +++++-------- tests/pl/test_upset.py | 37 ++++++++++++++----------------------- 2 files changed, 19 insertions(+), 31 deletions(-) diff --git a/proteopy/pl/upset.py b/proteopy/pl/upset.py index 20eb227..f479a01 100644 --- a/proteopy/pl/upset.py +++ b/proteopy/pl/upset.py @@ -529,10 +529,7 @@ def _print_verbose_report( ) -def _plot_upset( - counts: pd.Series, - names: list[str], -) -> dict[str, Axes]: +def _plot_upset(counts: pd.Series) -> dict[str, Axes]: """Render the UpSet plot of the intersection counts.""" upset = UpSet( counts, @@ -542,7 +539,6 @@ def _plot_upset( show_counts=True, include_empty_subsets=False, ) - upset.style_subsets(absent=names, label=_NO_CATEGORY_LABEL) return upset.plot() @@ -566,8 +562,9 @@ def var_detected_by_cat_upset( 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 - are shown as ``"No category"``. + 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 @@ -735,7 +732,7 @@ def var_detected_by_cat_upset( ) # -- Plot - axes = _plot_upset(counts, names) + axes = _plot_upset(counts) if save is not None: axes["matrix"].figure.savefig(save) diff --git a/tests/pl/test_upset.py b/tests/pl/test_upset.py index 17c4d89..33be5ce 100644 --- a/tests/pl/test_upset.py +++ b/tests/pl/test_upset.py @@ -390,16 +390,13 @@ def _matrix_labels(axes): return sorted(text for text in texts if text) -def _legend_texts(axes): +def _legends(axes): figure = axes["matrix"].figure - texts = [] - for legend in figure.legends: - texts += [entry.get_text() for entry in legend.get_texts()] - for axis in axes.values(): - legend = axis.get_legend() - if legend is not None: - texts += [entry.get_text() for entry in legend.get_texts()] - return texts + legends = list(figure.legends) + for axis in figure.axes: + if axis.get_legend() is not None: + legends.append(axis.get_legend()) + return legends def _bar_heights(axes): @@ -917,12 +914,9 @@ def test_T27_upset_constructor_options(self, spy): assert call["show_counts"] is True assert call["include_empty_subsets"] is False - def test_T28_no_category_styling(self, spy): + def test_T28_no_subset_styling(self, spy): var_detected_by_cat_upset(_h6(), "batch", show=False) - assert spy.style - call = spy.style[0] - assert set(call["absent"]) == {"zeta", "beta", "alpha", "mid"} - assert call["label"] == "No category" + assert spy.style == [] def test_T29_returns_plot_result(self, spy): axes = var_detected_by_cat_upset(_h1(), "organ", show=False) @@ -951,9 +945,9 @@ def test_T32_totals_bar_widths(self): ) assert _bar_widths(axes) == [5, 5, 6] - def test_T33_no_category_legend_entry(self): + def test_T33_no_legend(self): axes = var_detected_by_cat_upset(_h3(), "site", show=False) - assert "No category" in _legend_texts(axes) + assert _legends(axes) == [] # -- print_stats output @@ -1801,12 +1795,9 @@ def test_T27_upset_constructor_options_IMP(self, spy): assert call["show_counts"] is True assert call["include_empty_subsets"] is False - def test_T28_no_category_styling_IMP(self, spy): + def test_T28_no_subset_styling_IMP(self, spy): var_detected_by_cat_upset(_f1(), "tissue", show=False) - assert spy.style - call = spy.style[0] - assert set(call["absent"]) == {"A", "B", "C"} - assert call["label"] == "No category" + assert spy.style == [] def test_T29_returns_plot_result_IMP(self, spy): axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) @@ -1835,9 +1826,9 @@ def test_T32_totals_bar_widths_IMP(self): ) assert _bar_widths(axes) == [4, 4, 4] - def test_T33_no_category_legend_entry_IMP(self): + def test_T33_no_legend_IMP(self): axes = var_detected_by_cat_upset(_f1(), "tissue", show=False) - assert "No category" in _legend_texts(axes) + assert _legends(axes) == [] # -- print_stats output From 8ef588747a3efeaf9d8a7f12250f76a786e3e710 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Wed, 7 Oct 2026 22:49:21 +0200 Subject: [PATCH 5/6] fix: docstrings and history --- HISTORY.md | 5 +- proteopy/pl/upset.py | 122 +++++++++++++++---------------------------- 2 files changed, 42 insertions(+), 85 deletions(-) diff --git a/HISTORY.md b/HISTORY.md index a91f95f..9f62240 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -14,10 +14,7 @@ and this project adheres to - `peptide_intensities()`, `proteoform_intensities()`: new `facet_by` parameter splitting the samples (`.obs`) across a grid of subplots -- `var_detected_by_cat_upset()`: UpSet plot of which features (`.var`) - are detected in which categories of an `.obs` column, with - per-category `min_count` / `min_fraction` detection thresholds and a - `No category` set for features detected in none of them +- `var_detected_by_cat_upset()`: new feature **Preprocessing** (`pr.pp`) diff --git a/proteopy/pl/upset.py b/proteopy/pl/upset.py index f479a01..92be000 100644 --- a/proteopy/pl/upset.py +++ b/proteopy/pl/upset.py @@ -553,23 +553,12 @@ def var_detected_by_cat_upset( 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. + """Plot feature detection overlaps across observation categories. + + A feature belongs to a category when it meets ``min_count`` or + ``min_fraction``. Missing values do not count as detections; zeros + count unless ``zero_to_na=True``. Features belonging to no category + form the ``"No category"`` intersection. Parameters ---------- @@ -578,20 +567,21 @@ def var_detected_by_cat_upset( read from ``adata.X`` only. cat_key : str Column in ``adata.obs`` defining the categories. + Categories follow categorical order, otherwise lexicographic + order after conversion to strings. Use an ordered + :class:`pandas.Categorical` to customize the order. 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. + Minimum detections per category, as a nonnegative integer. + If neither threshold is supplied, one detection suffices. 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``. + Minimum fraction of observations with a detection per category, + between 0 and 1. Supply at most one of ``min_count`` and + ``min_fraction``. zero_to_na : bool If True, zeros in ``.X`` count as missing. print_stats : bool - If True, print the statistics underlying the plot. + If True, print global, intersection, and per-category + statistics. verbose : bool If True, print status messages about the input. show : bool @@ -609,85 +599,55 @@ def var_detected_by_cat_upset( 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. + If an argument has the wrong type. KeyError - When ``cat_key`` is not a column of ``adata.obs``. + If ``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. + If thresholds are invalid or both supplied, ``cat_key`` is + empty, either data axis is empty, or category labels are + missing or collide after conversion to strings. Warns ----- UserWarning - When ``adata.X`` is sparse and is densified. + If ``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. + P2 is measured in one liver sample; P3 is never measured. >>> 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, + >>> samples = ["S1", "S2", "S3", "S4"] + >>> proteins = ["P1", "P2", "P3"] + >>> adata = ad.AnnData( + ... X=np.array([ + ... [5.0, 2.0, np.nan], [4.0, np.nan, np.nan], + ... [6.0, np.nan, np.nan], [5.0, np.nan, np.nan], + ... ]), + ... obs=pd.DataFrame( + ... {"sample_id": samples, + ... "tissue": ["liver", "liver", "lung", "lung"]}, + ... index=samples, + ... ), + ... var=pd.DataFrame({"protein_id": proteins}, index=proteins), + ... ) + >>> axes = pr.pl.var_detected_by_cat_upset( + ... adata, cat_key="tissue", show=False, ... ) - >>> 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. + Require detection in every sample of a tissue: >>> axes = pr.pl.var_detected_by_cat_upset( ... adata, ... cat_key="tissue", ... min_fraction=1.0, - ... print_stats=True, + ... show=False, ... ) - 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( From a4164b5aaf051a0ca967aee9e236fd4b3a03d89d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Levente=20Temesv=C3=A1ri-Nagy?= <147416790+leventetn@users.noreply.github.com> Date: Thu, 8 Oct 2026 16:34:56 +0200 Subject: [PATCH 6/6] fix: Implemented sort_by for var_detected_by_cat_upset Implemented sort_by with "degree", "cardinality", "-degree" and "-cardinality". New tests cover all modes, equal counts and degrees, repeatability, bar/dot alignment, and invalid inputs. --- proteopy/pl/upset.py | 34 ++++- tests/pl/test_upset.py | 309 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 339 insertions(+), 4 deletions(-) diff --git a/proteopy/pl/upset.py b/proteopy/pl/upset.py index 92be000..78c3ba9 100644 --- a/proteopy/pl/upset.py +++ b/proteopy/pl/upset.py @@ -279,10 +279,21 @@ def _validate_upset_args( verbose, show, save, + sort_by, ) -> tuple[str, int | float]: """Validate all inputs and return the active threshold.""" check_proteodata(adata) + if not isinstance(sort_by, str): + raise TypeError( + f"`sort_by` must be a str, got {type(sort_by).__name__}." + ) + if sort_by not in ("degree", "cardinality", "-degree", "-cardinality"): + raise ValueError( + "`sort_by` must be one of 'degree', 'cardinality', " + "'-degree', '-cardinality'." + ) + for value, name in ( (zero_to_na, "zero_to_na"), (print_stats, "print_stats"), @@ -529,12 +540,12 @@ def _print_verbose_report( ) -def _plot_upset(counts: pd.Series) -> dict[str, Axes]: +def _plot_upset(counts: pd.Series, sort_by: str) -> dict[str, Axes]: """Render the UpSet plot of the intersection counts.""" upset = UpSet( counts, subset_size="sum", - sort_by="degree", + sort_by=sort_by, sort_categories_by="input", show_counts=True, include_empty_subsets=False, @@ -552,6 +563,8 @@ def var_detected_by_cat_upset( verbose: bool = False, show: bool = True, save: str | Path | None = None, + *, + sort_by: str = "degree", ) -> dict[str, Axes]: """Plot feature detection overlaps across observation categories. @@ -588,6 +601,11 @@ def var_detected_by_cat_upset( Call ``plt.show()`` at the end. save : str | Path | None Path to save the figure to; None skips saving. + sort_by : str + Order intersection bars by ``"degree"`` (fewest overlapping + categories first) or ``"cardinality"`` (largest feature counts + first). ``"-degree"`` and ``"-cardinality"`` reverse these + orders. Ties follow the plotting library's ordering. Returns ------- @@ -605,7 +623,8 @@ def var_detected_by_cat_upset( ValueError If thresholds are invalid or both supplied, ``cat_key`` is empty, either data axis is empty, or category labels are - missing or collide after conversion to strings. + missing or collide after conversion to strings, or ``sort_by`` + is unsupported. Warns ----- @@ -648,6 +667,12 @@ def var_detected_by_cat_upset( ... min_fraction=1.0, ... show=False, ... ) + + Show the largest intersections first: + + >>> axes = pr.pl.var_detected_by_cat_upset( + ... adata, cat_key="tissue", sort_by="cardinality", show=False, + ... ) """ # -- Validate inputs threshold_name, threshold_value = _validate_upset_args( @@ -660,6 +685,7 @@ def var_detected_by_cat_upset( verbose, show, save, + sort_by, ) # -- Derive categories and feature memberships @@ -692,7 +718,7 @@ def var_detected_by_cat_upset( ) # -- Plot - axes = _plot_upset(counts) + axes = _plot_upset(counts, sort_by) if save is not None: axes["matrix"].figure.savefig(save) diff --git a/tests/pl/test_upset.py b/tests/pl/test_upset.py index 33be5ce..c20d42a 100644 --- a/tests/pl/test_upset.py +++ b/tests/pl/test_upset.py @@ -320,6 +320,74 @@ def show_calls(monkeypatch): plt.close("all") +@pytest.fixture +def sorting_adata(): + """Return a builder for specified intersection counts.""" + + def build(counts: dict, categories=None) -> ad.AnnData: + columns = {} + for members, count in counts.items(): + for _ in range(count): + columns[f"p{len(columns)}"] = [ + 1.0 if category in members else np.nan + for category in ["C", "A", "B"] + ] + samples = ["s1", "s2", "s3"] + proteins = list(columns) + obs = pd.DataFrame( + {"sample_id": samples, "group": ["C", "A", "B"]}, + index=samples, + ) + if categories is not None: + obs["group"] = pd.Categorical( + obs["group"], categories=categories, ordered=True + ) + return ad.AnnData( + X=np.asarray(list(columns.values()), dtype=float).T, + obs=obs, + var=pd.DataFrame({"protein_id": proteins}, index=proteins), + ) + + return build + + +@pytest.fixture +def rendered_intersections(): + """Return a reader of memberships and counts from drawn artists.""" + + def read(axes: dict) -> list: + matrix = axes["matrix"] + matrix.figure.canvas.draw() + labels = { + int(tick): label.get_text() + for tick, label in zip( + matrix.get_yticks(), matrix.get_yticklabels() + ) + } + bars = sorted( + axes["intersections"].patches, + key=lambda patch: patch.get_x(), + ) + dots = matrix.collections[0] + offsets = dots.get_offsets() + colors = dots.get_facecolors() + assert offsets.shape == (len(bars) * len(labels), 2) + assert colors.shape == (len(offsets), 4) + rows = [] + for bar in bars: + center = bar.get_x() + bar.get_width() / 2 + members = [ + labels[int(y)] + for (x, y), color in zip(offsets, colors) + if np.isclose(x, center) + and np.allclose(color, bar.get_facecolor()) + ] + rows.append((tuple(sorted(members)), bar.get_height())) + return rows + + return read + + @pytest.fixture def spy(monkeypatch): recorder = _Spy() @@ -575,6 +643,122 @@ def _dose_imp(): class TestVarDetectedByCatUpset: + # ── Intersection sorting ──────────────────────────────────────── + + @pytest.mark.parametrize( + "sort_by, expected", + [ + ("degree", [1, 4, 2, 3]), + ("-degree", [3, 2, 4, 1]), + ("cardinality", [4, 3, 2, 1]), + ("-cardinality", [1, 2, 3, 4]), + ], + ) + def test_T77_sort_by_orders_rendered_intersections( + self, sorting_adata, rendered_intersections, sort_by, expected + ): + counts = {(): 1, ("A",): 4, ("A", "B"): 2, ("A", "B", "C"): 3} + axes = var_detected_by_cat_upset( + sorting_adata(counts), "group", sort_by=sort_by, show=False + ) + rows = rendered_intersections(axes) + assert len(rows) == len(counts) + np.testing.assert_array_equal([count for _, count in rows], expected) + assert dict(rows) == counts + + def test_T78_sort_by_default_matches_explicit_degree( + self, sorting_adata, rendered_intersections + ): + adata = sorting_adata({(): 1, ("B",): 4, ("A",): 2, ("A", "B"): 3}) + default = var_detected_by_cat_upset(adata, "group", show=False) + explicit = var_detected_by_cat_upset( + adata, "group", sort_by="degree", show=False + ) + assert rendered_intersections(default) == rendered_intersections( + explicit + ) + + @pytest.mark.parametrize( + "sort_by", ["degree", "-degree", "cardinality", "-cardinality"] + ) + @pytest.mark.parametrize( + "categories, expected", + [(None, ["A", "B", "C"]), (["C", "A", "B"], ["C", "A", "B"])], + ids=["lexicographic", "categorical"], + ) + def test_T79_sort_by_preserves_category_order( + self, sorting_adata, sort_by, categories, expected + ): + adata = sorting_adata({(): 1, ("A",): 2, ("B", "C"): 3}, categories) + axes = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + labels = [t.get_text() for t in axes["matrix"].get_yticklabels()] + assert labels == expected + + @pytest.mark.parametrize( + "sort_by", ["degree", "-degree", "cardinality", "-cardinality"] + ) + @pytest.mark.parametrize( + "counts", + [ + {(): 1, ("A",): 3, ("A", "B"): 3, ("A", "B", "C"): 2}, + {(): 1, ("A",): 2, ("B",): 4, ("C",): 3, ("A", "B"): 5}, + { + (): 1, + ("A",): 2, + ("B",): 2, + ("A", "B"): 4, + ("A", "C"): 4, + ("A", "B", "C"): 3, + }, + {(): 2, ("A",): 2, ("B",): 2, ("A", "B"): 2}, + ], + ids=["equal-counts", "equal-degrees", "equal-both", "all-equal"], + ) + def test_T80_sort_by_ties_preserve_order_alignment_and_repeatability( + self, sorting_adata, rendered_intersections, counts, sort_by + ): + adata = sorting_adata(counts) + first = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + rows = rendered_intersections(first) + assert len(rows) == len(counts) + assert dict(rows) == counts + if "degree" in sort_by: + metric = [len(members) for members, _ in rows] + descending = sort_by.startswith("-") + else: + metric = [count for _, count in rows] + descending = not sort_by.startswith("-") + assert metric == sorted(metric, reverse=descending) + second = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + assert rows == rendered_intersections(second) + + @pytest.mark.parametrize("sort_by", [None, True, 1, [], {}]) + def test_T81_sort_by_rejects_non_strings_before_plotting(self, sort_by): + before = set(plt.get_fignums()) + with pytest.raises(TypeError, match="`sort_by` must be a str"): + var_detected_by_cat_upset( + _h1(), "organ", sort_by=sort_by, show=False + ) + assert set(plt.get_fignums()) == before + + @pytest.mark.parametrize("sort_by", ["", "size", "input", "Degree"]) + def test_T82_sort_by_rejects_unknown_modes_before_plotting(self, sort_by): + before = set(plt.get_fignums()) + with pytest.raises( + ValueError, + match="`sort_by` must be one of .*degree.*cardinality", + ): + var_detected_by_cat_upset( + _h1(), "organ", sort_by=sort_by, show=False + ) + assert set(plt.get_fignums()) == before + # -- Core intersection counts and thresholds def test_T1_full_detection_membership(self, spy): @@ -1487,6 +1671,131 @@ def test_M4_detected_value_irrelevant(self, rng, report_input, spy): class TestVarDetectedByCatUpsetIMP: + # ── Intersection sorting ──────────────────────────────────────── + + @pytest.mark.parametrize( + "sort_by, expected", + [ + ("degree", [8, 11, 9, 10]), + ("-degree", [10, 9, 11, 8]), + ("cardinality", [11, 10, 9, 8]), + ("-cardinality", [8, 9, 10, 11]), + ], + ) + def test_T77_sort_by_orders_rendered_intersections_IMP( + self, sorting_adata, rendered_intersections, sort_by, expected + ): + counts = {(): 8, ("A",): 11, ("A", "B"): 9, ("A", "B", "C"): 10} + axes = var_detected_by_cat_upset( + sorting_adata(counts), "group", sort_by=sort_by, show=False + ) + rows = rendered_intersections(axes) + assert len(rows) == len(counts) + np.testing.assert_array_equal([count for _, count in rows], expected) + assert dict(rows) == counts + + def test_T78_sort_by_default_matches_explicit_degree_IMP( + self, sorting_adata, rendered_intersections + ): + adata = sorting_adata({(): 5, ("B",): 8, ("A",): 6, ("A", "B"): 7}) + default = var_detected_by_cat_upset(adata, "group", show=False) + explicit = var_detected_by_cat_upset( + adata, "group", sort_by="degree", show=False + ) + assert rendered_intersections(default) == rendered_intersections( + explicit + ) + + @pytest.mark.parametrize( + "sort_by", ["degree", "-degree", "cardinality", "-cardinality"] + ) + @pytest.mark.parametrize( + "categories, expected", + [(None, ["A", "B", "C"]), (["C", "A", "B"], ["C", "A", "B"])], + ids=["lexicographic", "categorical"], + ) + def test_T79_sort_by_preserves_category_order_IMP( + self, sorting_adata, sort_by, categories, expected + ): + adata = sorting_adata({(): 4, ("A",): 5, ("B", "C"): 6}, categories) + axes = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + labels = [t.get_text() for t in axes["matrix"].get_yticklabels()] + assert labels == expected + + @pytest.mark.parametrize( + "sort_by", ["degree", "-degree", "cardinality", "-cardinality"] + ) + @pytest.mark.parametrize( + "counts", + [ + {(): 1, ("A",): 3, ("A", "B"): 3, ("A", "B", "C"): 2}, + {(): 1, ("A",): 2, ("B",): 4, ("C",): 3, ("A", "B"): 5}, + { + (): 1, + ("A",): 2, + ("B",): 2, + ("A", "B"): 4, + ("A", "C"): 4, + ("A", "B", "C"): 3, + }, + {(): 2, ("A",): 2, ("B",): 2, ("A", "B"): 2}, + ], + ids=["equal-counts", "equal-degrees", "equal-both", "all-equal"], + ) + def test_T80_sort_by_ties_preserve_order_alignment_and_repeatability_IMP( + self, sorting_adata, rendered_intersections, counts, sort_by + ): + counts = {members: count + 7 for members, count in counts.items()} + adata = sorting_adata(counts) + first = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + rows = rendered_intersections(first) + assert len(rows) == len(counts) + assert dict(rows) == counts + if "degree" in sort_by: + metric = [len(members) for members, _ in rows] + descending = sort_by.startswith("-") + else: + metric = [count for _, count in rows] + descending = not sort_by.startswith("-") + assert metric == sorted(metric, reverse=descending) + second = var_detected_by_cat_upset( + adata, "group", sort_by=sort_by, show=False + ) + assert rows == rendered_intersections(second) + + @pytest.mark.parametrize( + "sort_by", [False, 2.5, (), np.array([1]), object()] + ) + def test_T81_sort_by_rejects_non_strings_before_plotting_IMP( + self, sort_by + ): + before = set(plt.get_fignums()) + with pytest.raises(TypeError, match="`sort_by` must be a str"): + var_detected_by_cat_upset( + _f1(), "tissue", sort_by=sort_by, show=False + ) + assert set(plt.get_fignums()) == before + + @pytest.mark.parametrize( + "sort_by", ["count", "-input", "degree ", "CARDINALITY"] + ) + def test_T82_sort_by_rejects_unknown_modes_before_plotting_IMP( + self, sort_by + ): + before = set(plt.get_fignums()) + with pytest.raises( + ValueError, + match="`sort_by` must be one of .*degree.*cardinality", + ): + var_detected_by_cat_upset( + _f1(), "tissue", sort_by=sort_by, show=False + ) + assert set(plt.get_fignums()) == before + # -- Core intersection counts and thresholds def test_T1_full_detection_membership_IMP(self, spy):