From dc1ea56c86fbd184c930cc41d63889205cd72b06 Mon Sep 17 00:00:00 2001 From: Stephan Breimann Date: Fri, 18 Sep 2026 19:25:05 +0200 Subject: [PATCH] feat(prediction): audit an evaluation setup for leakage risks Leakage inflates a score without ever raising an error, and on small windowed datasets it is easy to introduce and hard to see. aa.audit_leakage runs a set of cheap heuristics over whatever parts of a setup are supplied and returns a plain DataFrame of findings (check, severity, detail, ids), worst first, with the overall verdict in df_audit.attrs["status"] so no bespoke result type is needed. Without splits it audits the dataset, which is the check to run before choosing a split; with splits it also compares the training against the test part within each fold. It reports duplicate sequences, a row index, protein or group present on both sides of a fold, a feature correlating with the label at |r| >= 0.95, an empty or strongly uneven fold, and a single-class or skewed fold. Windowed input is handled correctly: the sampler emits both the repeated parent 'sequence' and the 'window' cut from it, so the window takes precedence and sibling windows of one protein are not mistaken for duplicate samples. raise_on turns findings at or above a severity into a bare ValueError naming them; the default reports only and never raises, so an audit can be dropped into a workflow without changing its control flow. The checks are heuristic and a clean report is not a proof that no leakage exists; the severities are labels for a human reader, not a machine taxonomy, since whether a finding must block a workflow is a policy decision left to the caller. Pairs with bind_groups, which prevents the group leak this reports: a split from bind_groups(...).split(X, y) feeds straight into splits=. Refs #479 Co-Authored-By: Claude Fable 5.1 --- CHANGELOG.md | 13 + aaanalysis/__init__.py | 3 +- aaanalysis/_constants.py | 27 + aaanalysis/prediction/__init__.py | 9 +- aaanalysis/prediction/_audit_leakage.py | 271 ++++++ .../prediction/_backend/audit_leakage.py | 285 +++++++ docs/source/api.rst | 1 + docs/source/index/references.rst | 4 + examples/prediction/audit_leakage.ipynb | 781 ++++++++++++++++++ .../prediction_tests/test_audit_leakage.py | 466 +++++++++++ 10 files changed, 1857 insertions(+), 3 deletions(-) create mode 100644 aaanalysis/prediction/_audit_leakage.py create mode 100644 aaanalysis/prediction/_backend/audit_leakage.py create mode 100644 examples/prediction/audit_leakage.ipynb create mode 100644 tests/unit/prediction_tests/test_audit_leakage.py diff --git a/CHANGELOG.md b/CHANGELOG.md index f4dcca68..0fcd7ec4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,6 +26,19 @@ notes — with cross-references and examples — live in rejected at construction with both counts in the message. Once the folds are consumed, `df_folds_` reports per-fold sample counts, group counts and class balance. The package performs the split; it runs no homology search or clustering itself (#478). +- `audit_leakage(df_seq=None, *, labels=None, groups=None, splits=None, X=None, names=None, + raise_on=None)`: inspects an evaluation setup for leakage risks and returns a plain `DataFrame` + of findings (`check`, `severity`, `detail`, `ids`), worst first, with the overall verdict in + `df_audit.attrs["status"]`. Without `splits` it audits the dataset, which is the check to run + before choosing a split; with `splits` it also compares the training against the test part + within each fold. It reports duplicate sequences, a row index, protein or group present on both + sides of a fold, a feature correlating with the label at `|r| >= 0.95`, an empty or strongly + uneven fold, and a single-class or skewed fold. Windowed input is handled correctly: `window` + takes precedence over the repeated parent `sequence`, so siblings are not mistaken for + duplicates. `raise_on` turns findings at or above a severity into a `ValueError`; the default + reports only and never raises. The checks are heuristic and a clean report is not a proof that + no leakage exists, and the severities are labels for a human reader, not a machine taxonomy + (#479). - `Docs Build (gate)` CI workflow (`.github/scripts/check_docs_build.py`): builds the documentation and fails on a docutils `ERROR`, so a broken reference cannot reach `master` silently. It parses the build log rather than Sphinx's exit code, which is `0` even with diff --git a/aaanalysis/__init__.py b/aaanalysis/__init__.py index dc7c1e35..e74bac40 100644 --- a/aaanalysis/__init__.py +++ b/aaanalysis/__init__.py @@ -9,7 +9,7 @@ from .pu_learning import dPULearn, dPULearnPlot from .explainable_ai import TreeModel from .prediction import (AAPred, AAPredPlot, ReliabilityModel, ReliabilityModelPlot, - ModelEvaluator, ModelEvaluatorPlot, bind_groups) + ModelEvaluator, ModelEvaluatorPlot, bind_groups, audit_leakage) from .protein_engineering import (DesignConstraints, AAMut, AAMutPlot, SeqMut, SeqMutPlot, SeqOpt, SeqOptPlot) from .plotting import (plot_get_clist, plot_get_cmap, plot_get_cdict, @@ -74,6 +74,7 @@ "ModelEvaluator", "ModelEvaluatorPlot", "bind_groups", + "audit_leakage", # "ShapModel" # SHAP "plot_get_clist", "plot_get_cmap", diff --git a/aaanalysis/_constants.py b/aaanalysis/_constants.py index 43563c92..7947f84d 100644 --- a/aaanalysis/_constants.py +++ b/aaanalysis/_constants.py @@ -639,6 +639,33 @@ def _folder_path(super_folder, folder_name): COLS_FOLDS_GROUPS = [COL_FOLD, COL_N_TRAIN, COL_N_TEST, COL_N_GROUPS_TRAIN, COL_N_GROUPS_TEST, COL_POS_RATE_TRAIN, COL_POS_RATE_TEST] +# audit_leakage (heuristic leakage diagnostic): one row per finding, worst severity in .attrs +COL_CHECK = "check" # name of the heuristic that produced the finding +COL_SEVERITY = "severity" # human-readable label (not a machine taxonomy): low, medium, high +COL_DETAIL = "detail" # one-sentence, human-readable description of the finding +COL_IDS = "ids" # affected identifiers (list, capped; the true count is in detail) +COLS_AUDIT_LEAKAGE = [COL_CHECK, COL_SEVERITY, COL_DETAIL, COL_IDS] +# Severities, ordered from least to most severe. A label for a human reader, deliberately NOT a +# stable machine code: deciding whether a finding blocks a workflow is decision-layer policy. +STR_SEVERITY_LOW = "low" +STR_SEVERITY_MEDIUM = "medium" +STR_SEVERITY_HIGH = "high" +LIST_SEVERITIES = [STR_SEVERITY_LOW, STR_SEVERITY_MEDIUM, STR_SEVERITY_HIGH] +STR_STATUS_OK = "ok" # df.attrs["status"] when no finding was made +# Check names, i.e. the values of the 'check' column +STR_CHECK_DUPLICATE_SEQ = "duplicate_sequences" # identical sequence strings in df_seq +STR_CHECK_TRAIN_TEST_OVERLAP = "train_test_overlap" # a row index in both parts of a fold +STR_CHECK_DUPLICATE_SEQ_FOLDS = "duplicate_sequences_across_folds" # identical sequence split apart +STR_CHECK_ENTRY_FOLDS = "same_protein_across_folds" # windows of one protein split apart +STR_CHECK_GROUP_FOLDS = "group_overlap_across_folds" # a group id in both parts of a fold +STR_CHECK_TARGET_LEAK = "target_derived_feature" # a feature near-perfectly tracking the label +STR_CHECK_FOLD_SIZE = "fold_size_anomaly" # a test fold far from the average size +STR_CHECK_CLASS_BALANCE = "class_balance_anomaly" # a fold's class balance far from the whole +LIST_CHECKS_LEAKAGE = [STR_CHECK_DUPLICATE_SEQ, STR_CHECK_TRAIN_TEST_OVERLAP, + STR_CHECK_DUPLICATE_SEQ_FOLDS, STR_CHECK_ENTRY_FOLDS, + STR_CHECK_GROUP_FOLDS, STR_CHECK_TARGET_LEAK, + STR_CHECK_FOLD_SIZE, STR_CHECK_CLASS_BALANCE] + # Labels LABEL_FEAT_VAL = "Feature value" LABEL_HIST_COUNT = "Number of proteins" diff --git a/aaanalysis/prediction/__init__.py b/aaanalysis/prediction/__init__.py index 8879379f..6a52e177 100644 --- a/aaanalysis/prediction/__init__.py +++ b/aaanalysis/prediction/__init__.py @@ -2,7 +2,7 @@ Prediction: evaluate and deploy sequence-based prediction models. Public objects: AAPred, AAPredPlot, ReliabilityModel, ReliabilityModelPlot, ModelEvaluator, -ModelEvaluatorPlot, bind_groups. +ModelEvaluatorPlot, bind_groups, audit_leakage. Downstream of feature engineering (``CPP`` / ``CPPGrid`` produce ``df_feat`` and the feature matrix ``X``): ``AAPred`` evaluates one or more scikit-learn models across metrics by cross-validation and an optional held-out set (``eval``), and fits them for deployment, then @@ -13,7 +13,10 @@ comparison with a signed delta and a Wilcoxon significance test (``eval``) — visualized by ``ModelEvaluatorPlot``; ``bind_groups`` binds group labels (protein accession, family, or an externally computed homology cluster) to any scikit-learn splitter, so dependent samples stay -within one fold of ``eval`` / ``run`` instead of leaking across them. Complements ``explainable_ai.TreeModel`` (tree-ensemble feature +within one fold of ``eval`` / ``run`` instead of leaking across them, while ``audit_leakage`` +inspects a dataset and its folds after the fact and reports the leakage risks it can see +(duplicate sequences, a protein or group split across a fold, a feature tracking the label, a +skewed fold) as a plain table of findings. Complements ``explainable_ai.TreeModel`` (tree-ensemble feature importance) — this subpackage owns the general evaluate-and-deploy path. See ``.claude/rules/code-conventions.md`` for conventions and ``CONTEXT.md`` for domain terms. @@ -25,6 +28,7 @@ from ._model_evaluator import ModelEvaluator from ._model_evaluator_plot import ModelEvaluatorPlot from ._bind_groups import bind_groups +from ._audit_leakage import audit_leakage __all__ = [ "AAPred", @@ -34,4 +38,5 @@ "ModelEvaluator", "ModelEvaluatorPlot", "bind_groups", + "audit_leakage", ] diff --git a/aaanalysis/prediction/_audit_leakage.py b/aaanalysis/prediction/_audit_leakage.py new file mode 100644 index 00000000..e82d87a2 --- /dev/null +++ b/aaanalysis/prediction/_audit_leakage.py @@ -0,0 +1,271 @@ +""" +This is a script for the frontend of the audit_leakage function for diagnosing leakage +risks in an evaluation setup. +""" +from typing import Optional, Union, List, Tuple, Iterable +import numpy as np +import pandas as pd + +import aaanalysis.utils as ut +from ._backend.audit_leakage import audit_leakage_, comp_status_, get_col_seq_ + + +# I Helper Functions +def check_groups(groups=None): + """Check the group labels and return them as a 1D array, or None.""" + if groups is None: + return None + groups = ut.check_array_like(name="groups", val=groups, accept_none=False) + groups = np.asarray(groups) + if groups.ndim != 1: + raise ValueError(f"'groups' (n dimensions={groups.ndim}) should be one-dimensional") + if len(groups) == 0: + raise ValueError("'groups' (length=0) should contain at least one group label") + return groups + + +def check_labels_audit(labels=None): + """Check the labels permissively, since a single-class setup is itself a finding.""" + if labels is None: + return None + labels = ut.check_array_like(name="labels", val=labels, accept_none=False) + labels = np.asarray(labels) + if labels.ndim != 1: + raise ValueError(f"'labels' (n dimensions={labels.ndim}) should be one-dimensional") + if len(labels) == 0: + raise ValueError("'labels' (length=0) should contain at least one label") + return labels + + +def check_splits(splits=None): + """Check the splits and return them as a list of (train_idx, test_idx) index arrays.""" + if splits is None: + return None + try: + splits = list(splits) + except TypeError: + raise ValueError(f"'splits' ({type(splits).__name__}) should be an iterable of " + f"(train_idx, test_idx) pairs, e.g. 'list(cv.split(X, labels))'") + if len(splits) == 0: + raise ValueError("'splits' (n folds=0) should contain at least one fold") + list_splits = [] + for fold, split in enumerate(splits): + if not isinstance(split, (tuple, list)) or len(split) != 2: + raise ValueError(f"'splits' should contain (train_idx, test_idx) pairs, but fold " + f"{fold} ({split}) is not a pair of two index arrays") + idx = [] + for name, val in zip(["train_idx", "test_idx"], split): + arr = np.asarray(val) + if arr.ndim != 1: + raise ValueError(f"'{name}' of fold {fold} (n dimensions={arr.ndim}) should be " + f"one-dimensional") + if len(arr) > 0 and not np.issubdtype(arr.dtype, np.integer): + raise ValueError(f"'{name}' of fold {fold} (dtype={arr.dtype}) should contain " + f"integer row positions, not labels or a boolean mask") + idx.append(arr.astype(int)) + list_splits.append((idx[0], idx[1])) + return list_splits + + +def check_names(names=None, n_features=None): + """Check the feature names and return them as a list of strings, or None.""" + if names is None: + return None + names = ut.check_list_like(name="names", val=names, accept_none=False, accept_str=False) + if n_features is not None and len(names) != n_features: + raise ValueError(f"'names' (length={len(names)}) should have one name per feature in 'X' " + f"(n features={n_features})") + return [str(n) for n in names] + + +def check_anything_to_audit(df_seq=None, groups=None, splits=None, X=None): + """Check that at least one input was given, since otherwise there is nothing to inspect.""" + if df_seq is None and groups is None and splits is None and X is None: + raise ValueError("'df_seq', 'groups', 'splits' and 'X' are all None; at least one should " + "be given, since there is otherwise nothing to audit") + + +def check_df_seq_audit(df_seq=None): + """Check the sequence frame permissively, since the audit inspects whatever columns exist.""" + if df_seq is None: + return None + ut.check_df(name="df_seq", df=df_seq, accept_none=False) + if len(df_seq) == 0: + raise ValueError("'df_seq' (n rows=0) should contain at least one sequence") + if get_col_seq_(df_seq=df_seq) is None and ut.COL_ENTRY not in df_seq: + raise ValueError(f"'df_seq' (columns={list(df_seq)}) should contain at least one of " + f"'{ut.COL_SEQ}', '{ut.COL_WINDOW}' (the sample sequence) or " + f"'{ut.COL_ENTRY}' (the protein it comes from)") + return df_seq + + +def check_match_splits_n_samples(splits=None, n_samples=None, name=None): + """Check that the row positions of every fold are within the given number of samples.""" + if splits is None or n_samples is None: + return None + for fold, (train_idx, test_idx) in enumerate(splits): + for str_part, idx in zip(["train_idx", "test_idx"], [train_idx, test_idx]): + if len(idx) == 0: + continue + if idx.min() < 0 or idx.max() >= n_samples: + raise ValueError(f"'{str_part}' of fold {fold} (min={idx.min()}, max={idx.max()}) " + f"should contain row positions within '{name}' " + f"(n samples={n_samples})") + + +def check_match_n_samples(df_seq=None, labels=None, groups=None, X=None): + """Check that every given input describes the same number of samples.""" + dict_n = {} + if df_seq is not None: + dict_n["df_seq"] = len(df_seq) + if labels is not None: + dict_n["labels"] = len(labels) + if groups is not None: + dict_n["groups"] = len(groups) + if X is not None: + dict_n["X"] = len(X) + if len(set(dict_n.values())) > 1: + str_n = ", ".join(f"'{k}' (n samples={v})" for k, v in dict_n.items()) + raise ValueError(f"{str_n} should all describe the same samples, but their lengths differ") + + +def raise_on_findings(df_audit: pd.DataFrame, raise_on=None): + """Raise when a finding reaches the severity the caller chose to hard-fail on.""" + if raise_on is None or len(df_audit) == 0: + return None + th = ut.LIST_SEVERITIES.index(raise_on) + mask = [ut.LIST_SEVERITIES.index(s) >= th for s in df_audit[ut.COL_SEVERITY]] + df_hit = df_audit[mask] + if len(df_hit) == 0: + return None + str_checks = "; ".join(f"{r[ut.COL_CHECK]} ({r[ut.COL_SEVERITY]}): {r[ut.COL_DETAIL]}" + for _, r in df_hit.iterrows()) + raise ValueError(f"'raise_on' ('{raise_on}') matched {len(df_hit)} leakage finding(s) at or " + f"above that severity: {str_checks}") + + +# II Main Functions +def audit_leakage(df_seq: Optional[pd.DataFrame] = None, + *, + labels: Optional[ut.ArrayLike1D] = None, + groups: Optional[Union[ut.ArrayLike1D, pd.Series]] = None, + splits: Optional[Iterable[Tuple[ut.ArrayLike1D, ut.ArrayLike1D]]] = None, + X: Optional[ut.ArrayLike2D] = None, + names: Optional[List[str]] = None, + raise_on: Optional[str] = None, + ) -> pd.DataFrame: + """ + Audit an evaluation setup for leakage risks and report them as a table of findings. + + Leakage inflates a score without ever raising an error: a duplicated sequence on both sides + of a split, two windows of one protein pulled apart, a feature quietly derived from the + label. On the small, windowed datasets typical of sequence-based prediction these are easy + to introduce and hard to see, and the only symptom is an implausibly good number. This + function runs a set of cheap heuristics over whatever parts of the setup are supplied and + lists what looks wrong, before a score is trusted [Kaufman12]_. + + It works in two stages, each optional. Without ``splits`` only the dataset itself is + inspected, which is the check to run *before* choosing a split. With ``splits`` the folds + are inspected too, comparing the training against the test part **within each fold**. + + .. versionadded:: 1.2.0 + + Parameters + ---------- + df_seq : pd.DataFrame, shape (n_samples, n_features), optional + DataFrame containing an ``entry`` column with protein identifiers, one row per sample. + An identifier repeats when several windows come from one protein, which is what the + per-protein check looks for. ``'sequence'``, or ``'window'`` when present, is compared + for duplicates; for windowed output the window is the sample, so it takes precedence + over the parent ``'sequence'``. + labels : array-like, shape (n_samples,), optional + Class labels per sample. Used for the class-balance and target-derived checks. Taken + from ``df_seq['label']`` when that column exists and ``labels`` is not given. Unlike + elsewhere in this package, a single class is accepted, since it is itself reported. + groups : array-like, shape (n_samples,), optional + Group label per sample, as passed to :func:`bind_groups`: a protein accession, a family + name, or an externally computed homology cluster. A group in both parts of a fold is a + leak. + splits : iterable of (train_idx, test_idx), optional + The folds to inspect, as integer **row positions**, e.g. ``list(cv.split(X, labels))`` + or the output of a splitter returned by :func:`bind_groups`. Without it only the + dataset-level checks run. + X : array-like, shape (n_samples, n_features), optional + Feature matrix to inspect for columns that track the label too closely. Requires + ``labels``. + names : list of str, optional + Feature name per column of ``X``, e.g. ``df_feat['feature'].to_list()``, so a finding + names the feature instead of its column position. + raise_on : str, optional + Severity at or above which a finding raises a ``ValueError`` instead of being reported: + one of ``'low'``, ``'medium'`` or ``'high'``. The default reports only and never raises, + so an audit can be added to a workflow without changing its control flow. + + Returns + ------- + df_audit : pd.DataFrame + Findings, one row per issue, worst first, with the columns ``'check'`` (which heuristic + fired), ``'severity'`` (``'low'``, ``'medium'`` or ``'high'``), ``'detail'`` (a sentence + describing the finding and why it matters) and ``'ids'`` (the affected identifiers, at + most 10 of them; the true count is in ``'detail'``). A clean setup gives an empty frame + with those columns. ``df_audit.attrs['status']`` carries the worst severity present, or + ``'ok'`` when there is no finding. + + Raises + ------ + ValueError + If the given inputs do not describe the same samples, if ``splits`` holds anything other + than pairs of integer row positions within range, or if ``raise_on`` is set and a finding + reaches that severity. + + Notes + ----- + * **The audit is heuristic, and a clean report is not a proof that no leakage exists.** It + inspects only the data and the folds it is given. It cannot see how a feature was built, + whether a scaler was fitted before the split, or that two sequences are near-identical + rather than byte-identical. Treat ``status='ok'`` as "nothing obvious", not as a guarantee. + * The severities are labels for a human reader, not a machine taxonomy: whether a finding + must block a workflow is a policy decision this package deliberately leaves to the caller. + * Thresholds are fixed so a report means the same thing everywhere: a feature correlates + with the label at :math:`|r| \\ge 0.95`, the largest test fold holds twice the rows of the + smallest, or a fold's class share departs from the whole dataset by more than 0.2. + * Row positions, not index labels, are expected in ``splits``, matching scikit-learn. A + ``df_seq`` with a non-default index is therefore audited by position. + + See Also + -------- + * :func:`bind_groups` to *prevent* the group leak this reports, by keeping every group + within one fold. + * :class:`AAPred` and :class:`ModelEvaluator`, whose ``cv`` argument takes the splitter + whose folds are audited here. + + Examples + -------- + .. include:: examples/audit_leakage.rst + """ + # Validate + df_seq = check_df_seq_audit(df_seq=df_seq) + if labels is None and df_seq is not None and ut.COL_LABEL in df_seq: + labels = df_seq[ut.COL_LABEL].to_numpy() + labels = check_labels_audit(labels=labels) + groups = check_groups(groups=groups) + splits = check_splits(splits=splits) + names_x = list(X.columns) if isinstance(X, pd.DataFrame) else None + X = ut.check_X(X, accept_none=True, min_n_samples=1, min_n_features=1) + names = check_names(names=names, n_features=None if X is None else X.shape[1]) + ut.check_str_options(name="raise_on", val=raise_on, accept_none=True, + list_str_options=ut.LIST_SEVERITIES) + check_anything_to_audit(df_seq=df_seq, groups=groups, splits=splits, X=X) + check_match_n_samples(df_seq=df_seq, labels=labels, groups=groups, X=X) + for name, val in zip(["df_seq", "labels", "groups", "X"], [df_seq, labels, groups, X]): + if val is not None: + check_match_splits_n_samples(splits=splits, n_samples=len(val), name=name) + if X is not None and labels is None: + raise ValueError("'labels' should be given together with 'X', since a feature can only " + "be checked against the target it might be derived from") + # Run the heuristics + df_audit = audit_leakage_(df_seq=df_seq, labels=labels, groups=groups, splits=splits, X=X, + names=names if names is not None else names_x) + df_audit.attrs["status"] = comp_status_(df_audit=df_audit) + raise_on_findings(df_audit=df_audit, raise_on=raise_on) + return df_audit diff --git a/aaanalysis/prediction/_backend/audit_leakage.py b/aaanalysis/prediction/_backend/audit_leakage.py new file mode 100644 index 00000000..b150b101 --- /dev/null +++ b/aaanalysis/prediction/_backend/audit_leakage.py @@ -0,0 +1,285 @@ +""" +This is a script for the backend of the audit_leakage function, i.e. the individual +leakage heuristics and the assembly of their findings into one table. +""" +from typing import List, Optional +import numpy as np +import pandas as pd + +import aaanalysis.utils as ut + + +# Heuristic thresholds. Deliberately fixed rather than exposed: they mark the point where a +# pattern is worth a human's attention, not a decision boundary the caller should tune. +TH_TARGET_CORR = 0.95 # |correlation| of a feature with the label above which it looks derived +TH_SIZE_RATIO = 2.0 # largest/smallest test fold above which the folds are called uneven +TH_BALANCE_DEV = 0.2 # absolute deviation of a fold's class share from the whole dataset +N_IDS_MAX = 10 # affected ids kept per finding (the true count stays in the detail text) + + +# I Helper Functions +def _format_ids(ids) -> List[str]: + """Render affected identifiers as a short, capped list of strings.""" + return [str(i) for i in list(ids)[:N_IDS_MAX]] + + +def _format_folds(folds) -> str: + """Render the affected fold indices as a readable enumeration.""" + return ", ".join(str(f) for f in folds) + + +def _finding(check=None, severity=None, detail=None, ids=None) -> list: + """Build one positional row of the findings frame, ordered as ut.COLS_AUDIT_LEAKAGE.""" + return [check, severity, detail, _format_ids(ids if ids is not None else [])] + + +def get_col_seq_(df_seq=None) -> Optional[str]: + """Obtain the column holding the sequence of the *sample*, or None when there is none. + + Windowed output carries both the parent protein sequence and the window cut from it, so the + window is the sample; comparing the parent instead would call every window of one protein a + duplicate of its siblings. + """ + if df_seq is None: + return None + for col in [ut.COL_WINDOW, ut.COL_SEQ]: + if col in df_seq: + return col + return None + + +def _get_ids(df_seq=None, n_rows=0) -> np.ndarray: + """Obtain the sample identifiers of a sequence frame, falling back to the positional index.""" + if df_seq is not None: + for col in [ut.COL_ENTRY_WIN, ut.COL_ENTRY]: + if col in df_seq: + return np.asarray(df_seq[col]) + return np.arange(n_rows) + + +def _comp_overlap(values=None, train_idx=None, test_idx=None) -> np.ndarray: + """Obtain the values shared by the training and the test part of one fold.""" + return np.intersect1d(np.unique(values[train_idx]), np.unique(values[test_idx])) + + +def _scan_folds(splits=None, values=None): + """Collect, over all folds, the values appearing in both parts and the folds they appear in.""" + shared_all, folds = [], [] + for fold, (train_idx, test_idx) in enumerate(splits): + shared = _comp_overlap(values=values, train_idx=train_idx, test_idx=test_idx) + if len(shared) > 0: + shared_all.extend(shared.tolist()) + folds.append(fold) + return sorted(set(shared_all), key=str), folds + + +def _comp_corr_with_labels(X=None, labels=None) -> np.ndarray: + """Compute the absolute Pearson correlation of every feature with the label vector.""" + y = np.asarray(labels, dtype=float) + x_cent = X - X.mean(axis=0) + y_cent = y - y.mean() + denom = np.sqrt((x_cent ** 2).sum(axis=0) * (y_cent ** 2).sum()) + corr = np.zeros(X.shape[1], dtype=float) + valid = denom > 0 + corr[valid] = (x_cent[:, valid] * y_cent[:, None]).sum(axis=0) / denom[valid] + return np.abs(corr) + + +def _comp_class_shares(labels=None, idx=None, classes=None) -> np.ndarray: + """Compute the share of every class within one part of a fold.""" + part = labels[idx] + if len(part) == 0: + return np.zeros(len(classes), dtype=float) + return np.array([np.mean(part == c) for c in classes], dtype=float) + + +# II Main Functions +def check_duplicate_seq_(df_seq=None) -> List[list]: + """Flag identical sequence strings, which become a leak as soon as the data is split.""" + col_seq = get_col_seq_(df_seq=df_seq) + if col_seq is None: + return [] + seqs = df_seq[col_seq].astype(str) + is_dup = seqs.duplicated(keep=False) + if not is_dup.any(): + return [] + ids = _get_ids(df_seq=df_seq, n_rows=len(df_seq)) + n_rows, n_seqs = int(is_dup.sum()), int(seqs[is_dup].nunique()) + detail = (f"{n_rows} rows share only {n_seqs} distinct '{col_seq}' value(s). Duplicates that " + f"end up on both sides of a split are scored on what the model already memorized") + return [_finding(check=ut.STR_CHECK_DUPLICATE_SEQ, severity=ut.STR_SEVERITY_MEDIUM, + detail=detail, ids=ids[np.asarray(is_dup)])] + + +def check_train_test_overlap_(splits=None) -> List[list]: + """Flag a row index used for both training and testing in the same fold.""" + shared_all, folds = [], [] + for fold, (train_idx, test_idx) in enumerate(splits): + shared = np.intersect1d(np.asarray(train_idx), np.asarray(test_idx)) + if len(shared) > 0: + shared_all.extend(shared.tolist()) + folds.append(fold) + if len(folds) == 0: + return [] + ids = sorted(set(shared_all)) + detail = (f"{len(ids)} row(s) are in both the training and the test part of fold(s) " + f"{_format_folds(folds)}, so the model is scored on rows it was fitted on") + return [_finding(check=ut.STR_CHECK_TRAIN_TEST_OVERLAP, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=ids)] + + +def check_duplicate_seq_folds_(df_seq=None, splits=None) -> List[list]: + """Flag an identical sequence whose copies are split across the two parts of a fold.""" + col_seq = get_col_seq_(df_seq=df_seq) + if col_seq is None: + return [] + seqs = np.asarray(df_seq[col_seq].astype(str)) + shared, folds = _scan_folds(splits=splits, values=seqs) + if len(folds) == 0: + return [] + detail = (f"{len(shared)} '{col_seq}' value(s) occur in both the training and the test part " + f"of fold(s) {_format_folds(folds)}, so an identical copy was seen during training") + # Report the affected samples, not the sequence strings, which are too long to read + ids = _get_ids(df_seq=df_seq, n_rows=len(df_seq))[np.isin(seqs, list(shared))] + return [_finding(check=ut.STR_CHECK_DUPLICATE_SEQ_FOLDS, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=ids)] + + +def check_entry_folds_(df_seq=None, splits=None) -> List[list]: + """Flag windows of one protein that are split across the two parts of a fold.""" + if df_seq is None or ut.COL_ENTRY not in df_seq: + return [] + entries = np.asarray(df_seq[ut.COL_ENTRY]) + # Only proteins contributing several rows can have their windows split apart; with one row + # per protein this would merely restate a train/test row overlap. + if not pd.Series(entries).duplicated().any(): + return [] + shared, folds = _scan_folds(splits=splits, values=entries) + if len(folds) == 0: + return [] + detail = (f"{len(shared)} protein(s) have windows in both the training and the test part of " + f"fold(s) {_format_folds(folds)}, so neighbouring windows leak across the split") + return [_finding(check=ut.STR_CHECK_ENTRY_FOLDS, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=shared)] + + +def check_group_folds_(groups=None, splits=None) -> List[list]: + """Flag a group label appearing in both parts of a fold.""" + if groups is None: + return [] + shared, folds = _scan_folds(splits=splits, values=groups) + if len(folds) == 0: + return [] + detail = (f"{len(shared)} group(s) appear in both the training and the test part of fold(s) " + f"{_format_folds(folds)}; bind the groups to a group-aware splitter to prevent it") + return [_finding(check=ut.STR_CHECK_GROUP_FOLDS, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=shared)] + + +def check_target_leak_(X=None, labels=None, names=None) -> List[list]: + """Flag a feature column that tracks the label closely enough to look derived from it.""" + if X is None or labels is None: + return [] + corr = _comp_corr_with_labels(X=X, labels=labels) + is_leaky = corr >= TH_TARGET_CORR + if not is_leaky.any(): + return [] + idx = np.flatnonzero(is_leaky) + ids = [names[i] for i in idx] if names is not None else [f"column {i}" for i in idx] + detail = (f"{len(idx)} feature(s) correlate with the label at |r| >= {TH_TARGET_CORR} " + f"(max {corr[idx].max():.3f}), which is what a feature derived from the target " + f"looks like. Confirm each is computed without the label") + return [_finding(check=ut.STR_CHECK_TARGET_LEAK, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=ids)] + + +def check_fold_size_(splits=None) -> List[list]: + """Flag an empty fold part, and test folds whose sizes are strongly uneven.""" + n_train = np.array([len(train_idx) for train_idx, _ in splits]) + n_test = np.array([len(test_idx) for _, test_idx in splits]) + empty = np.flatnonzero((n_train == 0) | (n_test == 0)) + if len(empty) > 0: + detail = (f"fold(s) {_format_folds(empty.tolist())} have an empty training or test part, " + f"so no score from them describes the model") + return [_finding(check=ut.STR_CHECK_FOLD_SIZE, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=empty.tolist())] + if len(n_test) < 2 or n_test.min() == 0: + return [] + ratio = n_test.max() / n_test.min() + if ratio < TH_SIZE_RATIO: + return [] + extremes = [int(np.argmin(n_test)), int(np.argmax(n_test))] + detail = (f"the largest test fold holds {ratio:.1f}x the rows of the smallest " + f"({n_test.max()} vs {n_test.min()}), so the folds are not comparable and the mean " + f"score is dominated by the large one. Group splitters do this legitimately") + return [_finding(check=ut.STR_CHECK_FOLD_SIZE, severity=ut.STR_SEVERITY_LOW, + detail=detail, ids=extremes)] + + +def check_class_balance_(labels=None, splits=None) -> List[list]: + """Flag a single-class test fold, and folds whose class balance departs from the dataset.""" + if labels is None: + return [] + labels = np.asarray(labels) + classes = np.unique(labels) + shares_all = _comp_class_shares(labels=labels, idx=np.arange(len(labels)), classes=classes) + single, skewed, devs = [], [], [] + for fold, (_, test_idx) in enumerate(splits): + test_idx = np.asarray(test_idx) + if len(test_idx) == 0: + continue + if len(np.unique(labels[test_idx])) < 2 and len(classes) > 1: + single.append(fold) + continue + shares = _comp_class_shares(labels=labels, idx=test_idx, classes=classes) + dev = float(np.abs(shares - shares_all).max()) + if dev > TH_BALANCE_DEV: + skewed.append(fold) + devs.append(dev) + findings = [] + if len(single) > 0: + detail = (f"the test part of fold(s) {_format_folds(single)} holds a single class, so " + f"every class-aware metric there is undefined or trivially satisfied") + findings.append(_finding(check=ut.STR_CHECK_CLASS_BALANCE, severity=ut.STR_SEVERITY_HIGH, + detail=detail, ids=single)) + if len(skewed) > 0: + detail = (f"the class balance of fold(s) {_format_folds(skewed)} departs from the whole " + f"dataset by up to {max(devs):.2f}, so their scores are not comparable. " + f"A stratified splitter keeps the balance") + findings.append(_finding(check=ut.STR_CHECK_CLASS_BALANCE, + severity=ut.STR_SEVERITY_MEDIUM, detail=detail, ids=skewed)) + return findings + + +def comp_status_(df_audit: pd.DataFrame) -> str: + """Summarise the findings as the worst severity present, or 'ok' when there are none.""" + if len(df_audit) == 0: + return ut.STR_STATUS_OK + severities = set(df_audit[ut.COL_SEVERITY]) + for severity in reversed(ut.LIST_SEVERITIES): + if severity in severities: + return severity + return ut.STR_STATUS_OK + + +def audit_leakage_(df_seq=None, labels=None, groups=None, splits=None, X=None, + names: Optional[list] = None) -> pd.DataFrame: + """Run every applicable heuristic and assemble the findings into one table.""" + rows = [] + rows.extend(check_duplicate_seq_(df_seq=df_seq)) + rows.extend(check_target_leak_(X=X, labels=labels, names=names)) + if splits is not None: + rows.extend(check_train_test_overlap_(splits=splits)) + rows.extend(check_duplicate_seq_folds_(df_seq=df_seq, splits=splits)) + rows.extend(check_entry_folds_(df_seq=df_seq, splits=splits)) + rows.extend(check_group_folds_(groups=groups, splits=splits)) + rows.extend(check_fold_size_(splits=splits)) + rows.extend(check_class_balance_(labels=labels, splits=splits)) + df_audit = pd.DataFrame(rows, columns=ut.COLS_AUDIT_LEAKAGE) + # Worst findings first, so the top of the table is what a reader must act on + order = {s: i for i, s in enumerate(reversed(ut.LIST_SEVERITIES))} + if len(df_audit) > 0: + df_audit = (df_audit.sort_values(by=ut.COL_SEVERITY, key=lambda x: x.map(order), + kind="stable") + .reset_index(drop=True)) + return df_audit diff --git a/docs/source/api.rst b/docs/source/api.rst index 9d1d10eb..4a968370 100755 --- a/docs/source/api.rst +++ b/docs/source/api.rst @@ -109,6 +109,7 @@ Prediction ModelEvaluator ModelEvaluatorPlot bind_groups + audit_leakage .. _protein_engineering_api: diff --git a/docs/source/index/references.rst b/docs/source/index/references.rst index e9b4671f..7f58e56c 100755 --- a/docs/source/index/references.rst +++ b/docs/source/index/references.rst @@ -236,6 +236,10 @@ Sampling Strategies *Cross-validation strategies for data with temporal, spatial, hierarchical, or phylogenetic structure*, `Ecography `__. +.. [Kaufman12] Kaufman *et al.* (2012), + *Leakage in data mining: formulation, detection, and avoidance*, + `ACM Transactions on Knowledge Discovery from Data `__. + .. [LiuDeber99] Liu L.-P., Deber C.M. (1999), *Combining hydrophobicity and helicity: a novel approach to membrane protein structure prediction*, `Bioorganic & Medicinal Chemistry `__. diff --git a/examples/prediction/audit_leakage.ipynb b/examples/prediction/audit_leakage.ipynb new file mode 100644 index 00000000..a37ad93b --- /dev/null +++ b/examples/prediction/audit_leakage.ipynb @@ -0,0 +1,781 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "14348aeb", + "metadata": {}, + "source": [ + "**Audit an evaluation setup for leakage risks.**\n", + "\n", + "Leakage inflates a score without ever raising an error: a duplicated sequence on both sides of a split, two windows of one protein pulled apart, a feature quietly derived from the label. The only symptom is an implausibly good number.\n", + "\n", + "`aa.audit_leakage` runs a set of cheap heuristics over whatever parts of a setup are given and returns a plain table of findings, worst first, with the overall verdict in `df_audit.attrs[\"status\"]`." + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "56a02397", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:38.850501Z", + "iopub.status.busy": "2026-09-18T17:17:38.850435Z", + "iopub.status.idle": "2026-09-18T17:17:40.297017Z", + "shell.execute_reply": "2026-09-18T17:17:40.296752Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "DataFrame shape: (48, 5)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 entryentry_winsequencewindowlabel
1P0P0_w0VPMGHACAETPWMPY...QMPGPSILYTYIQYQVPMGHACAETPW0
2P0P0_w1VPMGHACAETPWMPY...QMPGPSILYTYIQYQPMGHACAETPWM0
3P0P0_w2VPMGHACAETPWMPY...QMPGPSILYTYIQYQMGHACAETPWMP0
4P0P0_w3VPMGHACAETPWMPY...QMPGPSILYTYIQYQGHACAETPWMPY0
5P1P1_w0TQRIVDNRTMIHKLR...VCCQHNEVLVSRFSCTQRIVDNRTMIH1
6P1P1_w1TQRIVDNRTMIHKLR...VCCQHNEVLVSRFSCQRIVDNRTMIHK1
7P1P1_w2TQRIVDNRTMIHKLR...VCCQHNEVLVSRFSCRIVDNRTMIHKL1
8P1P1_w3TQRIVDNRTMIHKLR...VCCQHNEVLVSRFSCIVDNRTMIHKLR1
9P2P2_w0NKYEWCPNVGWQVES...MFSCKGRRRWWEDDRNKYEWCPNVGWQ0
10P2P2_w1NKYEWCPNVGWQVES...MFSCKGRRRWWEDDRKYEWCPNVGWQV0
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "import numpy as np\n", + "import pandas as pd\n", + "from sklearn.model_selection import KFold, GroupKFold, StratifiedKFold\n", + "\n", + "import aaanalysis as aa\n", + "\n", + "aa.options[\"verbose\"] = False\n", + "\n", + "# Twelve proteins, four windows each. The label is a property of the PROTEIN, so the\n", + "# windows of one protein are near-duplicates of each other.\n", + "rng = np.random.default_rng(0)\n", + "AA = list(\"ACDEFGHIKLMNPQRSTVWY\")\n", + "n_prot, n_win = 12, 4\n", + "\n", + "rows = []\n", + "for i in range(n_prot):\n", + " parent = \"\".join(rng.choice(AA, 60))\n", + " for w in range(n_win):\n", + " rows.append({\"entry\": f\"P{i}\", \"entry_win\": f\"P{i}_w{w}\",\n", + " \"sequence\": parent, \"window\": parent[w:w + 12], \"label\": i % 2})\n", + "df_seq = pd.DataFrame(rows)\n", + "aa.display_df(df_seq, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "markdown", + "id": "d136dcbf", + "metadata": {}, + "source": [ + "**`df_seq`: the dataset-level audit, before any split.**\n", + "\n", + "With only `df_seq` the audit inspects the data itself. This is the check to run *before* choosing a split. Here nothing is wrong yet, so the table is empty and the status is `ok`." + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "78b865a0", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.298000Z", + "iopub.status.busy": "2026-09-18T17:17:40.297929Z", + "iopub.status.idle": "2026-09-18T17:17:40.301325Z", + "shell.execute_reply": "2026-09-18T17:17:40.301106Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: ok\n", + "DataFrame shape: (0, 4)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 checkseveritydetailids
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "df_audit = aa.audit_leakage(df_seq)\n", + "\n", + "print(\"status:\", df_audit.attrs[\"status\"])\n", + "aa.display_df(df_audit, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "markdown", + "id": "efc2f278", + "metadata": {}, + "source": [ + "Note that the windows share a parent `sequence` but each `window` is distinct. The audit compares `window` when it is present, so the repeated parent is not mistaken for a pile of duplicate samples. Duplicating an actual window does get flagged:" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "id": "e33a7fa3", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.302235Z", + "iopub.status.busy": "2026-09-18T17:17:40.302158Z", + "iopub.status.idle": "2026-09-18T17:17:40.305887Z", + "shell.execute_reply": "2026-09-18T17:17:40.305693Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: medium\n", + "DataFrame shape: (1, 4)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 checkseveritydetailids
1duplicate_sequencesmedium2 rows share on...ready memorized['P0_w0', 'P1_w1']
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "df_dup = df_seq.copy()\n", + "df_dup.loc[5, \"window\"] = df_dup.loc[0, \"window\"] # P1_w1 is now a copy of P0_w0\n", + "\n", + "df_audit = aa.audit_leakage(df_dup)\n", + "print(\"status:\", df_audit.attrs[\"status\"])\n", + "aa.display_df(df_audit, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "markdown", + "id": "9c19d83e", + "metadata": {}, + "source": [ + "**`splits` and `labels`: the fold-level audit, after the split is built.**\n", + "\n", + "`splits` takes the folds as integer row positions, exactly what a scikit-learn splitter yields, so `cv.split(...)` can be passed straight in. `labels` adds the class-balance check; it defaults to `df_seq['label']` when that column exists.\n", + "\n", + "A plain `KFold` scatters the windows of each protein across the folds, which is the classic windowed-data leak." + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "484721c6", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.306728Z", + "iopub.status.busy": "2026-09-18T17:17:40.306668Z", + "iopub.status.idle": "2026-09-18T17:17:40.310896Z", + "shell.execute_reply": "2026-09-18T17:17:40.310700Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: high\n", + "DataFrame shape: (1, 4)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 checkseveritydetailids
1same_protein_across_foldshigh12 protein(s) h...cross the split['P0', 'P1', 'P10', 'P11', 'P2', 'P3', 'P4', 'P5', 'P6', 'P7']
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "labels = df_seq[\"label\"].to_numpy()\n", + "splits_leaky = list(KFold(n_splits=4, shuffle=True, random_state=0).split(df_seq))\n", + "\n", + "df_audit = aa.audit_leakage(df_seq, labels=labels, splits=splits_leaky)\n", + "print(\"status:\", df_audit.attrs[\"status\"])\n", + "aa.display_df(df_audit, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "4c874f7d", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.311690Z", + "iopub.status.busy": "2026-09-18T17:17:40.311631Z", + "iopub.status.idle": "2026-09-18T17:17:40.313275Z", + "shell.execute_reply": "2026-09-18T17:17:40.313080Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "[high] same_protein_across_folds\n", + " 12 protein(s) have windows in both the training and the test part of fold(s) 0, 1, 2, 3, so neighbouring windows leak across the split\n", + " ids: ['P0', 'P1', 'P10', 'P11', 'P2', 'P3', 'P4', 'P5', 'P6', 'P7']\n", + "\n" + ] + } + ], + "source": [ + "# The full sentence of the first finding, which is where the reasoning lives\n", + "for _, row in df_audit.iterrows():\n", + " print(f\"[{row['severity']}] {row['check']}\\n {row['detail']}\\n ids: {row['ids']}\\n\")" + ] + }, + { + "cell_type": "markdown", + "id": "3f1e8241", + "metadata": {}, + "source": [ + "**`groups`: the same audit in the group vocabulary.**\n", + "\n", + "`groups` is one label per sample, exactly as passed to `aa.bind_groups`: an accession, a family, or an externally computed homology cluster. Binding those groups to a group-aware splitter is the fix, and the audit confirms the leak is gone." + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "id": "5aeef14b", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.314022Z", + "iopub.status.busy": "2026-09-18T17:17:40.313966Z", + "iopub.status.idle": "2026-09-18T17:17:40.318276Z", + "shell.execute_reply": "2026-09-18T17:17:40.318082Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: high\n", + "DataFrame shape: (1, 4)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 checkseveritydetailids
1class_balance_anomalyhighthe test part o...ially satisfied['0', '1', '2', '3']
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "groups = df_seq[\"entry\"].to_numpy()\n", + "\n", + "cv = aa.bind_groups(GroupKFold(n_splits=4), groups=groups)\n", + "splits_clean = list(cv.split(df_seq))\n", + "\n", + "df_audit = aa.audit_leakage(df_seq, labels=labels, groups=groups, splits=splits_clean)\n", + "print(\"status:\", df_audit.attrs[\"status\"])\n", + "aa.display_df(df_audit, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "markdown", + "id": "5fe7b5ae", + "metadata": {}, + "source": [ + "The protein and group leaks are gone. What remains is a *true* property of this setup: grouping by protein while the label is a property of the protein leaves some folds single-class. That is worth knowing before trusting a score, which is exactly the point." + ] + }, + { + "cell_type": "markdown", + "id": "3df841ce", + "metadata": {}, + "source": [ + "**`X` and `names`: features that track the label too closely.**\n", + "\n", + "`X` is the feature matrix and `names` gives one name per column, e.g. `df_feat['feature'].to_list()`, so a finding names the feature instead of its column position. A feature correlating with the label at `|r| >= 0.95` is what a target-derived column looks like." + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "id": "41b44081", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.319160Z", + "iopub.status.busy": "2026-09-18T17:17:40.319104Z", + "iopub.status.idle": "2026-09-18T17:17:40.322553Z", + "shell.execute_reply": "2026-09-18T17:17:40.322353Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: high\n", + "DataFrame shape: (1, 4)\n" + ] + }, + { + "data": { + "text/html": [ + "\n", + "\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
 checkseveritydetailids
1target_derived_featurehigh1 feature(s) co...thout the label['leaky_feature']
\n" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "X = rng.random((len(df_seq), 4))\n", + "X[:, 2] = labels + rng.random(len(df_seq)) * 0.01 # quietly derived from the target\n", + "names = [\"hydrophobicity\", \"charge\", \"leaky_feature\", \"helix_propensity\"]\n", + "\n", + "df_audit = aa.audit_leakage(df_seq, labels=labels, X=X, names=names)\n", + "print(\"status:\", df_audit.attrs[\"status\"])\n", + "aa.display_df(df_audit, n_rows=10, show_shape=True)" + ] + }, + { + "cell_type": "markdown", + "id": "7d8a5e44", + "metadata": {}, + "source": [ + "**`raise_on`: turn a report into a hard failure.**\n", + "\n", + "By default the audit only reports, so it can be dropped into a workflow without changing its control flow. `raise_on` sets the severity at or above which a finding raises a `ValueError` instead: it fires on a high-severity finding, and stays quiet when the worst finding sits below the chosen level." + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "id": "0e3e8d5f", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.323351Z", + "iopub.status.busy": "2026-09-18T17:17:40.323283Z", + "iopub.status.idle": "2026-09-18T17:17:40.326034Z", + "shell.execute_reply": "2026-09-18T17:17:40.325845Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "ValueError: 'raise_on' ('high') matched 2 leakage finding(s) at or above that severity: same_protein_across_folds (high): 12 protein(s) have windows in both the training and the test part of fold(s) 0, 1, 2, 3, so neighbouring windows leak across the split; group_overlap_across_folds (high): 12 group(s) appear ...\n" + ] + } + ], + "source": [ + "try:\n", + " aa.audit_leakage(df_seq, labels=labels, groups=groups, splits=splits_leaky,\n", + " raise_on=\"high\")\n", + "except ValueError as e:\n", + " print(\"ValueError:\", str(e)[:300], \"...\")" + ] + }, + { + "cell_type": "code", + "execution_count": 9, + "id": "2552cfaf", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-18T17:17:40.326859Z", + "iopub.status.busy": "2026-09-18T17:17:40.326792Z", + "iopub.status.idle": "2026-09-18T17:17:40.329425Z", + "shell.execute_reply": "2026-09-18T17:17:40.329217Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "status: medium -> no exception raised\n", + "clean status: ok\n" + ] + } + ], + "source": [ + "# Only a medium finding here, so raise_on=\"high\" stays quiet and just returns the table\n", + "df_audit = aa.audit_leakage(df_dup, raise_on=\"high\")\n", + "print(\"status:\", df_audit.attrs[\"status\"], \"-> no exception raised\")\n", + "\n", + "# A clean setup never raises, whatever the threshold\n", + "df_clean = df_seq.drop(columns=[\"sequence\"])\n", + "print(\"clean status:\", aa.audit_leakage(df_clean, raise_on=\"low\").attrs[\"status\"])" + ] + }, + { + "cell_type": "markdown", + "id": "b9692acb", + "metadata": {}, + "source": [ + "**The audit is heuristic, and a clean report is not a proof that no leakage exists.**\n", + "\n", + "It inspects only the data and the folds it is given. It cannot see how a feature was built, whether a scaler was fitted before the split, or that two sequences are near-identical rather than byte-identical. Read `status='ok'` as *nothing obvious*, not as a guarantee. The severities are labels for a human reader, not a machine taxonomy: whether a finding must block a workflow is a decision this package leaves to you." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.11" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/tests/unit/prediction_tests/test_audit_leakage.py b/tests/unit/prediction_tests/test_audit_leakage.py new file mode 100644 index 00000000..32bfa123 --- /dev/null +++ b/tests/unit/prediction_tests/test_audit_leakage.py @@ -0,0 +1,466 @@ +"""This is a script to test the aa.audit_leakage() function.""" +import numpy as np +import pandas as pd +import pytest +from hypothesis import given, settings +import hypothesis.strategies as some +from sklearn.model_selection import KFold, StratifiedKFold, GroupKFold + +import aaanalysis as aa + +settings.register_profile("ci", deadline=None) +settings.load_profile("ci") + +AA = list("ACDEFGHIKLMNPQRSTVWY") +COLS = ["check", "severity", "detail", "ids"] + + +def _seqs(n=40, length=20, seed=0): + """Draw n distinct random sequences.""" + rng = np.random.default_rng(seed) + return ["".join(rng.choice(AA, length)) for _ in range(n)] + + +def _df_seq(n=40, seed=0): + """Clean dataset: one row per protein, all sequences distinct, labels balanced.""" + return pd.DataFrame({"entry": [f"P{i}" for i in range(n)], + "sequence": _seqs(n=n, seed=seed), + "label": [i % 2 for i in range(n)]}) + + +def _splits(n=40, n_splits=4, labels=None): + """Stratified folds over n rows, as a materialized list of index pairs.""" + labels = labels if labels is not None else [i % 2 for i in range(n)] + cv = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=0) + return list(cv.split(np.zeros((n, 1)), labels)) + + +def _df_windows(n_prot=10, n_win=4, seed=0): + """Windowed dataset: several distinct windows per protein, label per protein.""" + rng = np.random.default_rng(seed) + rows = [] + for i in range(n_prot): + parent = "".join(rng.choice(AA, 60)) + for w in range(n_win): + rows.append({"entry": f"P{i}", "entry_win": f"P{i}_w{w}", + "sequence": parent, "window": parent[w:w + 10], + "label": i % 2}) + return pd.DataFrame(rows) + + +def _checks(df_audit): + """Return the set of check names in a findings table.""" + return set(df_audit["check"]) + + +# I Normal cases, one parameter per test +class TestAuditLeakage: + """Test audit_leakage() parameter by parameter.""" + + # df_seq + def test_clean_dataset_is_ok(self): + df_audit = aa.audit_leakage(_df_seq()) + assert len(df_audit) == 0 + assert df_audit.attrs["status"] == "ok" + + def test_returns_plain_dataframe_with_schema(self): + df_audit = aa.audit_leakage(_df_seq()) + assert type(df_audit) is pd.DataFrame + assert list(df_audit.columns) == COLS + + @settings(max_examples=5) + @given(n=some.integers(min_value=4, max_value=30)) + def test_clean_dataset_of_any_size_is_ok(self, n): + assert aa.audit_leakage(_df_seq(n=n)).attrs["status"] == "ok" + + def test_duplicate_sequence_is_flagged(self): + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq) + assert "duplicate_sequences" in _checks(df_audit) + assert df_audit.attrs["status"] == "medium" + + def test_duplicate_sequence_reports_affected_entries(self): + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + row = aa.audit_leakage(df_seq).set_index("check").loc["duplicate_sequences"] + assert set(row["ids"]) == {"P0", "P1"} + + def test_df_seq_without_sequence_column_uses_entry(self): + df_seq = _df_seq()[["entry", "label"]] + assert aa.audit_leakage(df_seq).attrs["status"] == "ok" + + def test_df_seq_none_with_other_input_works(self): + df_audit = aa.audit_leakage(groups=np.repeat(["a", "b"], 5), + splits=[(np.arange(0, 5), np.arange(5, 10))]) + assert df_audit.attrs["status"] == "ok" + + def test_df_seq_is_not_mutated(self): + df_seq = _df_seq() + before = df_seq.copy() + aa.audit_leakage(df_seq, splits=_splits()) + pd.testing.assert_frame_equal(df_seq, before) + + def test_df_seq_empty_raises(self): + with pytest.raises(ValueError, match="at least one sequence"): + aa.audit_leakage(pd.DataFrame({"entry": [], "sequence": []})) + + def test_df_seq_wrong_type_raises(self): + for invalid in ["seq", 42, [1, 2, 3]]: + with pytest.raises(ValueError): + aa.audit_leakage(invalid) + + def test_df_seq_without_usable_columns_raises(self): + with pytest.raises(ValueError, match="should contain at least one of"): + aa.audit_leakage(pd.DataFrame({"foo": [1, 2, 3]})) + + # labels + def test_labels_taken_from_df_seq_label_column(self): + df_seq = _df_seq(n=20) + df_seq["label"] = [0] * 10 + [1] * 10 + df_audit = aa.audit_leakage(df_seq, splits=[(np.arange(0, 10), np.arange(10, 20))]) + assert "class_balance_anomaly" in _checks(df_audit) + + def test_labels_explicit_override_df_seq(self): + df_seq = _df_seq(n=20) + labels = [i % 2 for i in range(20)] + df_audit = aa.audit_leakage(df_seq, labels=labels, splits=_splits(n=20, labels=labels)) + assert "class_balance_anomaly" not in _checks(df_audit) + + def test_labels_single_class_is_accepted_not_raised(self): + df_seq = _df_seq(n=20) + df_audit = aa.audit_leakage(df_seq, labels=[1] * 20) + assert df_audit.attrs["status"] == "ok" + + def test_labels_wrong_length_raises(self): + with pytest.raises(ValueError, match="same samples"): + aa.audit_leakage(_df_seq(n=20), labels=[0, 1, 0]) + + def test_labels_two_dimensional_raises(self): + with pytest.raises(ValueError, match="one-dimensional"): + aa.audit_leakage(_df_seq(n=4), labels=np.zeros((4, 2))) + + # groups + def test_groups_overlap_is_flagged(self): + groups = np.tile([f"P{i}" for i in range(10)], 4) + splits = list(KFold(n_splits=4).split(np.zeros((40, 1)))) + df_audit = aa.audit_leakage(groups=groups, splits=splits) + assert "group_overlap_across_folds" in _checks(df_audit) + assert df_audit.attrs["status"] == "high" + + def test_groups_respected_by_group_splitter_is_ok(self): + groups = np.repeat([f"P{i}" for i in range(10)], 4) + splits = list(GroupKFold(n_splits=4).split(np.zeros((40, 1)), groups=groups)) + df_audit = aa.audit_leakage(groups=groups, splits=splits) + assert "group_overlap_across_folds" not in _checks(df_audit) + + def test_groups_as_list_and_series_agree(self): + groups = list(np.tile(["a", "b", "c", "d", "e"], 4)) + splits = list(KFold(n_splits=4).split(np.zeros((20, 1)))) + one = aa.audit_leakage(groups=groups, splits=splits) + two = aa.audit_leakage(groups=pd.Series(groups), splits=splits) + assert _checks(one) == _checks(two) + + def test_groups_wrong_length_raises(self): + with pytest.raises(ValueError, match="same samples"): + aa.audit_leakage(_df_seq(n=20), groups=["a", "b"]) + + def test_groups_empty_raises(self): + with pytest.raises(ValueError): + aa.audit_leakage(groups=[], splits=[(np.array([0]), np.array([1]))]) + + # splits + def test_splits_none_runs_dataset_checks_only(self): + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq) + assert _checks(df_audit) == {"duplicate_sequences"} + + def test_splits_accepts_a_generator(self): + df_seq = _df_seq() + cv = StratifiedKFold(n_splits=4, shuffle=True, random_state=0) + df_audit = aa.audit_leakage(df_seq, splits=cv.split(df_seq, df_seq["label"])) + assert df_audit.attrs["status"] == "ok" + + @settings(max_examples=5) + @given(n_splits=some.integers(min_value=2, max_value=5)) + def test_splits_any_feasible_fold_count_is_clean(self, n_splits): + df_audit = aa.audit_leakage(_df_seq(), splits=_splits(n_splits=n_splits)) + assert df_audit.attrs["status"] == "ok" + + def test_splits_train_test_row_overlap_is_flagged(self): + splits = [(np.array([0, 1, 2, 3]), np.array([3, 4]))] + df_audit = aa.audit_leakage(_df_seq(n=10), splits=splits) + assert "train_test_overlap" in _checks(df_audit) + assert df_audit.attrs["status"] == "high" + + def test_splits_not_a_pair_raises(self): + with pytest.raises(ValueError, match="pair of two index arrays"): + aa.audit_leakage(_df_seq(n=10), splits=[(np.array([0]), np.array([1]), np.array([2]))]) + + def test_splits_empty_raises(self): + with pytest.raises(ValueError, match="at least one fold"): + aa.audit_leakage(_df_seq(n=10), splits=[]) + + def test_splits_boolean_mask_raises(self): + mask = np.array([True] * 5 + [False] * 5) + with pytest.raises(ValueError, match="integer row positions"): + aa.audit_leakage(_df_seq(n=10), splits=[(mask, ~mask)]) + + def test_splits_out_of_range_raises(self): + with pytest.raises(ValueError, match="row positions within"): + aa.audit_leakage(_df_seq(n=10), splits=[(np.array([0, 1]), np.array([99]))]) + + # X and names + def test_x_target_derived_feature_is_flagged(self): + labels = np.array([i % 2 for i in range(40)]) + X = np.random.default_rng(0).random((40, 4)) + X[:, 2] = labels + df_audit = aa.audit_leakage(X=X, labels=labels) + assert "target_derived_feature" in _checks(df_audit) + assert df_audit.attrs["status"] == "high" + + def test_x_independent_features_are_ok(self): + labels = np.array([i % 2 for i in range(40)]) + X = np.random.default_rng(1).random((40, 4)) + assert aa.audit_leakage(X=X, labels=labels).attrs["status"] == "ok" + + def test_names_appear_in_the_finding(self): + labels = np.array([i % 2 for i in range(40)]) + X = np.random.default_rng(0).random((40, 3)) + X[:, 1] = labels + row = aa.audit_leakage(X=X, labels=labels, names=["a", "leaky", "c"]) + assert row.set_index("check").loc["target_derived_feature", "ids"] == ["leaky"] + + def test_names_default_to_dataframe_columns(self): + labels = np.array([i % 2 for i in range(40)]) + X = pd.DataFrame({"a": np.random.default_rng(0).random(40), "b": labels * 1.0}) + row = aa.audit_leakage(X=X, labels=labels) + assert row.set_index("check").loc["target_derived_feature", "ids"] == ["b"] + + def test_names_wrong_length_raises(self): + labels = np.array([i % 2 for i in range(40)]) + X = np.random.default_rng(0).random((40, 3)) + with pytest.raises(ValueError, match="one name per feature"): + aa.audit_leakage(X=X, labels=labels, names=["a", "b"]) + + def test_x_without_labels_raises(self): + X = np.random.default_rng(0).random((40, 3)) + with pytest.raises(ValueError, match="'labels' should be given together with 'X'"): + aa.audit_leakage(X=X) + + # raise_on + def test_raise_on_none_never_raises(self): + groups = np.tile(["a", "b"], 10) + splits = list(KFold(n_splits=2).split(np.zeros((20, 1)))) + df_audit = aa.audit_leakage(groups=groups, splits=splits) + assert df_audit.attrs["status"] == "high" + + def test_raise_on_high_raises_when_high_finding_exists(self): + groups = np.tile(["a", "b"], 10) + splits = list(KFold(n_splits=2).split(np.zeros((20, 1)))) + with pytest.raises(ValueError, match="group_overlap_across_folds"): + aa.audit_leakage(groups=groups, splits=splits, raise_on="high") + + def test_raise_on_high_does_not_raise_without_high_finding(self): + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq, raise_on="high") + assert df_audit.attrs["status"] == "medium" + + def test_raise_on_medium_raises_on_medium_finding(self): + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + with pytest.raises(ValueError, match="duplicate_sequences"): + aa.audit_leakage(df_seq, raise_on="medium") + + def test_raise_on_never_raises_on_clean_setup(self): + for severity in ["low", "medium", "high"]: + assert aa.audit_leakage(_df_seq(), raise_on=severity).attrs["status"] == "ok" + + def test_raise_on_invalid_option_raises(self): + for invalid in ["critical", "HIGH", "severe", 1]: + with pytest.raises(ValueError): + aa.audit_leakage(_df_seq(), raise_on=invalid) + + # general + def test_no_input_at_all_raises(self): + with pytest.raises(ValueError, match="nothing to audit"): + aa.audit_leakage() + + def test_every_check_name_is_registered(self): + import aaanalysis.utils as ut + df_seq = _df_seq() + df_seq.loc[1, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq, splits=[(np.arange(0, 39), np.array([39, 1]))]) + assert _checks(df_audit).issubset(set(ut.LIST_CHECKS_LEAKAGE)) + + +# II Complex cases +class TestAuditLeakageComplex: + """Test audit_leakage() with interacting parameters and edge cases.""" + + def test_duplicate_sequence_split_across_folds_is_high(self): + """KPI: a duplicated sequence placed in two folds is reported.""" + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + splits = [(np.arange(0, 19), np.array([19]))] + df_audit = aa.audit_leakage(df_seq, splits=splits) + assert "duplicate_sequences_across_folds" in _checks(df_audit) + row = df_audit.set_index("check").loc["duplicate_sequences_across_folds"] + assert row["severity"] == "high" + assert set(row["ids"]) == {"P0", "P19"} + + def test_clean_setup_with_every_input_given_is_ok(self): + """KPI: a clean setup returns an empty table with status ok.""" + df_seq = _df_seq(n=40) + labels = df_seq["label"].to_numpy() + groups = np.array([f"G{i}" for i in range(40)]) + X = np.random.default_rng(3).random((40, 5)) + df_audit = aa.audit_leakage(df_seq, labels=labels, groups=groups, + splits=_splits(n=40, labels=labels), X=X, + names=[f"f{i}" for i in range(5)]) + assert len(df_audit) == 0 + assert df_audit.attrs["status"] == "ok" + assert list(df_audit.columns) == COLS + + def test_windows_of_one_protein_split_apart_is_flagged(self): + df_seq = _df_windows(n_prot=10, n_win=4) + splits = list(KFold(n_splits=4, shuffle=True, random_state=0).split(df_seq)) + df_audit = aa.audit_leakage(df_seq, splits=splits) + assert "same_protein_across_folds" in _checks(df_audit) + + def test_windows_kept_whole_by_group_splitter_is_clean(self): + df_seq = _df_windows(n_prot=12, n_win=4) + groups = df_seq["entry"].to_numpy() + splits = list(GroupKFold(n_splits=4).split(df_seq, groups=groups)) + df_audit = aa.audit_leakage(df_seq, groups=groups, splits=splits) + assert "same_protein_across_folds" not in _checks(df_audit) + assert "group_overlap_across_folds" not in _checks(df_audit) + + def test_windowed_parent_sequence_is_not_called_a_duplicate(self): + """The parent 'sequence' repeats per window; only 'window' may flag duplicates.""" + df_seq = _df_windows(n_prot=10, n_win=4) + assert "duplicate_sequences" not in _checks(aa.audit_leakage(df_seq)) + + def test_repeated_window_is_flagged_as_duplicate(self): + df_seq = _df_windows(n_prot=10, n_win=4) + df_seq.loc[5, "window"] = df_seq.loc[0, "window"] + row = aa.audit_leakage(df_seq).set_index("check").loc["duplicate_sequences"] + assert set(row["ids"]) == {"P0_w0", "P1_w1"} + + def test_findings_are_sorted_worst_first(self): + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq, splits=[(np.arange(0, 19), np.array([19]))]) + order = {"high": 0, "medium": 1, "low": 2} + ranks = [order[s] for s in df_audit["severity"]] + assert ranks == sorted(ranks) + + def test_status_is_the_worst_severity_present(self): + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq, splits=[(np.arange(0, 19), np.array([19]))]) + assert df_audit.attrs["status"] == "high" + assert "high" in set(df_audit["severity"]) + + def test_single_class_test_fold_is_high(self): + df_seq = _df_seq(n=20) + labels = np.array([0] * 10 + [1] * 10) + splits = [(np.arange(0, 10), np.arange(10, 20))] + df_audit = aa.audit_leakage(df_seq, labels=labels, splits=splits) + row = df_audit.set_index("check").loc["class_balance_anomaly"] + assert row["severity"] == "high" + + def test_uneven_fold_sizes_are_low(self): + df_seq = _df_seq(n=30) + labels = df_seq["label"].to_numpy() + splits = [(np.arange(0, 20), np.arange(20, 30)), (np.arange(10, 30), np.arange(0, 2))] + df_audit = aa.audit_leakage(df_seq, labels=labels, splits=splits) + row = df_audit.set_index("check").loc["fold_size_anomaly"] + assert row["severity"] == "low" + + def test_empty_fold_part_is_high(self): + df_seq = _df_seq(n=20) + splits = [(np.arange(0, 20), np.array([], dtype=int))] + df_audit = aa.audit_leakage(df_seq, splits=splits) + row = df_audit.set_index("check").loc["fold_size_anomaly"] + assert row["severity"] == "high" + + def test_bind_groups_output_is_a_valid_splits_input(self): + """A split from bind_groups feeds straight into splits, and shows no group leak.""" + df_seq = _df_windows(n_prot=12, n_win=4) + groups = df_seq["entry"].to_numpy() + cv = aa.bind_groups(GroupKFold(n_splits=4), groups=groups) + df_audit = aa.audit_leakage(df_seq, groups=groups, splits=cv.split(df_seq)) + assert list(df_audit.columns) == COLS + assert _checks(df_audit).isdisjoint({"group_overlap_across_folds", + "same_protein_across_folds", + "train_test_overlap"}) + + def test_plain_kfold_on_windows_leaks_where_bind_groups_does_not(self): + """The audit separates a leaky ungrouped split from the group-aware one.""" + df_seq = _df_windows(n_prot=12, n_win=4) + groups = df_seq["entry"].to_numpy() + leaky = list(KFold(n_splits=4, shuffle=True, random_state=0).split(df_seq)) + clean = list(aa.bind_groups(GroupKFold(n_splits=4), groups=groups).split(df_seq)) + assert "same_protein_across_folds" in _checks(aa.audit_leakage(df_seq, splits=leaky)) + assert "same_protein_across_folds" not in _checks(aa.audit_leakage(df_seq, splits=clean)) + + def test_non_default_index_is_audited_by_position(self): + df_seq = _df_seq(n=20) + df_seq.index = [f"row{i}" for i in range(20)] + df_seq.loc["row19", "sequence"] = df_seq.loc["row0", "sequence"] + df_audit = aa.audit_leakage(df_seq, splits=[(np.arange(0, 19), np.array([19]))]) + assert "duplicate_sequences_across_folds" in _checks(df_audit) + + def test_ids_are_capped_at_ten(self): + df_seq = _df_seq(n=40) + df_seq["sequence"] = df_seq.loc[0, "sequence"] + row = aa.audit_leakage(df_seq).set_index("check").loc["duplicate_sequences"] + assert len(row["ids"]) == 10 + assert "40 rows" in row["detail"] + + def test_mismatched_lengths_across_inputs_raise(self): + with pytest.raises(ValueError, match="same samples"): + aa.audit_leakage(_df_seq(n=20), groups=["a"] * 20, + X=np.random.default_rng(0).random((10, 3)), + labels=[0, 1] * 5) + + def test_splits_out_of_range_for_groups_raises(self): + with pytest.raises(ValueError, match="row positions within"): + aa.audit_leakage(groups=["a", "b", "c"], splits=[(np.array([0]), np.array([7]))]) + + def test_several_independent_findings_are_separate_rows(self): + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + labels = np.array([0] * 10 + [1] * 10) + df_audit = aa.audit_leakage(df_seq, labels=labels, + splits=[(np.arange(0, 10), np.arange(10, 20))]) + assert len(_checks(df_audit)) >= 2 + assert len(df_audit) == len(df_audit.drop_duplicates(subset=["check", "detail"])) + + def test_repeated_calls_are_deterministic(self): + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + splits = [(np.arange(0, 19), np.array([19]))] + one = aa.audit_leakage(df_seq, splits=splits) + two = aa.audit_leakage(df_seq, splits=splits) + pd.testing.assert_frame_equal(one, two) + assert one.attrs["status"] == two.attrs["status"] + + @settings(max_examples=5) + @given(n_prot=some.integers(min_value=4, max_value=10)) + def test_group_splitter_never_flags_group_overlap(self, n_prot): + df_seq = _df_windows(n_prot=n_prot, n_win=3) + groups = df_seq["entry"].to_numpy() + splits = list(GroupKFold(n_splits=2).split(df_seq, groups=groups)) + df_audit = aa.audit_leakage(groups=groups, splits=splits) + assert "group_overlap_across_folds" not in _checks(df_audit) + + def test_status_survives_on_a_filtered_copy_only_via_attrs(self): + df_seq = _df_seq(n=20) + df_seq.loc[19, "sequence"] = df_seq.loc[0, "sequence"] + df_audit = aa.audit_leakage(df_seq) + assert df_audit.attrs["status"] == "medium" + assert "status" not in df_audit.columns