From f577ce5d5f092f981ee2f25aea3f7f5440044010 Mon Sep 17 00:00:00 2001 From: Alex Rubinsteyn Date: Fri, 10 Jul 2026 21:21:22 -0400 Subject: [PATCH] Add command-line predictor pair API --- mhctools/__init__.py | 2 +- mhctools/base_commandline_predictor.py | 274 ++++++++++++++++++------ tests/test_commandline_predict_pairs.py | 164 ++++++++++++++ 3 files changed, 372 insertions(+), 68 deletions(-) create mode 100644 tests/test_commandline_predict_pairs.py diff --git a/mhctools/__init__.py b/mhctools/__init__.py index 6a54314..f9fa6cd 100644 --- a/mhctools/__init__.py +++ b/mhctools/__init__.py @@ -88,7 +88,7 @@ def __getattr__(name): raise AttributeError( "module %r has no attribute %r" % (__name__, name)) -__version__ = "3.31.4" +__version__ = "3.31.5" __all__ = [ "Prediction", diff --git a/mhctools/base_commandline_predictor.py b/mhctools/base_commandline_predictor.py index 410cef5..49e008f 100644 --- a/mhctools/base_commandline_predictor.py +++ b/mhctools/base_commandline_predictor.py @@ -19,7 +19,11 @@ import tempfile from typechecks import require_string, require_integer, require_iterable_of -from .allele_normalization import normalize_allele_name, AlleleParseError +from .allele_normalization import ( + normalize_allele_name, + normalize_allele_name_or_raw, + AlleleParseError, +) from .base_predictor import BasePredictor from .unsupported_allele import UnsupportedAllele @@ -61,6 +65,10 @@ _normalized_supported_cache = {} +def _unique_in_order(values): + return list(dict.fromkeys(values)) + + def _prefer_allele_spelling(new, existing): """ When several raw -listMHC spellings normalize to the same allele, pick the @@ -363,11 +371,10 @@ def _normalized_supported_alleles(command, supported_allele_flag, raw_names): _normalized_supported_cache[cache_key] = mapping return mapping - def _resolve_supported_alleles(self): + def _resolve_supported_allele_cli_names(self, alleles): """ - Validate self.alleles against the predictor's supported-allele list and - record, for each, the exact spelling to pass on the command line - (self._allele_cli_names). + Validate alleles against the predictor's supported-allele list and + return the exact spelling to pass on the command line for each. Fast path: if prepare_allele_name(allele) is printed verbatim by the predictor, use it directly — no full-list normalization needed. Slow @@ -377,12 +384,12 @@ def _resolve_supported_alleles(self): non-human alleles) still round-trip. """ raw_names = self._supported_allele_names - self._allele_cli_names = {} + allele_cli_names = {} if raw_names is None: - return + return allele_cli_names unsupported = [] normalized_map = None - for allele in self.alleles: + for allele in _unique_in_order(alleles): # Fast path: the predictor's expected spelling is printed verbatim. # prepare_allele_name can itself reject an allele (e.g. netMHCIIpan # re-parses and raises on unexpected genes); if so, fall through to @@ -393,7 +400,7 @@ def _resolve_supported_alleles(self): except Exception: cli_name = None if cli_name is not None and cli_name in raw_names: - self._allele_cli_names[allele] = cli_name + allele_cli_names[allele] = cli_name continue if normalized_map is None: normalized_map = self._normalized_supported_alleles( @@ -402,7 +409,7 @@ def _resolve_supported_alleles(self): raw_names) original = normalized_map.get(allele) if original is not None: - self._allele_cli_names[allele] = original + allele_cli_names[allele] = original else: unsupported.append(allele) if unsupported: @@ -413,6 +420,14 @@ def _resolve_supported_alleles(self): unsupported, self.program_name, self.supported_alleles_flag)) + return allele_cli_names + + def _resolve_supported_alleles(self): + """ + Validate self.alleles and record command-line spellings for them. + """ + self._allele_cli_names = self._resolve_supported_allele_cli_names( + self.alleles) def prepare_allele_name(self, allele_name): """ @@ -420,7 +435,7 @@ def prepare_allele_name(self, allele_name): """ return allele_name.replace("*", "") - def _cli_allele_name(self, allele_name): + def _cli_allele_name(self, allele_name, allele_cli_names=None): """ The spelling to pass on the command line for a (normalized) allele. @@ -429,6 +444,10 @@ def _cli_allele_name(self, allele_name): rewritten to a rejected form); falls back to prepare_allele_name for predictors that don't enumerate their supported alleles. """ + if allele_cli_names: + resolved = allele_cli_names.get(allele_name) + if resolved is not None: + return resolved resolved = getattr(self, "_allele_cli_names", {}).get(allele_name) if resolved is not None: return resolved @@ -454,16 +473,18 @@ def _auto_allele_group_size(self, n_alleles, n_input_files): group_size = max(1, -(-n_alleles // groups_per_file)) return min(group_size, AUTO_MAX_ALLELES_PER_COMMAND) - def _allele_groups(self, n_input_files=1): + def _allele_groups(self, n_input_files=1, alleles=None): """ - Partition self.alleles into groups, one command (one predictor + Partition alleles into groups, one command (one predictor process) per group per input file. Group size is controlled by max_alleles_per_command (see __init__). Returns a list of lists of allele names. Empty if there are no alleles. """ - alleles = list(self.alleles) + if alleles is None: + alleles = self.alleles + alleles = list(alleles) if not alleles: return [] n = self.max_alleles_per_command @@ -482,7 +503,8 @@ def _build_command( alleles, length=None, temp_dirname=None, - peptide_mode=False): + peptide_mode=False, + allele_cli_names=None): # accept either a single allele name or a list of them if isinstance(alleles, str): alleles = [alleles] @@ -490,7 +512,8 @@ def _build_command( if peptide_mode: args.extend(self.peptide_mode_flags) allele_arg = ",".join( - self._cli_allele_name(allele) for allele in alleles) + self._cli_allele_name(allele, allele_cli_names) + for allele in alleles) args.extend([self.allele_flag, allele_arg]) if length: args.extend([self.length_flag, str(length)]) @@ -580,54 +603,11 @@ def _run_commands_and_collect_preds( groups[key].append(pred) return [PeptideResult(preds=tuple(preds)) for preds in groups.values()] - def predict(self, peptides, n_flanks=None, c_flanks=None): - """ - Predict for a list of peptide sequences. - - Returns list of PeptideResult. When a native parse_to_preds_fn is - available, parses directly to Pred objects. Otherwise falls back - to converting from BindingPrediction. - """ - peptides, n_flank_list, c_flank_list = self._check_flank_inputs( - peptides, n_flanks, c_flanks) - if self.parse_to_preds_fn is None: - return super().predict( - peptides, n_flanks=n_flank_list, c_flanks=c_flank_list) - - self._check_peptide_inputs(peptides) - input_filenames = create_input_peptides_files( + def _build_peptide_commands( + self, peptides, - max_peptides_per_file=self.max_peptides_per_file, - group_by_length=self.group_peptides_by_length) - commands = {} - dirs = [] - - for i, input_filename in enumerate(input_filenames): - for j, allele_group in enumerate( - self._allele_groups(n_input_files=len(input_filenames))): - if self.tempdir_flag: - temp_dirname = tempfile.mkdtemp( - prefix="tmp_%d_%d_%s" % (i, j, self.program_name), - suffix="XXXXXX") - dirs.append(temp_dirname) - else: - temp_dirname = None - output_file = tempfile.NamedTemporaryFile( - "w+", - prefix="%s_output_length_%d_%d" % ( - self.program_name, i, j), - delete=False) - commands[output_file] = self._build_command( - input_filename=input_filename, - alleles=allele_group, - peptide_mode=True, - temp_dirname=temp_dirname) - return self._run_commands_and_collect_preds( - commands=commands, - input_filenames=input_filenames, - temp_dir_list=dirs) - - def predict_peptides(self, peptides): + alleles, + allele_cli_names=None): self._check_peptide_inputs(peptides) input_filenames = create_input_peptides_files( peptides, @@ -639,14 +619,16 @@ def predict_peptides(self, peptides): for i, input_filename in enumerate(input_filenames): for j, allele_group in enumerate( - self._allele_groups(n_input_files=len(input_filenames))): + self._allele_groups( + n_input_files=len(input_filenames), + alleles=alleles)): if self.tempdir_flag: temp_dirname = tempfile.mkdtemp( prefix="tmp_%d_%d_%s" % ( i, j, self.program_name), - suffix="XXXXXX") + suffix="XXXXXX") logger.debug( "Created temporary directory %s for alleles %s", temp_dirname, @@ -663,7 +645,19 @@ def predict_peptides(self, peptides): input_filename=input_filename, alleles=allele_group, peptide_mode=True, - temp_dirname=temp_dirname) + temp_dirname=temp_dirname, + allele_cli_names=allele_cli_names) + return commands, input_filenames, dirs + + def _predict_binding_predictions_for_alleles( + self, + peptides, + alleles, + allele_cli_names=None): + commands, input_filenames, dirs = self._build_peptide_commands( + peptides=peptides, + alleles=alleles, + allele_cli_names=allele_cli_names) results = self._run_commands_and_collect_predictions( commands=commands, input_filenames=input_filenames, @@ -671,5 +665,151 @@ def predict_peptides(self, peptides): self._check_results( results, peptides=peptides, - alleles=self.alleles) + alleles=alleles) + return results + + def _predict_for_alleles(self, peptides, alleles, allele_cli_names=None): + if self.parse_to_preds_fn is None: + collection = self._predict_binding_predictions_for_alleles( + peptides=peptides, + alleles=alleles, + allele_cli_names=allele_cli_names) + return collection.to_peptide_preds(kind=self._default_pred_kind()) + + commands, input_filenames, dirs = self._build_peptide_commands( + peptides=peptides, + alleles=alleles, + allele_cli_names=allele_cli_names) + return self._run_commands_and_collect_preds( + commands=commands, + input_filenames=input_filenames, + temp_dir_list=dirs) + + def _check_pair_inputs(self, peptides, alleles=None): + if alleles is None: + pairs = list(peptides) + peptide_list = [] + allele_list = [] + for pair in pairs: + try: + peptide, allele = pair + except (TypeError, ValueError): + raise ValueError( + "Expected (peptide, allele) pairs, got %r" % (pair,)) + peptide_list.append(peptide) + allele_list.append(allele) + else: + peptide_list = list(peptides) + if isinstance(alleles, str): + allele_list = [alleles] + else: + allele_list = list(alleles) + if len(peptide_list) != len(allele_list): + raise ValueError( + "peptides length %d != alleles length %d" + % (len(peptide_list), len(allele_list))) + + self._check_peptide_inputs(peptide_list) + require_iterable_of(allele_list, str, "HLA alleles") + allele_list = [ + normalize_allele_name_or_raw(allele) + for allele in allele_list] + return peptide_list, allele_list + + @staticmethod + def _result_lookup_for_allele(results, allele): + lookup = {} + for peptide_result in results: + matching_preds = tuple( + pred for pred in peptide_result.preds + if pred.allele == allele) + if not matching_preds: + continue + peptide = matching_preds[0].peptide + existing = lookup.get(peptide) + if existing is None: + lookup[peptide] = PeptideResult(preds=matching_preds) + else: + lookup[peptide] = PeptideResult( + preds=existing.preds + matching_preds) + return lookup + + def predict_pairs(self, peptides, alleles=None): + """Predict explicit peptide/allele pairs. + + Parameters + ---------- + peptides : list of str or iterable of (str, str) + Peptides to score. If *alleles* is omitted, this must be an + iterable of ``(peptide, allele)`` pairs. + alleles : list of str, optional + Alleles parallel to *peptides*. + + Returns + ------- + list of PeptideResult + One entry per input pair, in order. + """ + 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) + 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: + indices = indices_by_allele[allele] + group_peptides = _unique_in_order( + peptide_list[index] for index in indices) + group_results = self._predict_for_alleles( + peptides=group_peptides, + alleles=[allele], + allele_cli_names=allele_cli_names) + lookup = self._result_lookup_for_allele(group_results, allele) + for index in indices: + peptide = peptide_list[index] + peptide_result = lookup.get(peptide) + if peptide_result is None: + raise ValueError( + "Missing predictions, example peptide='%s' allele='%s'" + % (peptide, allele)) + 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.""" + import pandas as pd + from .pred import COLUMNS + + dfs = [ + pp.to_dataframe(sample_name) + for pp in self.predict_pairs(peptides, alleles)] + if not dfs: + return pd.DataFrame(columns=COLUMNS) + return pd.concat(dfs, ignore_index=True) + + def predict(self, peptides, n_flanks=None, c_flanks=None): + """ + Predict for a list of peptide sequences. + + Returns list of PeptideResult. When a native parse_to_preds_fn is + available, parses directly to Pred objects. Otherwise falls back + to converting from BindingPrediction. + """ + peptides, _, _ = self._check_flank_inputs( + peptides, n_flanks, c_flanks) + return self._predict_for_alleles( + peptides=peptides, + alleles=self.alleles, + allele_cli_names=getattr(self, "_allele_cli_names", None)) + + def predict_peptides(self, peptides): + return self._predict_binding_predictions_for_alleles( + peptides=peptides, + alleles=self.alleles, + allele_cli_names=getattr(self, "_allele_cli_names", None)) diff --git a/tests/test_commandline_predict_pairs.py b/tests/test_commandline_predict_pairs.py new file mode 100644 index 0000000..1ecdb74 --- /dev/null +++ b/tests/test_commandline_predict_pairs.py @@ -0,0 +1,164 @@ +import os + +import pytest + +from mhctools.base_commandline_predictor import BaseCommandlinePredictor +from mhctools.binding_prediction import BindingPrediction +from mhctools.binding_prediction_collection import BindingPredictionCollection +from mhctools.pred import COLUMNS, Kind +from mhctools.unsupported_allele import UnsupportedAllele + + +class _PairStubPredictor(BaseCommandlinePredictor): + def __init__(self, supported_alleles=None): + if supported_alleles is None: + supported_alleles = ["HLA-A*02:01", "HLA-B*07:02"] + self.program_name = "fakepan" + self.peptide_mode_flags = ["-p"] + self.allele_flag = "-a" + self.length_flag = "-l" + self.input_file_flag = "-f" + self.supported_alleles_flag = "-listMHC" + self.tempdir_flag = None + self.extra_flags = [] + self.alleles = ["HLA-A*02:01"] + self._supported_allele_names = set(supported_alleles) + self._allele_cli_names = {"HLA-A*02:01": "HLA-A*02:01"} + self.max_alleles_per_command = 1 + self.max_peptides_per_file = 10 ** 4 + self.process_limit = -1 + self.group_peptides_by_length = False + self.default_peptide_lengths = [9] + self.allow_X_in_peptides = False + self.allow_lowercase_in_peptides = False + self.min_peptide_length = 8 + self.max_peptide_length = None + self.parse_output_fn = None + self.parse_to_preds_fn = None + self.calls = [] + + def prepare_allele_name(self, allele_name): + return allele_name + + def _run_commands_and_collect_predictions( + self, commands, input_filenames, temp_dir_list, + sequence_key_mapping=None): + predictions = [] + for output_file, command in commands.items(): + allele_arg = command[command.index("-a") + 1] + alleles = allele_arg.split(",") + input_filename = command[-1] + with open(input_filename) as input_file: + peptides = [ + line.strip() + for line in input_file + if line.strip()] + self.calls.append((alleles, peptides)) + for allele in alleles: + for peptide in peptides: + predictions.append(BindingPrediction( + peptide=peptide, + allele=allele, + affinity=float(len(peptide) + len(allele)), + prediction_method_name=self.program_name)) + output_file.close() + _silent_remove(output_file.name) + for input_filename in input_filenames: + _silent_remove(input_filename) + return BindingPredictionCollection(predictions) + + +def _silent_remove(path): + try: + os.remove(path) + except OSError: + pass + + +def test_predict_pairs_groups_by_allele_and_preserves_input_order(): + predictor = _PairStubPredictor() + + results = predictor.predict_pairs( + ["SIINFEKLL", "GILGFVFTL", "NLVPMVATV", "SIINFEKLL"], + ["HLA-B*07:02", "HLA-A*02:01", "HLA-B*07:02", "HLA-A*02:01"]) + + assert predictor.calls == [ + (["HLA-B*07:02"], ["SIINFEKLL", "NLVPMVATV"]), + (["HLA-A*02:01"], ["GILGFVFTL", "SIINFEKLL"]), + ] + assert len(results) == 4 + assert [result.preds[0].peptide for result in results] == [ + "SIINFEKLL", "GILGFVFTL", "NLVPMVATV", "SIINFEKLL"] + assert [result.preds[0].allele for result in results] == [ + "HLA-B*07:02", "HLA-A*02:01", "HLA-B*07:02", "HLA-A*02:01"] + assert all(result.preds[0].kind == Kind.pMHC_affinity + for result in results) + assert all(result.preds[0].predictor_name == "fakepan" + for result in results) + + +def test_predict_pairs_accepts_pair_iterable(): + predictor = _PairStubPredictor() + + results = predictor.predict_pairs([ + ("SIINFEKLL", "HLA-B*07:02"), + ("GILGFVFTL", "HLA-A*02:01"), + ]) + + assert [result.preds[0].peptide for result in results] == [ + "SIINFEKLL", "GILGFVFTL"] + assert [result.preds[0].allele for result in results] == [ + "HLA-B*07:02", "HLA-A*02:01"] + + +def test_predict_pairs_rejects_length_mismatch(): + predictor = _PairStubPredictor() + + with pytest.raises(ValueError) as e: + predictor.predict_pairs(["SIINFEKLL"], [ + "HLA-A*02:01", "HLA-B*07:02"]) + + assert "peptides length 1 != alleles length 2" in str(e.value) + assert predictor.calls == [] + + +def test_predict_pairs_rejects_malformed_pair_iterable(): + predictor = _PairStubPredictor() + + with pytest.raises(ValueError, match="Expected .* pairs"): + predictor.predict_pairs(["SIINFEKLL"]) + + assert predictor.calls == [] + + +def test_predict_pairs_dataframe_is_canonical_schema(): + predictor = _PairStubPredictor() + + df = predictor.predict_pairs_dataframe( + ["SIINFEKLL", "GILGFVFTL"], + ["HLA-B*07:02", "HLA-A*02:01"], + sample_name="sample-1") + + assert list(df.columns) == list(COLUMNS) + assert len(df) == 2 + assert df["sample_name"].tolist() == ["sample-1", "sample-1"] + assert df["peptide"].tolist() == ["SIINFEKLL", "GILGFVFTL"] + assert df["allele"].tolist() == ["HLA-B*07:02", "HLA-A*02:01"] + + +def test_predict_pairs_raises_unsupported_allele_before_running_commands(): + predictor = _PairStubPredictor(supported_alleles=["HLA-A*02:01"]) + + with pytest.raises(UnsupportedAllele, match="HLA-B"): + predictor.predict_pairs(["SIINFEKLL"], ["HLA-B*07:02"]) + + assert predictor.calls == [] + + +def test_predict_pairs_dataframe_empty_result_has_canonical_columns(): + predictor = _PairStubPredictor() + + df = predictor.predict_pairs_dataframe([], []) + + assert list(df.columns) == list(COLUMNS) + assert df.empty