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