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 ms2deepscore/__version__.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
__version__ = '2.11.0'
__version__ = '2.12.0'
56 changes: 45 additions & 11 deletions ms2deepscore/train_new_model/DataGeneratorEmbeddingEvaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,12 @@

import numpy as np
import pandas as pd
from torch import tensor
import torch
from matchms import Spectrum
from matchms.similarity.vector_similarity_functions import jaccard_similarity_matrix

from ms2deepscore.fingerprint_similarity_computations import (
compute_fingerprint_similarity_matrix,
is_dense_fingerprint_type,
)
from ms2deepscore.SettingsMS2Deepscore import SettingsEmbeddingEvaluator
from ms2deepscore.models import SiameseSpectralModel
from ms2deepscore.tensorize_spectra import tensorize_spectra
Expand Down Expand Up @@ -54,12 +56,17 @@ def __init__(
self.ms2ds_model.to(self.device)
self.indexes = np.arange(len(self.spectrums))
self.batch_size = self.settings.evaluator_distribution_size
self.fingerprint_df = self.compute_fingerprint_dataframe(
self.fingerprints, fingerprint_inchikeys = compute_fingerprints_for_training(
self.spectrums,
fingerprint_type=self.ms2ds_model.model_settings.fingerprint_type,
fingerprint_nbits=self.ms2ds_model.model_settings.fingerprint_nbits,
nbits=self.ms2ds_model.model_settings.fingerprint_nbits,
)

self.fingerprint_index = {
inchikey: idx
for idx, inchikey in enumerate(fingerprint_inchikeys)
}

# Initialize random number generator
self.rng = np.random.default_rng(self.settings.random_seed)

Expand Down Expand Up @@ -87,17 +94,44 @@ def _compute_embeddings_and_scores(self, batch_index: int):
spec_tensors, meta_tensors = tensorize_spectra(
[self.spectrums[i] for i in indexes], self.ms2ds_model.model_settings
)
embeddings = self.ms2ds_model.encoder(spec_tensors.to(self.device), meta_tensors.to(self.device))
with torch.no_grad():
embeddings = self.ms2ds_model.encoder(
spec_tensors.to(self.device),
meta_tensors.to(self.device),
)

ms2ds_scores = cosine_similarity_matrix(embeddings.cpu().detach().numpy(), embeddings.cpu().detach().numpy())
embeddings_numpy = embeddings.cpu().numpy()

ms2ds_scores = cosine_similarity_matrix(embeddings_numpy, embeddings_numpy)

# Compute true scores
inchikeys = [self.inchikey14s[i] for i in indexes]
fingerprints = self.fingerprint_df.loc[inchikeys].to_numpy()
fingerprint_type = self.ms2ds_model.model_settings.fingerprint_type

tanimoto_scores = jaccard_similarity_matrix(fingerprints, fingerprints)
inchikeys = [self.inchikey14s[i] for i in indexes]
fingerprint_indexes = [
self.fingerprint_index[inchikey]
for inchikey in inchikeys
]

if is_dense_fingerprint_type(fingerprint_type):
batch_fingerprints = self.fingerprints[fingerprint_indexes]
else:
batch_fingerprints = [
self.fingerprints[i]
for i in fingerprint_indexes
]

tanimoto_scores = compute_fingerprint_similarity_matrix(
batch_fingerprints,
batch_fingerprints,
fingerprint_type=fingerprint_type,
)

return tensor(tanimoto_scores), tensor(ms2ds_scores), embeddings.cpu().detach()
return (
torch.as_tensor(tanimoto_scores),
torch.as_tensor(ms2ds_scores),
embeddings.cpu(),
)

def on_epoch_end(self):
"""Updates indexes after each epoch."""
Expand Down
48 changes: 48 additions & 0 deletions tests/test_embedding_evaluator.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import os
from ms2deepscore.train_new_model.DataGeneratorEmbeddingEvaluation import DataGeneratorEmbeddingEvaluation
import pytest
import numpy as np
import torch
from sklearn.datasets import make_regression
from torch import randn
from ms2deepscore.models import EmbeddingEvaluationModel, LinearModel
Expand Down Expand Up @@ -131,3 +133,49 @@ def test_linear_model_save_load(tmp_path):
assert np.array_equal(model.model.coef_, loaded_model.model.coef_), "Coefficients do not match."
assert model.model.intercept_ == loaded_model.model.intercept_, "Intercepts do not match."
assert model.degree == loaded_model.degree == 3, "Degree does not match."


@pytest.mark.parametrize(
"fingerprint_type",
[
"rdkit_binary",
"rdkit_count",
"rdkit_logcount",
"rdkit_binary_unfolded",
"rdkit_count_unfolded",
"rdkit_logcount_unfolded",
],
)
def test_embedding_evaluator_generator_supports_fingerprint_types(
fingerprint_type,
):
spectra = create_test_spectra(
num_of_unique_inchikeys=10,
num_of_spectra_per_inchikey=1,
)

model = MockMS2DSModel()
model.model_settings.fingerprint_type = fingerprint_type
model.model_settings.fingerprint_nbits = 256

generator = DataGeneratorEmbeddingEvaluation(
spectrums=spectra,
ms2ds_model=model,
settings=SettingsEmbeddingEvaluator(
evaluator_distribution_size=5,
),
device="cpu",
)

tanimoto_scores, ms2ds_scores, embeddings = next(generator)

assert tanimoto_scores.shape == (5, 5)
assert ms2ds_scores.shape == (5, 5)
assert embeddings.shape[0] == 5

assert torch.isfinite(tanimoto_scores).all()

torch.testing.assert_close(
torch.diag(tanimoto_scores),
torch.ones(5),
)