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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion mhctools/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ def __getattr__(name):
raise AttributeError(
"module %r has no attribute %r" % (__name__, name))

__version__ = "3.31.5"
__version__ = "3.31.6"

__all__ = [
"Prediction",
Expand Down
63 changes: 52 additions & 11 deletions mhctools/base_commandline_predictor.py
Original file line number Diff line number Diff line change
Expand Up @@ -371,22 +371,25 @@ def _normalized_supported_alleles(command, supported_allele_flag, raw_names):
_normalized_supported_cache[cache_key] = mapping
return mapping

def _resolve_supported_allele_cli_names(self, alleles):
def _partition_supported_allele_cli_names(self, alleles):
"""
Validate alleles against the predictor's supported-allele list and
return the exact spelling to pass on the command line for each.
Resolve alleles against the predictor's supported-allele list without
raising, returning ``(allele_cli_names, unsupported)``.

Fast path: if prepare_allele_name(allele) is printed verbatim by the
predictor, use it directly — no full-list normalization needed. Slow
path (only on a miss): normalize the whole supported list once and look
the allele up by its normalized name, using the predictor's own
spelling on the command line so names our normalizer rewrites (many
non-human alleles) still round-trip.

A predictor with no known supported-allele list (``raw_names is None``)
supports everything, so ``unsupported`` is always empty for it.
"""
raw_names = self._supported_allele_names
allele_cli_names = {}
if raw_names is None:
return allele_cli_names
return allele_cli_names, []
unsupported = []
normalized_map = None
for allele in _unique_in_order(alleles):
Expand All @@ -412,6 +415,16 @@ def _resolve_supported_allele_cli_names(self, alleles):
allele_cli_names[allele] = original
else:
unsupported.append(allele)
return allele_cli_names, unsupported

def _resolve_supported_allele_cli_names(self, alleles):
"""
Validate alleles against the predictor's supported-allele list and
return the exact spelling to pass on the command line for each. Raise
UnsupportedAllele if any allele is not supported.
"""
allele_cli_names, unsupported = \
self._partition_supported_allele_cli_names(alleles)
if unsupported:
raise UnsupportedAllele(
"Unsupported alleles for %s: %s\n"
Expand Down Expand Up @@ -734,7 +747,7 @@ def _result_lookup_for_allele(results, allele):
preds=existing.preds + matching_preds)
return lookup

def predict_pairs(self, peptides, alleles=None):
def predict_pairs(self, peptides, alleles=None, skip_unsupported=False):
"""Predict explicit peptide/allele pairs.

Parameters
Expand All @@ -744,25 +757,44 @@ def predict_pairs(self, peptides, alleles=None):
iterable of ``(peptide, allele)`` pairs.
alleles : list of str, optional
Alleles parallel to *peptides*.
skip_unsupported : bool
By default one unsupported allele raises UnsupportedAllele and no
pair is scored. Set True to score the supported pairs anyway: pairs
whose allele the predictor does not support are left as ``None`` in
the returned list (its length and order are preserved), and the
skipped alleles are logged at WARNING.

Returns
-------
list of PeptideResult
One entry per input pair, in order.
One entry per input pair, in order. With ``skip_unsupported`` a pair
whose allele is unsupported is ``None`` rather than a PeptideResult.
"""
peptide_list, allele_list = self._check_pair_inputs(peptides, alleles)
if not peptide_list:
return []

unique_alleles = _unique_in_order(allele_list)
allele_cli_names = self._resolve_supported_allele_cli_names(
unique_alleles)
if skip_unsupported:
allele_cli_names, unsupported = \
self._partition_supported_allele_cli_names(unique_alleles)
if unsupported:
logger.warning(
"Skipping %d unsupported allele(s) for %s: %s",
len(unsupported), self.program_name, unsupported)
else:
allele_cli_names = self._resolve_supported_allele_cli_names(
unique_alleles)
unsupported = []
skipped = set(unsupported)
indices_by_allele = defaultdict(list)
for index, allele in enumerate(allele_list):
indices_by_allele[allele].append(index)

results = [None] * len(peptide_list)
for allele in unique_alleles:
if allele in skipped:
continue
indices = indices_by_allele[allele]
group_peptides = _unique_in_order(
peptide_list[index] for index in indices)
Expand All @@ -781,14 +813,23 @@ def predict_pairs(self, peptides, alleles=None):
results[index] = PeptideResult(preds=tuple(peptide_result.preds))
return results

def predict_pairs_dataframe(self, peptides, alleles=None, sample_name=""):
"""``predict_pairs()`` flattened to a DataFrame."""
def predict_pairs_dataframe(
self, peptides, alleles=None, sample_name="",
skip_unsupported=False):
"""``predict_pairs()`` flattened to a DataFrame.

With ``skip_unsupported`` (see :meth:`predict_pairs`) pairs whose allele
the predictor does not support are simply absent from the result rather
than raising, so the DataFrame can be shorter than the input.
"""
import pandas as pd
from .pred import COLUMNS

dfs = [
pp.to_dataframe(sample_name)
for pp in self.predict_pairs(peptides, alleles)]
for pp in self.predict_pairs(
peptides, alleles, skip_unsupported=skip_unsupported)
if pp is not None]
if not dfs:
return pd.DataFrame(columns=COLUMNS)
return pd.concat(dfs, ignore_index=True)
Expand Down
64 changes: 64 additions & 0 deletions tests/test_commandline_predict_pairs.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,3 +162,67 @@ def test_predict_pairs_dataframe_empty_result_has_canonical_columns():

assert list(df.columns) == list(COLUMNS)
assert df.empty


def test_predict_pairs_skip_unsupported_scores_supported_and_nones_the_rest():
predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"])

results = predictor.predict_pairs(
["SIINFEKLL", "GILGFVFTL", "NLVPMVATV"],
["HLA-A*02:01", "HLA-B*07:02", "HLA-A*02:01"],
skip_unsupported=True)

# length and order preserved; the unsupported allele's pair is None
assert len(results) == 3
assert results[1] is None
assert [r.preds[0].peptide for r in results if r is not None] == [
"SIINFEKLL", "NLVPMVATV"]
assert [r.preds[0].allele for r in results if r is not None] == [
"HLA-A*02:01", "HLA-A*02:01"]
# the unsupported allele never reached the command runner
assert predictor.calls == [(["HLA-A*02:01"], ["SIINFEKLL", "NLVPMVATV"])]


def test_predict_pairs_skip_unsupported_warns_which_alleles(caplog):
predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"])

with caplog.at_level("WARNING"):
predictor.predict_pairs(
["SIINFEKLL"], ["HLA-B*07:02"], skip_unsupported=True)

assert "HLA-B*07:02" in caplog.text


def test_predict_pairs_skip_unsupported_all_unsupported_is_all_none():
predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"])

results = predictor.predict_pairs(
["SIINFEKLL", "GILGFVFTL"],
["HLA-B*07:02", "HLA-C*07:02"],
skip_unsupported=True)

assert results == [None, None]
assert predictor.calls == []


def test_predict_pairs_dataframe_skip_unsupported_omits_unsupported_rows():
predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"])

df = predictor.predict_pairs_dataframe(
["SIINFEKLL", "GILGFVFTL", "NLVPMVATV"],
["HLA-A*02:01", "HLA-B*07:02", "HLA-A*02:01"],
skip_unsupported=True)

assert list(df.columns) == list(COLUMNS)
assert df["peptide"].tolist() == ["SIINFEKLL", "NLVPMVATV"]
assert df["allele"].tolist() == ["HLA-A*02:01", "HLA-A*02:01"]


def test_predict_pairs_dataframe_skip_unsupported_all_unsupported_is_empty():
predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"])

df = predictor.predict_pairs_dataframe(
["SIINFEKLL"], ["HLA-B*07:02"], skip_unsupported=True)

assert list(df.columns) == list(COLUMNS)
assert df.empty
Loading