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
122 changes: 95 additions & 27 deletions ms2deepscore/MS2DeepScore.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,24 @@
from typing import List

import numpy as np
from matchms import Spectrum
from matchms.similarity.BaseSimilarity import BaseSimilarity

from ms2deepscore.matchms_compat import (
MATCHMS_V1_API,
as_matchms_scores,
normalize_score_fields,
assert_legacy_symmetric_inputs,
)
if MATCHMS_V1_API:
from matchms.similarity.base_similarity import BaseSimilarity
else:
from matchms.similarity.BaseSimilarity import BaseSimilarity
from ms2deepscore.models.SiameseSpectralModel import SiameseSpectralModel
from .vector_operations import cosine_similarity, cosine_similarity_matrix


class MS2DeepScore(BaseSimilarity):
"""Calculate MS2DeepScore similarity scores between a reference and a query.
"""Calculate MS2DeepScore similarity scores between spectra.

Using a trained model, binned spectrums will be converted into spectrum
vectors using a deep neural network. The MS2DeepScore similarity is then
Expand All @@ -19,13 +30,13 @@ class MS2DeepScore(BaseSimilarity):
.. code-block:: python

from matchms import calculate_scores()
from matchms.importing import load_from_json
from matchms.importing import load_ms2_dataset
from ms2deepscore import MS2DeepScore
from ms2deepscore.models import load_model

# Import data
references = load_from_json("abc.json")
queries = load_from_json("xyz.json")
references = load_ms2_dataset("abc.mgf")
queries = load_ms2_dataset("xyz.mgf")

# Load pretrained model
model = load_model("model_file_123.pt")
Expand All @@ -34,9 +45,12 @@ class MS2DeepScore(BaseSimilarity):
# Calculate scores and get matchms.Scores object
scores = calculate_scores(references, queries, similarity_measure)


This implementation supports both the matchms <=0.33 matrix API (NumPy
return values) and the matchms >=1.0 API (``matchms.Scores`` return values).
"""

score_fields = ("score",)

def __init__(self, model: SiameseSpectralModel, progress_bar: bool = True):
"""

Expand All @@ -54,17 +68,24 @@ def __init__(self, model: SiameseSpectralModel, progress_bar: bool = True):
self.output_vector_dim = self.model.model_settings.embedding_dim
self.progress_bar = progress_bar

def get_embedding_array(self, spectrums, datatype: str = "numpy", batch_size: int = 1024) -> np.ndarray:
"""Calculate the spectrum embeddings for a list of spectrums."""
def get_embedding_array(
self,
spectra,
datatype: str = "numpy",
batch_size: int = 1024,
progress_bar: bool | None = None,
) -> np.ndarray:
"""Calculate embeddings for a collection of spectra."""
show_progress = self.progress_bar if progress_bar is None else progress_bar
return self.model.compute_embedding_array(
spectrums,
spectra,
datatype=datatype,
progress_bar=self.progress_bar,
batch_size=batch_size
)
progress_bar=show_progress,
batch_size=batch_size,
)

def pair(self, reference: Spectrum, query: Spectrum) -> float:
"""Calculate the MS2DeepScore similaritiy between a reference and a query spectrum.
"""Calculate a single MS2DeepScore similarity.

Parameters
----------
Expand All @@ -78,13 +99,17 @@ def pair(self, reference: Spectrum, query: Spectrum) -> float:
ms2ds_similarity
MS2DeepScore similarity score.
"""
embedding_reference = self.get_embedding_array([reference])
embedding_query = self.get_embedding_array([query])
return cosine_similarity(embedding_reference[0, :], embedding_query[0, :])

def matrix(self, references: List[Spectrum], queries: List[Spectrum],
array_type: str = "numpy",
is_symmetric: bool = False) -> np.ndarray:
embeddings = self.get_embedding_array([reference, query])
return cosine_similarity(embeddings[0, :], embeddings[1, :])

def _matrix_numpy(
self,
references: List[Spectrum],
queries: List[Spectrum],
*,
is_symmetric: bool,
progress_bar: bool,
) -> np.ndarray:
"""Calculate the MS2DeepScore similarities between all references and queries.

Parameters
Expand All @@ -106,13 +131,56 @@ def matrix(self, references: List[Spectrum], queries: List[Spectrum],
ms2ds_similarity
Array of MS2DeepScore similarity scores.
"""
embeddings_reference = self.get_embedding_array(references)
embeddings_reference = self.get_embedding_array(
references, progress_bar=progress_bar
)
if is_symmetric:
assert np.all(references == queries), \
"Expected references to be equal to queries for is_symmetric=True"
embeddings_query = embeddings_reference
else:
embeddings_query = self.get_embedding_array(queries)

ms2ds_similarity = cosine_similarity_matrix(embeddings_reference, embeddings_query)
return ms2ds_similarity
embeddings_query = self.get_embedding_array(
queries, progress_bar=progress_bar
)
return cosine_similarity_matrix(embeddings_reference, embeddings_query)

if MATCHMS_V1_API:

def matrix(
self,
spectra_1: List[Spectrum],
spectra_2: List[Spectrum] | None = None,
score_fields=None,
progress_bar: bool = True,
):
"""Return a matchms >=1.0 ``Scores`` object."""
normalize_score_fields(score_fields, self.score_fields)
is_symmetric = spectra_2 is None or spectra_2 is spectra_1
queries = spectra_1 if spectra_2 is None else spectra_2
score_matrix = self._matrix_numpy(
spectra_1,
queries,
is_symmetric=is_symmetric,
progress_bar=progress_bar,
)
return as_matchms_scores({"score": score_matrix})

else:

def matrix(
self,
references: List[Spectrum],
queries: List[Spectrum],
array_type: str = "numpy",
is_symmetric: bool = False,
progress_bar: bool = True,
) -> np.ndarray:
"""Return a NumPy matrix for the matchms <=0.33 API."""
if array_type != "numpy":
raise NotImplementedError("MS2DeepScore currently supports only array_type='numpy'.")
if is_symmetric:
assert_legacy_symmetric_inputs(references, queries)
return self._matrix_numpy(
references,
queries,
is_symmetric=is_symmetric,
progress_bar=progress_bar,
)
Loading